From 3b70f8bb2ec37d0bb3c9b787ab9647adf13c41d9 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 28 Sep 2026 21:01:37 -0700 Subject: [PATCH 1/7] feat(tracking): add a reset_tracking command that undoes tracking for a session Splits every multi-detection occurrence of a session back into one occurrence per detection, recomputes each touched occurrence's determination from its own detections in batches, clears chain links inside the session, grouping confirmations and the tracking task's recorded classifications, then refreshes the session and station cached counts. Refuses sessions with identifications unless --force; --dry-run reports the plan without writing. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01C7Xf6VPbwWtTumhjjF15g8 --- .../management/commands/reset_tracking.py | 64 +++++ ami/main/models_future/session_reset.py | 254 ++++++++++++++++++ .../tests/test_session_reset.py | 119 ++++++++ 3 files changed, 437 insertions(+) create mode 100644 ami/main/management/commands/reset_tracking.py create mode 100644 ami/main/models_future/session_reset.py create mode 100644 ami/ml/post_processing/tests/test_session_reset.py diff --git a/ami/main/management/commands/reset_tracking.py b/ami/main/management/commands/reset_tracking.py new file mode 100644 index 000000000..45754651a --- /dev/null +++ b/ami/main/management/commands/reset_tracking.py @@ -0,0 +1,64 @@ +""" +Undo tracking for one or more sessions, so they can be tracked again from scratch. + +Splits every multi-detection occurrence back into one occurrence per detection, clears +the chain links and grouping confirmations, and prints the counts before and after. See +``ami/main/models_future/session_reset.py`` for exactly what changes. Each session is +reset in its own transaction. +""" + +import dataclasses +import logging + +from django.core.management.base import BaseCommand, CommandError + +from ami.main.models import Event +from ami.main.models_future.session_reset import SessionResetRefused, reset_session_tracking + +logger = logging.getLogger(__name__) + + +class Command(BaseCommand): + help = "Split a session's tracked occurrences back into one occurrence per detection." + + def add_arguments(self, parser): + parser.add_argument("--project", type=int, required=True, help="Project the sessions belong to.") + parser.add_argument( + "--event", + type=int, + action="append", + dest="events", + required=True, + help="Session (event) ID to reset. Repeat to reset several.", + ) + parser.add_argument("--dry-run", action="store_true", help="Report what would change and write nothing.") + parser.add_argument( + "--force", + action="store_true", + help="Reset even when occurrences in the session carry human identifications.", + ) + + def handle(self, *args, **options): + project_id: int = options["project"] + event_ids: list[int] = options["events"] + events = {event.pk: event for event in Event.objects.filter(project_id=project_id, pk__in=event_ids)} + missing = [pk for pk in event_ids if pk not in events] + if missing: + raise CommandError(f"Session(s) {missing} not found in project {project_id}") + + refused = [] + for pk in event_ids: + try: + result = reset_session_tracking(events[pk], force=options["force"], dry_run=options["dry_run"]) + except SessionResetRefused as error: + self.stderr.write(str(error)) + refused.append(pk) + continue + summary = {k: v for k, v in dataclasses.asdict(result).items() if k not in ("before", "after")} + logger.info(f"reset_tracking: {summary}") + self.stdout.write(f"Session {pk}{' (dry run)' if result.dry_run else ''}: {summary}") + self.stdout.write(f" before: {dataclasses.asdict(result.before)}") + if result.after is not None: + self.stdout.write(f" after: {dataclasses.asdict(result.after)}") + if refused: + raise CommandError(f"Refused to reset session(s) {refused}; see above.") diff --git a/ami/main/models_future/session_reset.py b/ami/main/models_future/session_reset.py new file mode 100644 index 000000000..261cec441 --- /dev/null +++ b/ami/main/models_future/session_reset.py @@ -0,0 +1,254 @@ +"""Undo tracking for a session: put every detection back in an occurrence of its own. + +Tracking only runs on a session whose occurrences each hold a single detection (see +``tracking_task.event_is_fresh``). Resetting returns a tracked session to that state, +so the same night can be tracked again with other settings, or looked at as it was +before tracking. It is a staff tool: it throws away every grouping in the session, +including groupings a person confirmed. + +After a reset, for every detection of the session that shared an occurrence: + +- The earliest detection (in capture order) keeps the occurrence; each other + detection in the session moves to a new occurrence of its own. Detections of the + same occurrence in other sessions stay where they are. +- Every occurrence the split touched takes its determination from its own + detections' best prediction, the rule ``Occurrence.best_prediction`` applies. An + occurrence with a human identification keeps the determination it has. +- Chain links between two detections of the session are cleared. A link that crosses + into another session records one animal across a regroup boundary, which tracking + never follows, so it is kept. +- Grouping verification is cleared on every occurrence of the session. +- Classifications the tracking task recorded on the session's detections are deleted. + Each is a copy of a prediction on the same detection, made when a merge changed a + determination, so it describes a merge that no longer exists. +""" + +from __future__ import annotations + +import dataclasses +from collections.abc import Iterable + +from django.db import transaction +from django.db.models import Count, Prefetch, Q + +from ami.main.models import ( + Classification, + Detection, + Event, + Identification, + Occurrence, + update_calculated_fields_for_sessions_and_stations, +) +from ami.main.models_future.occurrence import best_prediction_from_prefetch +from ami.main.models_future.track_stats import refresh_track_stats_for_ids +from ami.main.models_future.tracks import _capture_order_key + +# Occurrences whose determination is recomputed per round trip. Each batch reads the +# occurrences, their detections and those detections' classifications: three queries. +DETERMINATION_BATCH_SIZE = 2000 + +# The Algorithm key the tracking task records its determination changes under. +TRACKING_ALGORITHM_KEY = "tracking" + + +class SessionResetRefused(ValueError): + """The session holds work a reset would discard, and the caller did not force it.""" + + +@dataclasses.dataclass +class SessionTrackingCounts: + """How grouped a session is, read through its detections' captures.""" + + occurrences: int + multi_detection_occurrences: int + determinations: int + links: int + grouping_verified: int + + +@dataclasses.dataclass +class SessionResetResult: + event_id: int + dry_run: bool + identifications: int + occurrences_split: int + occurrences_created: int + links_cleared: int + verifications_cleared: int + tracking_classifications_deleted: int + determinations_updated: int + before: SessionTrackingCounts + after: SessionTrackingCounts | None + + +def _session_detections(event: Event): + return Detection.objects.filter(source_image__event=event) + + +def _session_occurrence_ids(event: Event): + """Occurrences reached through the session's captures, as ``event_is_fresh`` finds them.""" + return _session_detections(event).filter(occurrence__isnull=False).order_by().values("occurrence_id") + + +def session_tracking_counts(event: Event) -> SessionTrackingCounts: + """Occurrences, multi-frame occurrences, distinct determinations, links and confirmations. Three queries.""" + occurrences = Occurrence.objects.filter(pk__in=_session_occurrence_ids(event)) + totals = occurrences.aggregate( + occurrences=Count("pk"), + determinations=Count("determination", distinct=True), + grouping_verified=Count("pk", filter=Q(grouping_verified_at__isnull=False)), + ) + multi = occurrences.annotate(_n=Count("detections")).filter(_n__gt=1).count() + links = _session_detections(event).filter(next_detection__source_image__event=event).count() + return SessionTrackingCounts( + occurrences=totals["occurrences"], + multi_detection_occurrences=multi, + determinations=totals["determinations"], + links=links, + grouping_verified=totals["grouping_verified"], + ) + + +def _plan_split(event: Event) -> tuple[list[int], list[int]]: + """The multi-detection occurrences of the session, and the detections that leave them. + + A detection leaves when it is in this session and is not its occurrence's first + detection in capture order. + """ + multi_ids = list( + Occurrence.objects.filter(pk__in=_session_occurrence_ids(event)) + .annotate(_n=Count("detections")) + .filter(_n__gt=1) + .values_list("pk", flat=True) + ) + if not multi_ids: + return [], [] + rows = list( + Detection.objects.filter(occurrence_id__in=multi_ids) + .order_by() + .values("pk", "occurrence_id", "source_image_id", "source_image__timestamp", "source_image__event_id") + ) + by_occurrence: dict[int, list[dict]] = {} + for row in rows: + by_occurrence.setdefault(row["occurrence_id"], []).append(row) + movers: list[int] = [] + for members in by_occurrence.values(): + members.sort(key=_capture_order_key) + movers.extend(row["pk"] for row in members[1:] if row["source_image__event_id"] == event.pk) + return multi_ids, movers + + +def _refresh_determinations(occurrence_ids: Iterable[int]) -> int: + """Set each occurrence's determination to its best prediction, in batches; returns how many changed. + + The batched form of ``update_occurrence_determination`` for occurrences with no + identifications, using the same choice as ``Occurrence.best_prediction``. + """ + ids = list(occurrence_ids) + classifications = Classification.objects.select_related(None).only( + "pk", "detection_id", "algorithm_id", "taxon_id", "score", "terminal" + ) + detections = ( + Detection.objects.select_related(None) + .only("pk", "occurrence_id") + .prefetch_related(Prefetch("classifications", queryset=classifications)) + ) + changed: list[Occurrence] = [] + for start in range(0, len(ids), DETERMINATION_BATCH_SIZE): + batch = ( + Occurrence.objects.select_related(None) + .filter(pk__in=ids[start : start + DETERMINATION_BATCH_SIZE]) + .only("pk", "determination_id", "determination_score") + ) + for occurrence in batch.prefetch_related(Prefetch("detections", queryset=detections)): + best = best_prediction_from_prefetch(occurrence) + if best is None or best.taxon_id is None: + continue + if (occurrence.determination_id, occurrence.determination_score) != (best.taxon_id, best.score): + occurrence.determination_id = best.taxon_id + occurrence.determination_score = best.score + changed.append(occurrence) + Occurrence.objects.bulk_update(changed, ["determination", "determination_score"], batch_size=1000) + return len(changed) + + +def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = False) -> SessionResetResult: + """Return ``event`` to one occurrence per detection, with no links and no confirmations. + + Raises ``SessionResetRefused`` when the session's occurrences carry human + identifications, unless ``force`` is set; forced, each identification stays on the + occurrence that keeps the first detection. With ``dry_run`` the counts are those a + reset would produce and nothing is written. The query count depends on the number of + occurrences only through the determination batches, not per row. + """ + before = session_tracking_counts(event) + occurrence_ids = _session_occurrence_ids(event) + identified = set( + Identification.objects.filter(occurrence_id__in=occurrence_ids).values_list("occurrence_id", flat=True) + ) + identification_count = Identification.objects.filter(occurrence_id__in=identified).count() if identified else 0 + if identification_count and not force: + raise SessionResetRefused( + f"Session {event.pk} has {identification_count} identification(s) on {len(identified)} occurrence(s). " + "Resetting would split occurrences a person identified; pass force to reset anyway." + ) + + multi_ids, movers = _plan_split(event) + links = _session_detections(event).filter(next_detection__source_image__event=event) + tracking_classifications = Classification.objects.filter( + detection__source_image__event=event, algorithm__key=TRACKING_ALGORITHM_KEY + ) + verified = Occurrence.objects.filter(pk__in=occurrence_ids, grouping_verified_at__isnull=False) + + if dry_run: + return SessionResetResult( + event_id=event.pk, + dry_run=True, + identifications=identification_count, + occurrences_split=len(multi_ids), + occurrences_created=len(movers), + links_cleared=before.links, + verifications_cleared=before.grouping_verified, + tracking_classifications_deleted=tracking_classifications.count(), + determinations_updated=0, + before=before, + after=None, + ) + + with transaction.atomic(): + links_cleared = links.update(next_detection=None) + verifications_cleared = verified.update(grouping_verified_at=None, grouping_verified_by=None) + _, deleted_by_model = tracking_classifications.delete() + tracking_deleted = deleted_by_model.get(Classification._meta.label, 0) + + new_occurrences = Occurrence.objects.bulk_create( + [Occurrence(event=event, deployment_id=event.deployment_id, project_id=event.project_id) for _ in movers], + batch_size=1000, + ) + Detection.objects.bulk_update( + [ + Detection(pk=detection_id, occurrence_id=occurrence.pk) + for detection_id, occurrence in zip(movers, new_occurrences) + ], + ["occurrence"], + batch_size=1000, + ) + + touched = [*multi_ids, *(o.pk for o in new_occurrences)] + determinations_updated = _refresh_determinations(pk for pk in touched if pk not in identified) + refresh_track_stats_for_ids(touched) + + update_calculated_fields_for_sessions_and_stations([event.pk], stations_async=False) + return SessionResetResult( + event_id=event.pk, + dry_run=False, + identifications=identification_count, + occurrences_split=len(multi_ids), + occurrences_created=len(new_occurrences), + links_cleared=links_cleared, + verifications_cleared=verifications_cleared, + tracking_classifications_deleted=tracking_deleted, + determinations_updated=determinations_updated, + before=before, + after=session_tracking_counts(event), + ) diff --git a/ami/ml/post_processing/tests/test_session_reset.py b/ami/ml/post_processing/tests/test_session_reset.py new file mode 100644 index 000000000..75886b969 --- /dev/null +++ b/ami/ml/post_processing/tests/test_session_reset.py @@ -0,0 +1,119 @@ +import io +import logging + +from django.core.management import CommandError, call_command +from django.db import connection +from django.test import TestCase +from django.test.utils import CaptureQueriesContext +from django.utils import timezone + +from ami.main.models import Classification, Detection, Identification, Occurrence, Taxon +from ami.main.models_future.session_reset import SessionResetRefused, reset_session_tracking, session_tracking_counts +from ami.main.tests import cachalot_disabled +from ami.ml.models import Algorithm +from ami.ml.post_processing.tracking_task import assign_occurrences_from_detection_chains, event_is_fresh +from ami.tests.fixtures.main import create_captures, create_occurrences, create_taxa, setup_test_project +from ami.users.tests.factories import UserFactory + +logger = logging.getLogger(__name__) + + +class TestResetSessionTracking(TestCase): + """A reset returns a tracked session to one occurrence per detection, so it can be tracked again.""" + + def _build_tracked_session(self, images: int, boxes_per_image: int = 2): + """A session whose detections each predict their own taxon, tracked into one chain per box + position, with one grouping confirmation and the tracking task's recorded classifications.""" + project, deployment = setup_test_project(reuse=False) + create_captures(deployment=deployment, num_nights=1, images_per_night=images, interval_minutes=1) + create_taxa(project) + create_occurrences(deployment=deployment, num=images * boxes_per_image) + event = project.events.get() + captures = list(event.captures.order_by("timestamp")) + taxa = list(Taxon.objects.filter(projects=project).order_by("pk")) + + detections_by_capture = [list(c.detections.order_by("pk")) for c in captures] + for i, detection in enumerate(det for dets in detections_by_capture for det in dets): + Classification.objects.filter(detection=detection).update( + taxon=taxa[i % len(taxa)], score=0.5 + (i % 7) / 20, terminal=True + ) + for position in range(boxes_per_image): + chain = [dets[position] for dets in detections_by_capture] + for current, following in zip(chain, chain[1:]): + current.next_detection = following + current.save(update_fields=["next_detection"]) + + tracking = Algorithm.objects.get_or_create(key="tracking", defaults={"name": "Occurrence Tracking"})[0] + assign_occurrences_from_detection_chains(captures, logger, record_as=tracking) + keeper = Occurrence.objects.filter(event=event).order_by("pk").first() + Occurrence.objects.filter(pk=keeper.pk).update( + grouping_verified_at=timezone.now(), grouping_verified_by=UserFactory() + ) + return project, event + + def test_splits_every_track_and_clears_links_verification_and_tracking_records(self): + project, event = self._build_tracked_session(images=4) + tracked = session_tracking_counts(event) + self.assertEqual((tracked.occurrences, tracked.multi_detection_occurrences), (2, 2)) + self.assertEqual((tracked.links, tracked.grouping_verified), (6, 1)) + + result = reset_session_tracking(event) + + self.assertEqual((result.occurrences_split, result.occurrences_created), (2, 6)) + after = session_tracking_counts(event) + self.assertEqual((after.occurrences, after.multi_detection_occurrences), (8, 0)) + self.assertEqual((after.links, after.grouping_verified), (0, 0)) + self.assertTrue(event_is_fresh(event)[0]) + self.assertFalse(Classification.objects.filter(algorithm__key="tracking").exists()) + for occurrence in Occurrence.objects.filter(detections__source_image__event=event): + best = occurrence.best_prediction + self.assertEqual( + (occurrence.determination_id, occurrence.determination_score), (best.taxon_id, best.score) + ) + + def test_refuses_a_session_with_identifications_unless_forced(self): + project, event = self._build_tracked_session(images=3) + occurrence = Occurrence.objects.filter(event=event).order_by("pk").first() + Identification.objects.create(occurrence=occurrence, user=UserFactory(), taxon=occurrence.determination) + + with self.assertRaises(SessionResetRefused): + reset_session_tracking(event) + self.assertEqual(session_tracking_counts(event).multi_detection_occurrences, 2) + + result = reset_session_tracking(event, force=True) + self.assertEqual(result.after.multi_detection_occurrences, 0) + self.assertEqual(Identification.objects.filter(occurrence=occurrence).count(), 1) + + def test_dry_run_reports_the_plan_and_writes_nothing(self): + project, event = self._build_tracked_session(images=3) + links = dict(Detection.objects.filter(source_image__event=event).values_list("pk", "next_detection_id")) + memberships = dict(Detection.objects.filter(source_image__event=event).values_list("pk", "occurrence_id")) + before = session_tracking_counts(event) + + out = io.StringIO() + call_command("reset_tracking", project=project.pk, events=[event.pk], dry_run=True, stdout=out) + + self.assertIn("'occurrences_created': 4", out.getvalue()) + self.assertEqual(session_tracking_counts(event), before) + self.assertEqual( + dict(Detection.objects.filter(source_image__event=event).values_list("pk", "next_detection_id")), links + ) + self.assertEqual( + dict(Detection.objects.filter(source_image__event=event).values_list("pk", "occurrence_id")), + memberships, + ) + + def test_command_rejects_a_session_from_another_project(self): + project, event = self._build_tracked_session(images=2) + with self.assertRaises(CommandError): + call_command("reset_tracking", project=project.pk + 1000, events=[event.pk], stdout=io.StringIO()) + + def test_query_count_does_not_grow_with_the_session(self): + """Doubling the detections must not add queries: every write is a bulk statement.""" + counts = [] + for images in (3, 6): + project, event = self._build_tracked_session(images=images) + with cachalot_disabled(), CaptureQueriesContext(connection) as context: + reset_session_tracking(event) + counts.append(len(context.captured_queries)) + self.assertEqual(counts[0], counts[1], counts) From 54c2ef3c6afee6c95dca6bef0b343c6fbc033fc9 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 28 Sep 2026 21:23:53 -0700 Subject: [PATCH 2/7] fix(tracking): leave a track that reaches into another session with that session on reset Resetting a session kept the earliest detection of every multi-detection occurrence. When that detection was in the session being reset and the occurrence also held detections of a later session, the occurrence still spanned more than one detection afterwards, so the session did not pass the tracking task's freshness check and could not be tracked again. The occurrence also stayed filed under the reset session even when none of its detections remained there. An occurrence that holds detections of another session now stays whole in that session: every detection of the reset session leaves it, and if it was filed under the reset session it moves to the session of its first remaining detection. A new test covers an occurrence spanning two sessions and checks that the reset session is fresh and the other session's detections are unchanged. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01C7Xf6VPbwWtTumhjjF15g8 --- ami/main/models_future/session_reset.py | 50 +++++++++++++------ .../tests/test_session_reset.py | 50 ++++++++++++++++--- 2 files changed, 76 insertions(+), 24 deletions(-) diff --git a/ami/main/models_future/session_reset.py b/ami/main/models_future/session_reset.py index 261cec441..efc71f598 100644 --- a/ami/main/models_future/session_reset.py +++ b/ami/main/models_future/session_reset.py @@ -9,8 +9,9 @@ After a reset, for every detection of the session that shared an occurrence: - The earliest detection (in capture order) keeps the occurrence; each other - detection in the session moves to a new occurrence of its own. Detections of the - same occurrence in other sessions stay where they are. + detection in the session moves to a new occurrence of its own. An occurrence that + also holds detections of another session stays with those detections, untouched + there, and every detection of this session leaves it. - Every occurrence the split touched takes its determination from its own detections' best prediction, the rule ``Occurrence.best_prediction`` applies. An occurrence with a human identification keeps the determination it has. @@ -109,22 +110,26 @@ def session_tracking_counts(event: Event) -> SessionTrackingCounts: ) -def _plan_split(event: Event) -> tuple[list[int], list[int]]: - """The multi-detection occurrences of the session, and the detections that leave them. +def _plan_split(event: Event) -> tuple[list[int], list[int], dict[int, int]]: + """The multi-detection occurrences of the session, the detections that leave them, and + the occurrences whose session changes. - A detection leaves when it is in this session and is not its occurrence's first - detection in capture order. + An occurrence that lies wholly in this session keeps its first detection in capture + order, and its other detections leave. An occurrence that also holds detections of + another session stays whole there: every detection of this session leaves it, and if + it was filed under this session it moves to the session of its first remaining + detection. """ - multi_ids = list( + multi = dict( Occurrence.objects.filter(pk__in=_session_occurrence_ids(event)) .annotate(_n=Count("detections")) .filter(_n__gt=1) - .values_list("pk", flat=True) + .values_list("pk", "event_id") ) - if not multi_ids: - return [], [] + if not multi: + return [], [], {} rows = list( - Detection.objects.filter(occurrence_id__in=multi_ids) + Detection.objects.filter(occurrence_id__in=multi) .order_by() .values("pk", "occurrence_id", "source_image_id", "source_image__timestamp", "source_image__event_id") ) @@ -132,10 +137,17 @@ def _plan_split(event: Event) -> tuple[list[int], list[int]]: for row in rows: by_occurrence.setdefault(row["occurrence_id"], []).append(row) movers: list[int] = [] - for members in by_occurrence.values(): + new_sessions: dict[int, int] = {} + for occurrence_id, members in by_occurrence.items(): members.sort(key=_capture_order_key) - movers.extend(row["pk"] for row in members[1:] if row["source_image__event_id"] == event.pk) - return multi_ids, movers + outside = [row for row in members if row["source_image__event_id"] != event.pk] + if outside: + movers.extend(row["pk"] for row in members if row["source_image__event_id"] == event.pk) + if multi[occurrence_id] == event.pk: + new_sessions[occurrence_id] = outside[0]["source_image__event_id"] + else: + movers.extend(row["pk"] for row in members[1:]) + return list(multi), movers, new_sessions def _refresh_determinations(occurrence_ids: Iterable[int]) -> int: @@ -193,7 +205,7 @@ def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = "Resetting would split occurrences a person identified; pass force to reset anyway." ) - multi_ids, movers = _plan_split(event) + multi_ids, movers, new_sessions = _plan_split(event) links = _session_detections(event).filter(next_detection__source_image__event=event) tracking_classifications = Classification.objects.filter( detection__source_image__event=event, algorithm__key=TRACKING_ALGORITHM_KEY @@ -234,11 +246,17 @@ def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = batch_size=1000, ) + Occurrence.objects.bulk_update( + [Occurrence(pk=pk, event_id=session_id) for pk, session_id in new_sessions.items()], ["event"] + ) + touched = [*multi_ids, *(o.pk for o in new_occurrences)] determinations_updated = _refresh_determinations(pk for pk in touched if pk not in identified) refresh_track_stats_for_ids(touched) - update_calculated_fields_for_sessions_and_stations([event.pk], stations_async=False) + update_calculated_fields_for_sessions_and_stations( + [event.pk, *set(new_sessions.values())], stations_async=False + ) return SessionResetResult( event_id=event.pk, dry_run=False, diff --git a/ami/ml/post_processing/tests/test_session_reset.py b/ami/ml/post_processing/tests/test_session_reset.py index 75886b969..68fb1d6cb 100644 --- a/ami/ml/post_processing/tests/test_session_reset.py +++ b/ami/ml/post_processing/tests/test_session_reset.py @@ -7,7 +7,7 @@ from django.test.utils import CaptureQueriesContext from django.utils import timezone -from ami.main.models import Classification, Detection, Identification, Occurrence, Taxon +from ami.main.models import Classification, Detection, Identification, Occurrence, SourceImage, Taxon from ami.main.models_future.session_reset import SessionResetRefused, reset_session_tracking, session_tracking_counts from ami.main.tests import cachalot_disabled from ami.ml.models import Algorithm @@ -21,22 +21,28 @@ class TestResetSessionTracking(TestCase): """A reset returns a tracked session to one occurrence per detection, so it can be tracked again.""" - def _build_tracked_session(self, images: int, boxes_per_image: int = 2): + def _build_tracked_session(self, images: int, boxes_per_image: int = 2, nights: int = 1): """A session whose detections each predict their own taxon, tracked into one chain per box - position, with one grouping confirmation and the tracking task's recorded classifications.""" + position, with one grouping confirmation and the tracking task's recorded classifications. + + With more than one night, only the first session is tracked; the later ones keep one + occurrence per detection. + """ project, deployment = setup_test_project(reuse=False) - create_captures(deployment=deployment, num_nights=1, images_per_night=images, interval_minutes=1) + create_captures(deployment=deployment, num_nights=nights, images_per_night=images, interval_minutes=1) create_taxa(project) - create_occurrences(deployment=deployment, num=images * boxes_per_image) - event = project.events.get() + create_occurrences(deployment=deployment, num=nights * images * boxes_per_image) + event = project.events.order_by("start").first() + self.assertEqual(project.events.count(), nights) captures = list(event.captures.order_by("timestamp")) taxa = list(Taxon.objects.filter(projects=project).order_by("pk")) - detections_by_capture = [list(c.detections.order_by("pk")) for c in captures] - for i, detection in enumerate(det for dets in detections_by_capture for det in dets): + all_captures = SourceImage.objects.filter(deployment=deployment).order_by("timestamp") + for i, detection in enumerate(det for c in all_captures for det in c.detections.order_by("pk")): Classification.objects.filter(detection=detection).update( taxon=taxa[i % len(taxa)], score=0.5 + (i % 7) / 20, terminal=True ) + detections_by_capture = [list(c.detections.order_by("pk")) for c in captures] for position in range(boxes_per_image): chain = [dets[position] for dets in detections_by_capture] for current, following in zip(chain, chain[1:]): @@ -108,6 +114,34 @@ def test_command_rejects_a_session_from_another_project(self): with self.assertRaises(CommandError): call_command("reset_tracking", project=project.pk + 1000, events=[event.pk], stdout=io.StringIO()) + def test_an_occurrence_reaching_into_another_session_stays_there(self): + """A track that also holds a later session's detection keeps only that detection, filed under + the later session, and the reset session comes out fresh.""" + project, event = self._build_tracked_session(images=3, nights=2) + later_event = project.events.exclude(pk=event.pk).get() + spanning = Occurrence.objects.filter(event=event).order_by("pk").first() + later_detection = Detection.objects.filter(source_image__event=later_event).order_by("pk").first() + orphaned = later_detection.occurrence + later_detection.occurrence = spanning + later_detection.save(update_fields=["occurrence"]) + orphaned.delete() + later_memberships = dict( + Detection.objects.filter(source_image__event=later_event).values_list("pk", "occurrence_id") + ) + + reset_session_tracking(event) + + self.assertTrue(event_is_fresh(event)[0]) + self.assertEqual( + dict(Detection.objects.filter(source_image__event=later_event).values_list("pk", "occurrence_id")), + later_memberships, + ) + spanning.refresh_from_db() + self.assertEqual(list(spanning.detections.values_list("pk", flat=True)), [later_detection.pk]) + self.assertEqual(spanning.event_id, later_event.pk) + best = spanning.best_prediction + self.assertEqual((spanning.determination_id, spanning.determination_score), (best.taxon_id, best.score)) + def test_query_count_does_not_grow_with_the_session(self): """Doubling the detections must not add queries: every write is a bulk statement.""" counts = [] From da6c917a9912f44921475f7b82a7438f3307e018 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 28 Sep 2026 21:24:27 -0700 Subject: [PATCH 3/7] fix(tracking): clear a determination on reset when no prediction is left to support it After a reset, an occurrence whose remaining detection had no scored classification kept the determination and score it had inherited from the merged track, which described detections that had moved to other occurrences. Such an occurrence, when it has no identification, now has its determination and score cleared. The split test removes the predictions from one kept detection and checks this. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01C7Xf6VPbwWtTumhjjF15g8 --- ami/main/models_future/session_reset.py | 16 +++++++++------- .../post_processing/tests/test_session_reset.py | 9 ++++++--- 2 files changed, 15 insertions(+), 10 deletions(-) diff --git a/ami/main/models_future/session_reset.py b/ami/main/models_future/session_reset.py index efc71f598..32ca1ba93 100644 --- a/ami/main/models_future/session_reset.py +++ b/ami/main/models_future/session_reset.py @@ -13,7 +13,8 @@ also holds detections of another session stays with those detections, untouched there, and every detection of this session leaves it. - Every occurrence the split touched takes its determination from its own - detections' best prediction, the rule ``Occurrence.best_prediction`` applies. An + detections' best prediction, the rule ``Occurrence.best_prediction`` applies, or + none if they have no scored prediction. An occurrence with a human identification keeps the determination it has. - Chain links between two detections of the session are cleared. A link that crosses into another session records one animal across a regroup boundary, which tracking @@ -154,7 +155,8 @@ def _refresh_determinations(occurrence_ids: Iterable[int]) -> int: """Set each occurrence's determination to its best prediction, in batches; returns how many changed. The batched form of ``update_occurrence_determination`` for occurrences with no - identifications, using the same choice as ``Occurrence.best_prediction``. + identifications, using the same choice as ``Occurrence.best_prediction``. Unlike it, + an occurrence with no scored prediction gets no determination. """ ids = list(occurrence_ids) classifications = Classification.objects.select_related(None).only( @@ -174,11 +176,11 @@ def _refresh_determinations(occurrence_ids: Iterable[int]) -> int: ) for occurrence in batch.prefetch_related(Prefetch("detections", queryset=detections)): best = best_prediction_from_prefetch(occurrence) - if best is None or best.taxon_id is None: - continue - if (occurrence.determination_id, occurrence.determination_score) != (best.taxon_id, best.score): - occurrence.determination_id = best.taxon_id - occurrence.determination_score = best.score + # With no scored prediction left, a determination inherited from the merged + # track would describe detections that have moved away, so it is cleared. + wanted = (best.taxon_id, best.score) if best is not None and best.taxon_id else (None, None) + if (occurrence.determination_id, occurrence.determination_score) != wanted: + occurrence.determination_id, occurrence.determination_score = wanted changed.append(occurrence) Occurrence.objects.bulk_update(changed, ["determination", "determination_score"], batch_size=1000) return len(changed) diff --git a/ami/ml/post_processing/tests/test_session_reset.py b/ami/ml/post_processing/tests/test_session_reset.py index 68fb1d6cb..474d01b4b 100644 --- a/ami/ml/post_processing/tests/test_session_reset.py +++ b/ami/ml/post_processing/tests/test_session_reset.py @@ -59,6 +59,8 @@ def _build_tracked_session(self, images: int, boxes_per_image: int = 2, nights: def test_splits_every_track_and_clears_links_verification_and_tracking_records(self): project, event = self._build_tracked_session(images=4) + unscored_keeper = Occurrence.objects.filter(event=event).order_by("pk").first() + Classification.objects.filter(detection=unscored_keeper.detections.order_by("timestamp", "pk").first()).delete() tracked = session_tracking_counts(event) self.assertEqual((tracked.occurrences, tracked.multi_detection_occurrences), (2, 2)) self.assertEqual((tracked.links, tracked.grouping_verified), (6, 1)) @@ -73,9 +75,10 @@ def test_splits_every_track_and_clears_links_verification_and_tracking_records(s self.assertFalse(Classification.objects.filter(algorithm__key="tracking").exists()) for occurrence in Occurrence.objects.filter(detections__source_image__event=event): best = occurrence.best_prediction - self.assertEqual( - (occurrence.determination_id, occurrence.determination_score), (best.taxon_id, best.score) - ) + expected = (best.taxon_id, best.score) if best else (None, None) + self.assertEqual((occurrence.determination_id, occurrence.determination_score), expected) + unscored_keeper.refresh_from_db() + self.assertIsNone(unscored_keeper.determination_id) def test_refuses_a_session_with_identifications_unless_forced(self): project, event = self._build_tracked_session(images=3) From 615af9116826c249d2c445ae6a83063a0b05fce7 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 28 Sep 2026 21:25:14 -0700 Subject: [PATCH 4/7] fix(tracking): plan a session reset inside the transaction that writes it The reset read the session's counts, checked for identifications and chose which detections move before opening its transaction, and locked nothing. A tracking run, a track edit or a new identification landing in between could leave the writes acting on a stale plan, or split an occurrence a person had just identified. The reset now locks the session's occurrence rows first, then reads the counts, checks for identifications and plans the split inside the same transaction as the writes. Anything that deletes those occurrences, points a detection at them or adds an identification to them waits until the reset commits. A dry run still plans outside any transaction and takes no lock. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01C7Xf6VPbwWtTumhjjF15g8 --- ami/main/models_future/session_reset.py | 82 ++++++++++++++++++------- 1 file changed, 59 insertions(+), 23 deletions(-) diff --git a/ami/main/models_future/session_reset.py b/ami/main/models_future/session_reset.py index 32ca1ba93..2a6e741cb 100644 --- a/ami/main/models_future/session_reset.py +++ b/ami/main/models_future/session_reset.py @@ -186,19 +186,23 @@ def _refresh_determinations(occurrence_ids: Iterable[int]) -> int: return len(changed) -def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = False) -> SessionResetResult: - """Return ``event`` to one occurrence per detection, with no links and no confirmations. +@dataclasses.dataclass +class _ResetPlan: + before: SessionTrackingCounts + identified: set[int] + identification_count: int + multi_ids: list[int] + movers: list[int] + new_sessions: dict[int, int] - Raises ``SessionResetRefused`` when the session's occurrences carry human - identifications, unless ``force`` is set; forced, each identification stays on the - occurrence that keeps the first detection. With ``dry_run`` the counts are those a - reset would produce and nothing is written. The query count depends on the number of - occurrences only through the determination batches, not per row. - """ + +def _plan_reset(event: Event, force: bool) -> _ResetPlan: + """Read the session's state and decide what the reset changes; raises ``SessionResetRefused``.""" before = session_tracking_counts(event) - occurrence_ids = _session_occurrence_ids(event) identified = set( - Identification.objects.filter(occurrence_id__in=occurrence_ids).values_list("occurrence_id", flat=True) + Identification.objects.filter(occurrence_id__in=_session_occurrence_ids(event)).values_list( + "occurrence_id", flat=True + ) ) identification_count = Identification.objects.filter(occurrence_id__in=identified).count() if identified else 0 if identification_count and not force: @@ -206,30 +210,62 @@ def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = f"Session {event.pk} has {identification_count} identification(s) on {len(identified)} occurrence(s). " "Resetting would split occurrences a person identified; pass force to reset anyway." ) - multi_ids, movers, new_sessions = _plan_split(event) + return _ResetPlan(before, identified, identification_count, multi_ids, movers, new_sessions) + + +def _lock_session_occurrences(event: Event) -> None: + """Lock the session's occurrence rows until the transaction ends. + + A tracking run, a track edit or a new identification on these occurrences waits for + the reset instead of changing them between the plan and the writes. + """ + list( + Occurrence.objects.select_related(None) + .select_for_update(of=("self",)) + .filter(pk__in=_session_occurrence_ids(event)) + .order_by("pk") + .values_list("pk", flat=True) + ) + + +def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = False) -> SessionResetResult: + """Return ``event`` to one occurrence per detection, with no links and no confirmations. + + Raises ``SessionResetRefused`` when the session's occurrences carry human + identifications, unless ``force`` is set; forced, each identification stays on the + occurrence it was made on, which keeps its first detection in the session (or its + detections in another session). The plan and the writes run in one transaction that + holds the session's occurrences locked. With ``dry_run`` the counts are those a reset + would produce and nothing is written or locked. The query count depends on the number + of occurrences only through the determination batches, not per row. + """ links = _session_detections(event).filter(next_detection__source_image__event=event) tracking_classifications = Classification.objects.filter( detection__source_image__event=event, algorithm__key=TRACKING_ALGORITHM_KEY ) - verified = Occurrence.objects.filter(pk__in=occurrence_ids, grouping_verified_at__isnull=False) + verified = Occurrence.objects.filter(pk__in=_session_occurrence_ids(event), grouping_verified_at__isnull=False) if dry_run: + plan = _plan_reset(event, force) return SessionResetResult( event_id=event.pk, dry_run=True, - identifications=identification_count, - occurrences_split=len(multi_ids), - occurrences_created=len(movers), - links_cleared=before.links, - verifications_cleared=before.grouping_verified, + identifications=plan.identification_count, + occurrences_split=len(plan.multi_ids), + occurrences_created=len(plan.movers), + links_cleared=plan.before.links, + verifications_cleared=plan.before.grouping_verified, tracking_classifications_deleted=tracking_classifications.count(), determinations_updated=0, - before=before, + before=plan.before, after=None, ) with transaction.atomic(): + _lock_session_occurrences(event) + plan = _plan_reset(event, force) + movers, new_sessions = plan.movers, plan.new_sessions links_cleared = links.update(next_detection=None) verifications_cleared = verified.update(grouping_verified_at=None, grouping_verified_by=None) _, deleted_by_model = tracking_classifications.delete() @@ -252,8 +288,8 @@ def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = [Occurrence(pk=pk, event_id=session_id) for pk, session_id in new_sessions.items()], ["event"] ) - touched = [*multi_ids, *(o.pk for o in new_occurrences)] - determinations_updated = _refresh_determinations(pk for pk in touched if pk not in identified) + touched = [*plan.multi_ids, *(o.pk for o in new_occurrences)] + determinations_updated = _refresh_determinations(pk for pk in touched if pk not in plan.identified) refresh_track_stats_for_ids(touched) update_calculated_fields_for_sessions_and_stations( @@ -262,13 +298,13 @@ def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = return SessionResetResult( event_id=event.pk, dry_run=False, - identifications=identification_count, - occurrences_split=len(multi_ids), + identifications=plan.identification_count, + occurrences_split=len(plan.multi_ids), occurrences_created=len(new_occurrences), links_cleared=links_cleared, verifications_cleared=verifications_cleared, tracking_classifications_deleted=tracking_deleted, determinations_updated=determinations_updated, - before=before, + before=plan.before, after=session_tracking_counts(event), ) From 50f8a939050f55c8f38fa6542e533db730041272 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 28 Sep 2026 21:25:49 -0700 Subject: [PATCH 5/7] test(tracking): cover a forced reset's determinations and a link into another session The forced reset test now gives the identified occurrence a taxon that differs from its own prediction, and checks that the occurrence keeps the identified determination while every occurrence split off from it takes its own best prediction. A new test links the last detection of one session to the first detection of the next and checks that the reset keeps that link while clearing the links inside the session. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01C7Xf6VPbwWtTumhjjF15g8 --- .../tests/test_session_reset.py | 26 ++++++++++++++++++- 1 file changed, 25 insertions(+), 1 deletion(-) diff --git a/ami/ml/post_processing/tests/test_session_reset.py b/ami/ml/post_processing/tests/test_session_reset.py index 474d01b4b..5a5c158d6 100644 --- a/ami/ml/post_processing/tests/test_session_reset.py +++ b/ami/ml/post_processing/tests/test_session_reset.py @@ -83,7 +83,10 @@ def test_splits_every_track_and_clears_links_verification_and_tracking_records(s def test_refuses_a_session_with_identifications_unless_forced(self): project, event = self._build_tracked_session(images=3) occurrence = Occurrence.objects.filter(event=event).order_by("pk").first() - Identification.objects.create(occurrence=occurrence, user=UserFactory(), taxon=occurrence.determination) + kept_detection = occurrence.detections.order_by("timestamp", "pk").first() + own_prediction = kept_detection.classifications.order_by("-score").first().taxon + identified_taxon = Taxon.objects.filter(projects=project).exclude(pk=own_prediction.pk).first() + Identification.objects.create(occurrence=occurrence, user=UserFactory(), taxon=identified_taxon) with self.assertRaises(SessionResetRefused): reset_session_tracking(event) @@ -92,6 +95,12 @@ def test_refuses_a_session_with_identifications_unless_forced(self): result = reset_session_tracking(event, force=True) self.assertEqual(result.after.multi_detection_occurrences, 0) self.assertEqual(Identification.objects.filter(occurrence=occurrence).count(), 1) + occurrence.refresh_from_db() + self.assertEqual(occurrence.determination_id, identified_taxon.pk) + split_off = Occurrence.objects.filter(detections__source_image__event=event).exclude(pk=occurrence.pk) + for other in split_off: + best = other.best_prediction + self.assertEqual((other.determination_id, other.determination_score), (best.taxon_id, best.score)) def test_dry_run_reports_the_plan_and_writes_nothing(self): project, event = self._build_tracked_session(images=3) @@ -145,6 +154,21 @@ def test_an_occurrence_reaching_into_another_session_stays_there(self): best = spanning.best_prediction self.assertEqual((spanning.determination_id, spanning.determination_score), (best.taxon_id, best.score)) + def test_keeps_a_link_into_another_session(self): + """A link that crosses a session boundary records one animal across a regroup, so it survives.""" + project, event = self._build_tracked_session(images=3, nights=2) + later_event = project.events.exclude(pk=event.pk).get() + last = Detection.objects.filter(source_image__event=event).order_by("-timestamp", "pk").first() + later_first = Detection.objects.filter(source_image__event=later_event).order_by("timestamp", "pk").first() + last.next_detection = later_first + last.save(update_fields=["next_detection"]) + + reset_session_tracking(event) + + last.refresh_from_db() + self.assertEqual(last.next_detection_id, later_first.pk) + self.assertEqual(session_tracking_counts(event).links, 0) + def test_query_count_does_not_grow_with_the_session(self): """Doubling the detections must not add queries: every write is a bulk statement.""" counts = [] From 06c60112c5cb3ebd16406eab0dcbb58069b082b5 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 28 Sep 2026 21:25:57 -0700 Subject: [PATCH 6/7] style(tracking): apply black to the session reset module and its tests Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01C7Xf6VPbwWtTumhjjF15g8 --- ami/main/models_future/session_reset.py | 4 +--- ami/ml/post_processing/tests/test_session_reset.py | 4 +++- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/ami/main/models_future/session_reset.py b/ami/main/models_future/session_reset.py index 2a6e741cb..046d5fafc 100644 --- a/ami/main/models_future/session_reset.py +++ b/ami/main/models_future/session_reset.py @@ -292,9 +292,7 @@ def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = determinations_updated = _refresh_determinations(pk for pk in touched if pk not in plan.identified) refresh_track_stats_for_ids(touched) - update_calculated_fields_for_sessions_and_stations( - [event.pk, *set(new_sessions.values())], stations_async=False - ) + update_calculated_fields_for_sessions_and_stations([event.pk, *set(new_sessions.values())], stations_async=False) return SessionResetResult( event_id=event.pk, dry_run=False, diff --git a/ami/ml/post_processing/tests/test_session_reset.py b/ami/ml/post_processing/tests/test_session_reset.py index 5a5c158d6..60bf5e129 100644 --- a/ami/ml/post_processing/tests/test_session_reset.py +++ b/ami/ml/post_processing/tests/test_session_reset.py @@ -60,7 +60,9 @@ def _build_tracked_session(self, images: int, boxes_per_image: int = 2, nights: def test_splits_every_track_and_clears_links_verification_and_tracking_records(self): project, event = self._build_tracked_session(images=4) unscored_keeper = Occurrence.objects.filter(event=event).order_by("pk").first() - Classification.objects.filter(detection=unscored_keeper.detections.order_by("timestamp", "pk").first()).delete() + Classification.objects.filter( + detection=unscored_keeper.detections.order_by("timestamp", "pk").first() + ).delete() tracked = session_tracking_counts(event) self.assertEqual((tracked.occurrences, tracked.multi_detection_occurrences), (2, 2)) self.assertEqual((tracked.links, tracked.grouping_verified), (6, 1)) From e9194019a74c3c99a8a69c93a9a41c9329a82972 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Tue, 29 Sep 2026 01:10:58 -0700 Subject: [PATCH 7/7] fix(tracking): drop undone tracking results from history when a session is reset A reset undoes every merge in a session, but the tracking results recorded in the occurrences' history stayed, so after a new run an occurrence listed merges from runs that no longer exist. The reset now deletes the tracking results on the session's occurrences, as it already deletes the tracking task's recorded classifications. Reviews stay, and an occurrence that reaches into another session keeps its whole history because a result does not say which session's run wrote it. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01C7Xf6VPbwWtTumhjjF15g8 --- ami/main/models_future/session_reset.py | 22 +++++++++ .../tests/test_session_reset.py | 48 +++++++++++++++++-- 2 files changed, 67 insertions(+), 3 deletions(-) diff --git a/ami/main/models_future/session_reset.py b/ami/main/models_future/session_reset.py index 046d5fafc..1370f5c7c 100644 --- a/ami/main/models_future/session_reset.py +++ b/ami/main/models_future/session_reset.py @@ -23,6 +23,10 @@ - Classifications the tracking task recorded on the session's detections are deleted. Each is a copy of a prediction on the same detection, made when a merge changed a determination, so it describes a merge that no longer exists. +- For the same reason, the tracking results in the history of the session's + occurrences are deleted. Reviews stay, since the history keeps every review. An + occurrence that also holds detections of another session keeps its whole history, + because a tracking result does not say which session's run wrote it. """ from __future__ import annotations @@ -39,6 +43,7 @@ Event, Identification, Occurrence, + OccurrenceHistoryRecord, update_calculated_fields_for_sessions_and_stations, ) from ami.main.models_future.occurrence import best_prediction_from_prefetch @@ -78,6 +83,7 @@ class SessionResetResult: links_cleared: int verifications_cleared: int tracking_classifications_deleted: int + tracking_history_deleted: int determinations_updated: int before: SessionTrackingCounts after: SessionTrackingCounts | None @@ -245,6 +251,17 @@ def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = detection__source_image__event=event, algorithm__key=TRACKING_ALGORITHM_KEY ) verified = Occurrence.objects.filter(pk__in=_session_occurrence_ids(event), grouping_verified_at__isnull=False) + spanning = ( + Detection.objects.filter(occurrence_id__in=_session_occurrence_ids(event)) + .exclude(source_image__event=event) + .order_by() + .values("occurrence_id") + ) + tracking_history = OccurrenceHistoryRecord.objects.filter( + occurrence_id__in=_session_occurrence_ids(event), + kind=OccurrenceHistoryRecord.Kind.ALGORITHM_RESULT, + subtype=TRACKING_ALGORITHM_KEY, + ).exclude(occurrence_id__in=spanning) if dry_run: plan = _plan_reset(event, force) @@ -257,6 +274,7 @@ def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = links_cleared=plan.before.links, verifications_cleared=plan.before.grouping_verified, tracking_classifications_deleted=tracking_classifications.count(), + tracking_history_deleted=tracking_history.count(), determinations_updated=0, before=plan.before, after=None, @@ -270,6 +288,9 @@ def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = verifications_cleared = verified.update(grouping_verified_at=None, grouping_verified_by=None) _, deleted_by_model = tracking_classifications.delete() tracking_deleted = deleted_by_model.get(Classification._meta.label, 0) + # Before the split, while each occurrence still holds the detections that decide + # whether it reaches into another session. + history_deleted, _ = tracking_history.delete() new_occurrences = Occurrence.objects.bulk_create( [Occurrence(event=event, deployment_id=event.deployment_id, project_id=event.project_id) for _ in movers], @@ -302,6 +323,7 @@ def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = links_cleared=links_cleared, verifications_cleared=verifications_cleared, tracking_classifications_deleted=tracking_deleted, + tracking_history_deleted=history_deleted, determinations_updated=determinations_updated, before=plan.before, after=session_tracking_counts(event), diff --git a/ami/ml/post_processing/tests/test_session_reset.py b/ami/ml/post_processing/tests/test_session_reset.py index 60bf5e129..fb7e74232 100644 --- a/ami/ml/post_processing/tests/test_session_reset.py +++ b/ami/ml/post_processing/tests/test_session_reset.py @@ -7,11 +7,23 @@ from django.test.utils import CaptureQueriesContext from django.utils import timezone -from ami.main.models import Classification, Detection, Identification, Occurrence, SourceImage, Taxon +from ami.main.models import ( + Classification, + Detection, + Identification, + Occurrence, + OccurrenceHistoryRecord, + SourceImage, + Taxon, +) from ami.main.models_future.session_reset import SessionResetRefused, reset_session_tracking, session_tracking_counts from ami.main.tests import cachalot_disabled from ami.ml.models import Algorithm -from ami.ml.post_processing.tracking_task import assign_occurrences_from_detection_chains, event_is_fresh +from ami.ml.post_processing.tracking_task import ( + TrackingHistory, + assign_occurrences_from_detection_chains, + event_is_fresh, +) from ami.tests.fixtures.main import create_captures, create_occurrences, create_taxa, setup_test_project from ami.users.tests.factories import UserFactory @@ -50,7 +62,9 @@ def _build_tracked_session(self, images: int, boxes_per_image: int = 2, nights: current.save(update_fields=["next_detection"]) tracking = Algorithm.objects.get_or_create(key="tracking", defaults={"name": "Occurrence Tracking"})[0] - assign_occurrences_from_detection_chains(captures, logger, record_as=tracking) + assign_occurrences_from_detection_chains( + captures, logger, record_as=tracking, history=TrackingHistory(settings={}, algorithm=tracking) + ) keeper = Occurrence.objects.filter(event=event).order_by("pk").first() Occurrence.objects.filter(pk=keeper.pk).update( grouping_verified_at=timezone.now(), grouping_verified_by=UserFactory() @@ -65,6 +79,7 @@ def test_splits_every_track_and_clears_links_verification_and_tracking_records(s ).delete() tracked = session_tracking_counts(event) self.assertEqual((tracked.occurrences, tracked.multi_detection_occurrences), (2, 2)) + self.assertEqual(OccurrenceHistoryRecord.objects.filter(subtype="tracking").count(), 2) self.assertEqual((tracked.links, tracked.grouping_verified), (6, 1)) result = reset_session_tracking(event) @@ -75,6 +90,8 @@ def test_splits_every_track_and_clears_links_verification_and_tracking_records(s self.assertEqual((after.links, after.grouping_verified), (0, 0)) self.assertTrue(event_is_fresh(event)[0]) self.assertFalse(Classification.objects.filter(algorithm__key="tracking").exists()) + self.assertEqual(result.tracking_history_deleted, 2) + self.assertFalse(OccurrenceHistoryRecord.objects.filter(subtype="tracking").exists()) for occurrence in Occurrence.objects.filter(detections__source_image__event=event): best = occurrence.best_prediction expected = (best.taxon_id, best.score) if best else (None, None) @@ -82,6 +99,31 @@ def test_splits_every_track_and_clears_links_verification_and_tracking_records(s unscored_keeper.refresh_from_db() self.assertIsNone(unscored_keeper.determination_id) + def test_tracking_again_after_a_reset_records_only_the_new_run(self): + """An occurrence's history lists the merges of the run that made it, not those of runs a reset undid.""" + project, event = self._build_tracked_session(images=3) + review = OccurrenceHistoryRecord.objects.create( + occurrence=Occurrence.objects.filter(event=event).order_by("pk").first(), + kind=OccurrenceHistoryRecord.Kind.REVIEW, + subtype="track_complete", + timestamp=timezone.now(), + payload={"detection_ids": [], "frames_count": 3}, + ) + reset_session_tracking(event) + self.assertEqual(OccurrenceHistoryRecord.objects.filter(subtype="tracking").count(), 0) + + captures = list(event.captures.order_by("timestamp")) + for dets in zip(*(list(c.detections.order_by("pk")) for c in captures)): + for current, following in zip(dets, dets[1:]): + current.next_detection = following + current.save(update_fields=["next_detection"]) + assign_occurrences_from_detection_chains(captures, logger, history=TrackingHistory(settings={"run": 2})) + + records = OccurrenceHistoryRecord.objects.filter(subtype="tracking") + self.assertEqual(records.count(), 2) + self.assertEqual({r.payload["settings"]["run"] for r in records}, {2}) + self.assertTrue(OccurrenceHistoryRecord.objects.filter(pk=review.pk).exists()) + def test_refuses_a_session_with_identifications_unless_forced(self): project, event = self._build_tracked_session(images=3) occurrence = Occurrence.objects.filter(event=event).order_by("pk").first()