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..1370f5c7c --- /dev/null +++ b/ami/main/models_future/session_reset.py @@ -0,0 +1,330 @@ +"""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. 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, 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 + 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. +- 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 + +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, + OccurrenceHistoryRecord, + 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 + tracking_history_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], dict[int, int]]: + """The multi-detection occurrences of the session, the detections that leave them, and + the occurrences whose session changes. + + 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 = dict( + Occurrence.objects.filter(pk__in=_session_occurrence_ids(event)) + .annotate(_n=Count("detections")) + .filter(_n__gt=1) + .values_list("pk", "event_id") + ) + if not multi: + return [], [], {} + rows = list( + Detection.objects.filter(occurrence_id__in=multi) + .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] = [] + new_sessions: dict[int, int] = {} + for occurrence_id, members in by_occurrence.items(): + members.sort(key=_capture_order_key) + 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: + """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``. Unlike it, + an occurrence with no scored prediction gets no determination. + """ + 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) + # 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) + + +@dataclasses.dataclass +class _ResetPlan: + before: SessionTrackingCounts + identified: set[int] + identification_count: int + multi_ids: list[int] + movers: list[int] + new_sessions: dict[int, int] + + +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) + identified = set( + 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: + 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, 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=_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) + return SessionResetResult( + event_id=event.pk, + dry_run=True, + 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(), + tracking_history_deleted=tracking_history.count(), + determinations_updated=0, + 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() + 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], + 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, + ) + + Occurrence.objects.bulk_update( + [Occurrence(pk=pk, event_id=session_id) for pk, session_id in new_sessions.items()], ["event"] + ) + + 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([event.pk, *set(new_sessions.values())], stations_async=False) + return SessionResetResult( + event_id=event.pk, + dry_run=False, + 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, + 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 new file mode 100644 index 000000000..fb7e74232 --- /dev/null +++ b/ami/ml/post_processing/tests/test_session_reset.py @@ -0,0 +1,224 @@ +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, + 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 ( + 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 + +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, 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. + + 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=nights, images_per_night=images, interval_minutes=1) + create_taxa(project) + 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")) + + 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:]): + 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, 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() + ) + return project, event + + 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(OccurrenceHistoryRecord.objects.filter(subtype="tracking").count(), 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()) + 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) + self.assertEqual((occurrence.determination_id, occurrence.determination_score), expected) + 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() + 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) + 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) + 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) + 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_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_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 = [] + 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)