diff --git a/ami/base/permissions.py b/ami/base/permissions.py index 328aa8384..efe444662 100644 --- a/ami/base/permissions.py +++ b/ami/base/permissions.py @@ -172,6 +172,24 @@ def has_object_permission(self, request, view, obj: BaseModel): return obj.check_permission(request.user, view.action) +TRACKING_NOT_ENABLED_MESSAGE = "Tracking is not enabled for this project." + + +class TrackingEnabled(permissions.BasePermission): + """ + Refuse an action on an object whose project has not opted into tracking. + + Pair it after ObjectPermission, so a user without rights on the object gets the + usual refusal rather than a hint about the project's settings. + """ + + message = TRACKING_NOT_ENABLED_MESSAGE + + def has_object_permission(self, request, view, obj: BaseModel): + project = obj.get_project() if hasattr(obj, "get_project") else None + return bool(project and project.feature_flags.tracking) + + class ProjectPipelineConfigPermission(ObjectPermission): """ Permission for the nested project pipelines route (/projects/{pk}/pipelines/). diff --git a/ami/exports/base.py b/ami/exports/base.py index 389480d5e..7a5c25bbf 100644 --- a/ami/exports/base.py +++ b/ami/exports/base.py @@ -11,6 +11,8 @@ class BaseExporter(ABC): """Base class for all data export handlers.""" file_format = "" # To be defined in child classes + # A project feature flag that must be on before this format can be requested. + required_feature_flag: str | None = None serializer_class = None filter_backends = [] diff --git a/ami/exports/format_types.py b/ami/exports/format_types.py index a3f4c82d0..b7542e013 100644 --- a/ami/exports/format_types.py +++ b/ami/exports/format_types.py @@ -15,10 +15,19 @@ def get_export_serializer(): - from ami.main.api.serializers import OccurrenceSerializer + from ami.main.api.serializers import DetectionNestedSerializer, OccurrenceSerializer + + # The detail response pages its detections; an export carries every one of them. + detail_only_fields = ("grouping_summary", "first_detection", "last_detection") class OccurrenceExportSerializer(OccurrenceSerializer): detection_images = serializers.SerializerMethodField() + detections = DetectionNestedSerializer(many=True, read_only=True) + + class Meta(OccurrenceSerializer.Meta): + # The grouping summary describes the detections as they stand at read + # time; it is not part of the occurrence record and is never exported. + fields = [name for name in OccurrenceSerializer.Meta.fields if name not in detail_only_fields] def get_detection_images(self, obj: Occurrence): """Convert the generator field to a list before serialization""" @@ -110,6 +119,10 @@ class OccurrenceTabularSerializer(serializers.ModelSerializer): agreed_with_user = serializers.SerializerMethodField() determination_matches_machine_prediction = serializers.SerializerMethodField() + # Track grouping confirmation; who confirmed it is user data and is not exported. + grouping_verified = serializers.BooleanField(read_only=True) + grouping_verified_at = serializers.DateTimeField(allow_null=True, read_only=True) + # Detection fields best_detection_url = serializers.SerializerMethodField() best_detection_bbox = serializers.SerializerMethodField() @@ -139,6 +152,8 @@ class Meta: "agreed_with_algorithm", "agreed_with_user", "determination_matches_machine_prediction", + "grouping_verified", + "grouping_verified_at", "detections_count", "first_appearance_timestamp", "last_appearance_timestamp", @@ -252,3 +267,25 @@ def export(self): self.update_job_progress(records_exported) self.update_export_stats(file_temp_path=temp_file.name) return temp_file.name # Return the file path + + +class TracksCSVExporter(BaseExporter): + """One row per detection of every occurrence in scope; see ami/exports/tracks.py for the columns.""" + + file_format = "csv" + required_feature_flag = "tracking" + + def get_queryset(self): + return Occurrence.objects.with_real_detections().filter(project=self.project) # type: ignore[union-attr] + + def export(self): + from ami.exports.tracks import write_tracks_csv + + temp_file = tempfile.NamedTemporaryFile(delete=False, suffix=".csv", mode="w", newline="", encoding="utf-8") + with open(temp_file.name, "w", newline="", encoding="utf-8") as csvfile: + rows = write_tracks_csv(self.queryset, csvfile, on_chunk=self.update_job_progress) + self.update_export_stats(file_temp_path=temp_file.name) + # A tracks file holds one record per detection, not per occurrence. + self.data_export.record_count = rows + self.data_export.save(update_fields=["record_count"]) + return temp_file.name diff --git a/ami/exports/migrations/0002_add_tracks_csv_format.py b/ami/exports/migrations/0002_add_tracks_csv_format.py new file mode 100644 index 000000000..e2af0d8e3 --- /dev/null +++ b/ami/exports/migrations/0002_add_tracks_csv_format.py @@ -0,0 +1,24 @@ +# Generated by Django 4.2.10 on 2026-09-22 17:00 + +from django.db import migrations, models + + +class Migration(migrations.Migration): + dependencies = [ + ("exports", "0001_initial"), + ] + + operations = [ + migrations.AlterField( + model_name="dataexport", + name="format", + field=models.CharField( + choices=[ + ("occurrences_api_json", "occurrences_api_json"), + ("occurrences_simple_csv", "occurrences_simple_csv"), + ("tracks_csv", "tracks_csv"), + ], + max_length=255, + ), + ), + ] diff --git a/ami/exports/registry.py b/ami/exports/registry.py index 29a4cc0e7..8574c77c6 100644 --- a/ami/exports/registry.py +++ b/ami/exports/registry.py @@ -27,3 +27,4 @@ def get_supported_formats(cls): ExportRegistry.register("occurrences_api_json")(format_types.JSONExporter) ExportRegistry.register("occurrences_simple_csv")(format_types.CSVExporter) +ExportRegistry.register("tracks_csv")(format_types.TracksCSVExporter) diff --git a/ami/exports/serializers.py b/ami/exports/serializers.py index e16f63025..3574e7cb1 100644 --- a/ami/exports/serializers.py +++ b/ami/exports/serializers.py @@ -62,6 +62,17 @@ def validate_format(self, value): raise serializers.ValidationError(f"Invalid format. Supported formats are: {supported_formats}") return value + def validate(self, attrs): + attrs = super().validate(attrs) + exporter = ExportRegistry.get_exporter(attrs.get("format")) + flag = getattr(exporter, "required_feature_flag", None) + project = attrs.get("project") + if flag and project is not None and not getattr(project.feature_flags, flag, False): + raise serializers.ValidationError( + {"format": f"The {attrs['format']} export needs the {flag} feature, which is off for this project."} + ) + return attrs + def get_file_url(self, obj): return obj.get_absolute_url(request=self.context.get("request")) diff --git a/ami/exports/tests.py b/ami/exports/tests.py index 866b1af61..3a8e6502f 100644 --- a/ami/exports/tests.py +++ b/ami/exports/tests.py @@ -1,6 +1,8 @@ import csv +import io import json import logging +from unittest import mock from django.core.files.base import ContentFile from django.core.files.storage import default_storage @@ -551,3 +553,262 @@ def test_csv_has_all_new_fields(self): ] for field in expected_fields: self.assertIn(field, headers, f"Missing CSV field: {field}") + + +class TracksExportTest(TestCase): + """The tracks CSV: one row per detection, a fixed column contract, and bounded queries.""" + + EXPECTED_HEADER = ( + "occurrence_id,detection_id,event_id,deployment_id,source_image_id,timestamp,frame_index,frame_count," + "bbox_x1,bbox_y1,bbox_x2,bbox_y2,image_width,image_height,detection_label,detection_score," + "occurrence_determination,occurrence_determination_score,grouping_verified,grouping_verified_at," + "has_feature_vector,next_detection_id" + ) + + def setUp(self): + self.project, self.deployment = setup_test_project(reuse=False) + self.user = self.project.owner + create_captures(deployment=self.deployment, num_nights=1, images_per_night=3, interval_minutes=1) + group_images_into_events(self.deployment) + create_taxa(self.project) + self.taxon = Taxon.objects.filter(projects=self.project).first() + self.algorithm, _ = Algorithm.objects.get_or_create( + name="test-classifier", defaults={"key": "test-classifier"} + ) + self.captures = list(self.project.captures.order_by("timestamp")) + self.captures[0].width, self.captures[0].height = 4096, 2160 + self.captures[0].save() + # Three occurrences of three detections each, created latest frame first so that + # the export cannot get frame order from detection pks. + self.occurrences = [self._make_track(offset=i * 100) for i in range(3)] + + def _make_track(self, offset: int) -> Occurrence: + occurrence = Occurrence.objects.create( + project=self.project, + deployment=self.deployment, + event=self.captures[0].event, + determination=self.taxon, + determination_score=0.9, + ) + for capture in reversed(self.captures): + detection = Detection.objects.create( + source_image=capture, + timestamp=capture.timestamp, + bbox=[offset, offset, offset + 10, offset + 20], + occurrence=occurrence, + ) + detection.classifications.create( + taxon=self.taxon, score=0.8, timestamp=capture.timestamp, algorithm=self.algorithm, terminal=True + ) + return occurrence + + def _rows(self, occurrences=None, **kwargs) -> list[dict[str, str]]: + from ami.exports.tracks import iter_track_rows + + return list(iter_track_rows(occurrences or Occurrence.objects.filter(project=self.project), **kwargs)) + + def _run_format_export(self) -> tuple[str, DataExport]: + data_export = DataExport.objects.create(user=self.user, project=self.project, format="tracks_csv") + file_path = data_export.run_export().replace("/media/", "") + with default_storage.open(file_path, "r") as f: + content = f.read() + default_storage.delete(file_path) + data_export.refresh_from_db() + return content, data_export + + def _request_export(self): + client = APIClient() + client.force_authenticate(user=self.user) + with mock.patch("ami.jobs.models.Job.enqueue"): + return client.post( + f"/api/v2/exports/?project_id={self.project.pk}", + {"project": self.project.pk, "format": "tracks_csv"}, + format="json", + ) + + def test_tracks_export_is_refused_while_tracking_is_off(self): + self.assertFalse(self.project.feature_flags.tracking) + response = self._request_export() + self.assertEqual(response.status_code, 400, response.data) + self.assertIn("format", response.data) + self.assertFalse(DataExport.objects.filter(project=self.project, format="tracks_csv").exists()) + + def test_tracks_export_is_accepted_once_tracking_is_on(self): + self.project.feature_flags.tracking = True + self.project.save(update_fields=["feature_flags"]) + response = self._request_export() + self.assertEqual(response.status_code, 201, response.data) + + def test_format_export_header_is_the_contract(self): + content, data_export = self._run_format_export() + lines = content.splitlines() + self.assertEqual(lines[0], self.EXPECTED_HEADER) + self.assertEqual(len(lines) - 1, 9) + self.assertEqual(data_export.record_count, 9, "A tracks export counts detection rows") + + def test_frame_index_follows_capture_time(self): + rows = [row for row in self._rows() if row["occurrence_id"] == str(self.occurrences[0].pk)] + by_capture = {int(row["source_image_id"]): row for row in rows} + for index, capture in enumerate(self.captures): + row = by_capture[capture.pk] + self.assertEqual(row["frame_index"], str(index)) + self.assertEqual(row["frame_count"], "3") + self.assertEqual(row["timestamp"], capture.timestamp.isoformat()) + first = by_capture[self.captures[0].pk] + self.assertEqual((first["image_width"], first["image_height"]), ("4096", "2160")) + self.assertEqual((first["bbox_x1"], first["bbox_y2"]), ("0", "20")) + self.assertEqual((first["detection_label"], first["detection_score"]), (self.taxon.name, "0.8")) + self.assertEqual(by_capture[self.captures[1].pk]["image_width"], "") + + def test_grouping_verified_flag(self): + from django.utils import timezone + + verified = self.occurrences[1] + verified.grouping_verified_at = timezone.now() + verified.save(update_fields=["grouping_verified_at"]) + + rows = self._rows() + flags = {row["occurrence_id"]: (row["grouping_verified"], row["grouping_verified_at"]) for row in rows} + self.assertEqual(flags[str(verified.pk)], ("true", verified.grouping_verified_at.isoformat())) + self.assertEqual(flags[str(self.occurrences[0].pk)], ("false", "")) + + def test_feature_vector_and_next_detection(self): + from ami.tests.fixtures.tracking import pgvector_is_available + + detections = list(self.occurrences[0].detections.order_by("source_image__timestamp")) + detections[0].next_detection = detections[1] + detections[0].save(update_fields=["next_detection"]) + if pgvector_is_available(): + detections[0].classifications.update(features_2048=[0.1] * 2048) + + rows = {int(row["detection_id"]): row for row in self._rows()} + self.assertEqual(rows[detections[0].pk]["next_detection_id"], str(detections[1].pk)) + self.assertEqual(rows[detections[1].pk]["next_detection_id"], "") + self.assertEqual(rows[detections[1].pk]["has_feature_vector"], "false") + if pgvector_is_available(): + self.assertEqual(rows[detections[0].pk]["has_feature_vector"], "true") + + def test_query_count_is_one_pair_per_chunk(self): + from django.db import connection + from django.test.utils import CaptureQueriesContext + + from ami.main.tests import cachalot_disabled + + # Three occurrences in chunks of two: occurrences, detections, occurrences, detections. + with cachalot_disabled(), CaptureQueriesContext(connection) as ctx: + rows = self._rows(chunk_size=2) + self.assertEqual(len(rows), 9) + self.assertEqual(len(ctx.captured_queries), 4) + + # Doubling the detections per occurrence adds no queries. + for occurrence in self.occurrences: + for detection in list(occurrence.detections.all()): + detection.pk = None + detection.next_detection = None + detection.save() + with cachalot_disabled(), CaptureQueriesContext(connection) as ctx: + rows = self._rows(chunk_size=2) + self.assertEqual(len(rows), 18) + self.assertEqual(len(ctx.captured_queries), 4) + + def _run_command(self, **options) -> tuple[list[dict[str, str]], str]: + import io + + from django.core.management import call_command + + stdout, stderr = io.StringIO(), io.StringIO() + call_command("export_tracks", project=self.project.pk, stdout=stdout, stderr=stderr, **options) + lines = stdout.getvalue().splitlines() + self.assertEqual(lines[0], self.EXPECTED_HEADER) + return list(csv.DictReader(lines)), stderr.getvalue() + + def test_management_command_writes_the_same_csv(self): + rows, summary = self._run_command() + self.assertEqual(len(rows), 9) + self.assertIn("Wrote 9 detection rows", summary) + + def test_management_command_event_filter(self): + import datetime + + from ami.main.models import Event + + first_event = self.captures[0].event + other_event = Event.objects.create( + project=self.project, + deployment=self.deployment, + group_by="2030-01-01", + start=datetime.datetime(2030, 1, 1, 22, 0), + ) + moved = self.occurrences[2] + moved.event = other_event + moved.save(update_fields=["event"]) + + rows, _ = self._run_command(events=[first_event.pk]) + self.assertEqual( + {row["occurrence_id"] for row in rows}, {str(self.occurrences[0].pk), str(self.occurrences[1].pk)} + ) + self.assertEqual(len(rows), 6) + + rows, _ = self._run_command(events=[first_event.pk, other_event.pk]) + self.assertEqual({row["occurrence_id"] for row in rows}, {str(o.pk) for o in self.occurrences}) + self.assertEqual(len(rows), 9) + + def test_management_command_verified_only(self): + from django.utils import timezone + + rows, _ = self._run_command(verified_only=True) + self.assertEqual(rows, [], "No confirmed tracks yet, so only the header comes out") + + verified = self.occurrences[1] + verified.grouping_verified_at = timezone.now() + verified.save(update_fields=["grouping_verified_at"]) + rows, _ = self._run_command(verified_only=True) + self.assertEqual({row["occurrence_id"] for row in rows}, {str(verified.pk)}) + self.assertEqual(len(rows), 3) + self.assertEqual({row["grouping_verified"] for row in rows}, {"true"}) + + def test_undetermined_tracks_are_exported_and_scored_alike(self): + """On a project where only a detector ran, no occurrence has a determination, and its + confirmed tracks must still reach the benchmark file that the scorer reads them from.""" + from django.core.management import call_command + from django.utils import timezone + + undetermined = Occurrence.objects.create( + project=self.project, deployment=self.deployment, event=self.captures[0].event + ) + for capture in self.captures: + Detection.objects.create( + source_image=capture, timestamp=capture.timestamp, bbox=[500, 500, 510, 520], occurrence=undetermined + ) + self.assertIsNone(Occurrence.objects.get(pk=undetermined.pk).determination_id) + Occurrence.objects.filter(pk__in=[undetermined.pk, self.occurrences[0].pk]).update( + grouping_verified_at=timezone.now() + ) + + content, _ = self._run_format_export() + exported = {row["occurrence_id"] for row in csv.DictReader(content.splitlines())} + self.assertIn(str(undetermined.pk), exported) + + rows, _ = self._run_command(verified_only=True) + self.assertEqual({row["occurrence_id"] for row in rows}, {str(undetermined.pk), str(self.occurrences[0].pk)}) + output = io.StringIO() + args = ["--project", str(self.project.pk), "--format", "json", "--no-require-features"] + call_command("evaluate_tracking", *args, stdout=output) + self.assertEqual(len(rows), json.loads(output.getvalue())["overall"]["detections"]) + + def test_occurrence_csv_carries_grouping_confirmation(self): + from django.utils import timezone + + verified = self.occurrences[0] + verified.grouping_verified_at = timezone.now() + verified.save(update_fields=["grouping_verified_at"]) + data_export = DataExport.objects.create(user=self.user, project=self.project, format="occurrences_simple_csv") + file_path = data_export.run_export().replace("/media/", "") + with default_storage.open(file_path, "r") as f: + rows = {row["id"]: row for row in csv.DictReader(f)} + default_storage.delete(file_path) + self.assertNotIn("grouping_verified_by", next(iter(rows.values()))) + self.assertEqual(rows[str(verified.pk)]["grouping_verified"], "True") + self.assertTrue(rows[str(verified.pk)]["grouping_verified_at"]) + self.assertEqual(rows[str(self.occurrences[1].pk)]["grouping_verified"], "False") + self.assertEqual(rows[str(self.occurrences[1].pk)]["grouping_verified_at"], "") diff --git a/ami/exports/tracks.py b/ami/exports/tracks.py new file mode 100644 index 000000000..76150c2ef --- /dev/null +++ b/ami/exports/tracks.py @@ -0,0 +1,174 @@ +""" +One row per detection of every occurrence (track) in scope, for taking tracks out as a benchmark. + +The export format and the ``export_tracks`` management command both write through +``iter_track_rows()``, so ``TRACKS_CSV_COLUMNS`` is the single definition of the columns. +No user data is exported: a confirmed grouping carries its time, not who confirmed it. +""" + +import csv +import datetime +import typing +from collections.abc import Callable, Iterator + +from django.db import models +from django.db.models import Exists, OuterRef, Subquery + +from ami.main.models import BEST_MACHINE_PREDICTION_ORDER, Classification, Detection, Occurrence +from ami.main.models_future.tracks import CAPTURE_ORDER + +TRACKS_CSV_COLUMNS: typing.Final = ( + "occurrence_id", + "detection_id", + "event_id", + "deployment_id", + "source_image_id", + "timestamp", + "frame_index", + "frame_count", + "bbox_x1", + "bbox_y1", + "bbox_x2", + "bbox_y2", + "image_width", + "image_height", + "detection_label", + "detection_score", + "occurrence_determination", + "occurrence_determination_score", + "grouping_verified", + "grouping_verified_at", + "has_feature_vector", + "next_detection_id", +) + +DEFAULT_CHUNK_SIZE: typing.Final = 500 + + +def _cell(value) -> str: + """Render one value the way every tracks CSV writes it: blank for None, lowercase booleans.""" + if value is None: + return "" + if isinstance(value, bool): + return "true" if value else "false" + if isinstance(value, datetime.datetime): + return value.isoformat() + return str(value) + + +def _detections_for(occurrence_ids: list[int]) -> models.QuerySet: + """Real detections of the given occurrences, in frame order, with their own best label.""" + best_classification = Classification.objects.filter(detection=OuterRef("pk")).order_by( + *BEST_MACHINE_PREDICTION_ORDER + ) + return ( + Detection.objects.valid() # type: ignore[attr-defined] Custom queryset method + .filter(occurrence_id__in=occurrence_ids) + .annotate( + label=Subquery(best_classification.values("taxon__name")[:1]), + label_score=Subquery(best_classification.values("score")[:1]), + # Same notion as ClassificationQuerySet.with_has_features(), per detection. + has_feature_vector=Exists( + Classification.objects.filter(detection=OuterRef("pk"), features_2048__isnull=False) + ), + ) + .order_by("occurrence_id", *CAPTURE_ORDER) + .values( + "pk", + "occurrence_id", + "bbox", + "next_detection_id", + "source_image_id", + "source_image__timestamp", + "source_image__width", + "source_image__height", + "source_image__event_id", + "source_image__deployment_id", + "label", + "label_score", + "has_feature_vector", + ) + ) + + +def iter_track_rows( + occurrences: models.QuerySet[Occurrence], + chunk_size: int = DEFAULT_CHUNK_SIZE, + on_chunk: Callable[[int], None] | None = None, +) -> Iterator[dict[str, str]]: + """ + Yield one row per detection of each occurrence in ``occurrences``, keyed by TRACKS_CSV_COLUMNS. + + Occurrences are read in pk order, ``chunk_size`` at a time, with one detection query per + chunk, so the query count grows with the number of chunks and never with the rows. + ``on_chunk`` receives the running count of occurrences read, for progress reporting. + """ + scope = ( + occurrences.order_by("pk") + .values("pk", "determination__name", "determination_score", "grouping_verified_at") + .distinct() # A capture-set filter joins through detections and repeats occurrences. + ) + occurrences_read = 0 + last_pk = 0 + while True: + chunk = list(scope.filter(pk__gt=last_pk)[:chunk_size]) + if not chunk: + break + last_pk = chunk[-1]["pk"] + occurrences_read += len(chunk) + + detections_by_occurrence: dict[int, list[dict]] = {} + for detection in _detections_for([row["pk"] for row in chunk]): + detections_by_occurrence.setdefault(detection["occurrence_id"], []).append(detection) + + for occurrence in chunk: + detections = detections_by_occurrence.get(occurrence["pk"], []) + verified_at = occurrence["grouping_verified_at"] + for frame_index, detection in enumerate(detections): + bbox = detection["bbox"] + x1, y1, x2, y2 = bbox if isinstance(bbox, list) and len(bbox) == 4 else (None,) * 4 + yield { + "occurrence_id": _cell(occurrence["pk"]), + "detection_id": _cell(detection["pk"]), + "event_id": _cell(detection["source_image__event_id"]), + "deployment_id": _cell(detection["source_image__deployment_id"]), + "source_image_id": _cell(detection["source_image_id"]), + "timestamp": _cell(detection["source_image__timestamp"]), + "frame_index": _cell(frame_index), + "frame_count": _cell(len(detections)), + "bbox_x1": _cell(x1), + "bbox_y1": _cell(y1), + "bbox_x2": _cell(x2), + "bbox_y2": _cell(y2), + "image_width": _cell(detection["source_image__width"]), + "image_height": _cell(detection["source_image__height"]), + "detection_label": _cell(detection["label"]), + "detection_score": _cell(detection["label_score"]), + "occurrence_determination": _cell(occurrence["determination__name"]), + "occurrence_determination_score": _cell(occurrence["determination_score"]), + "grouping_verified": _cell(verified_at is not None), + "grouping_verified_at": _cell(verified_at), + "has_feature_vector": _cell(detection["has_feature_vector"]), + "next_detection_id": _cell(detection["next_detection_id"]), + } + + if on_chunk: + on_chunk(occurrences_read) + if len(chunk) < chunk_size: + break + + +def write_tracks_csv( + occurrences: models.QuerySet[Occurrence], + stream: typing.TextIO, + chunk_size: int = DEFAULT_CHUNK_SIZE, + on_chunk: Callable[[int], None] | None = None, +) -> int: + """Write the header and every track row to ``stream``; return the number of detection rows.""" + writer = csv.DictWriter(stream, fieldnames=TRACKS_CSV_COLUMNS) + writer.writeheader() + rows = 0 + for row in iter_track_rows(occurrences, chunk_size=chunk_size, on_chunk=on_chunk): + writer.writerow(row) + rows += 1 + return rows diff --git a/ami/jobs/models.py b/ami/jobs/models.py index ff65f31f2..791677e1c 100644 --- a/ami/jobs/models.py +++ b/ami/jobs/models.py @@ -19,7 +19,12 @@ from ami.jobs.tasks import cleanup_async_job_if_needed, run_job from ami.main.models import Deployment, Project, SourceImage, SourceImageCollection from ami.ml.models import Pipeline -from ami.ml.post_processing.registry import get_postprocessing_task +from ami.ml.post_processing.registry import ( + MEMBER_POST_PROCESSING_TASKS, + get_postprocessing_task, + staff_only_config_fields, +) +from ami.ml.post_processing.tracking_task import TrackingTask from ami.utils.schemas import OrderedEnum logger = logging.getLogger(__name__) @@ -572,7 +577,7 @@ def process_images(cls, job, images): total_classifications = 0 config = job.pipeline.get_config(project_id=job.project.pk) - chunk_size = config.get("request_source_image_batch_size", 1) + chunk_size = config.request_source_image_batch_size chunks = [images[i : i + chunk_size] for i in range(0, image_count, chunk_size)] # noqa request_failed_images = [] job.logger.info(f"Processing {image_count} images in {len(chunks)} batches of up to {chunk_size}") @@ -691,6 +696,7 @@ def process_images(cls, job, images): "Events created", "Events touched", "Empty events deleted", + "Occurrences split at a session boundary", "Duplicate timestamps", "Ungrouped captures", "Captures missing timestamp", @@ -1395,8 +1401,28 @@ def check_custom_permission(self, user, action: str) -> bool: permission_codename = f"{action}_{job_type}_job" project = self.get_project() if hasattr(self, "get_project") else None + if job_type == PostProcessingJob.key and action in ("run", "retry") and not user.is_superuser: + if not self._members_may_run_post_processing(project): + return False return user.has_perm(permission_codename, project) + def _members_may_run_post_processing(self, project: Project | None) -> bool: + """Whether a project role may run this post-processing job, as opposed to staff only. + + Holds for the tasks and config a member could have created through the API, so a + staff task or a staff-only guard setting cannot be re-run by a project role. + """ + params = self.params or {} + task_key = params.get("task") + if task_key not in MEMBER_POST_PROCESSING_TASKS: + return False + config = params.get("config") or {} + if not isinstance(config, dict) or staff_only_config_fields(task_key, config): + return False + if task_key == TrackingTask.key and not (project and project.feature_flags.tracking): + return False + return True + def get_custom_user_permissions(self, user) -> list[str]: project = self.get_project() if not project: diff --git a/ami/jobs/serializers.py b/ami/jobs/serializers.py index f53199e73..16025fafa 100644 --- a/ami/jobs/serializers.py +++ b/ami/jobs/serializers.py @@ -1,7 +1,9 @@ +import pydantic from django_pydantic_field.rest_framework import SchemaField from drf_spectacular.utils import extend_schema_field -from rest_framework import serializers +from rest_framework import exceptions, serializers +from ami.base.permissions import TRACKING_NOT_ENABLED_MESSAGE from ami.exports.models import DataExport from ami.main.api.serializers import ( DefaultSerializer, @@ -9,12 +11,26 @@ SourceImageCollectionNestedSerializer, SourceImageNestedSerializer, ) -from ami.main.models import Deployment, Project, SourceImage, SourceImageCollection +from ami.main.models import Deployment, Event, Project, SourceImage, SourceImageCollection from ami.ml.models import Pipeline +from ami.ml.post_processing.registry import ( + MEMBER_POST_PROCESSING_TASKS, + get_postprocessing_task, + staff_only_config_fields, +) +from ami.ml.post_processing.tracking_task import TrackingConfig, TrackingTask from ami.ml.schemas import PipelineProcessingTask, PipelineTaskResult, ProcessingServiceClientInfo from ami.ml.serializers import PipelineNestedSerializer -from .models import JOB_LOGS_DEFAULT_LIMIT, Job, JobProgress, MLJob, _legacy_logs_shape, serialize_job_logs +from .models import ( + JOB_LOGS_DEFAULT_LIMIT, + Job, + JobProgress, + MLJob, + PostProcessingJob, + _legacy_logs_shape, + serialize_job_logs, +) from .schemas import QueuedTaskAcknowledgment @@ -41,6 +57,73 @@ class JobTypeSerializer(serializers.Serializer): key = serializers.SlugField(read_only=True) +def _pydantic_messages(exc: pydantic.ValidationError) -> list[str]: + messages = [] + for err in exc.errors(): + field = ".".join(str(part) for part in err.get("loc", ()) if part != "__root__") + messages.append(f"{field}: {err['msg']}" if field else err["msg"]) + return messages + + +def validate_post_processing_params(project: Project | None, params, user=None) -> dict: + """Check a post-processing job's ``{"task": ..., "config": {...}}`` before it is saved. + + Returns the params with the config normalized by the task's schema, so the stored + job carries every default the worker will run with. Raises a 400 otherwise. Only a + superuser may set the staff-only config fields, matching who may run such a job. + """ + if not isinstance(params, dict) or set(params) - {"task", "config"}: + raise serializers.ValidationError( + {"params": 'Post-processing jobs take params of the form {"task": , "config": {...}}.'} + ) + task_key = params.get("task") + task_cls = get_postprocessing_task(task_key) if isinstance(task_key, str) else None + if task_cls is None: + raise serializers.ValidationError({"params": {"task": f"Unknown post-processing task {task_key!r}."}}) + if task_key not in MEMBER_POST_PROCESSING_TASKS: + raise serializers.ValidationError( + {"params": {"task": f"The {task_cls.name} task cannot be started through the API."}} + ) + if task_cls is TrackingTask and not (project and project.feature_flags.tracking): + raise serializers.ValidationError({"project_id": TRACKING_NOT_ENABLED_MESSAGE}) + + config = params.get("config") or {} + if not isinstance(config, dict): + raise serializers.ValidationError({"params": {"config": "Must be an object."}}) + staff_only = staff_only_config_fields(task_key, config) + if staff_only and not (user is not None and user.is_superuser): + raise serializers.ValidationError( + {"params": {"config": [f"{name}: Only a superuser can change this setting." for name in staff_only]}} + ) + try: + model = task_cls.config_schema(**config) + except pydantic.ValidationError as exc: + raise serializers.ValidationError({"params": {"config": _pydantic_messages(exc)}}) + + if isinstance(model, TrackingConfig): + if model.event_ids: + found = set(Event.objects.filter(pk__in=model.event_ids, project=project).values_list("pk", flat=True)) + missing = sorted(set(model.event_ids) - found) + if missing: + raise serializers.ValidationError( + {"params": {"config": [f"event_ids: Session(s) {missing} were not found in this project."]}} + ) + if model.source_image_collection_id is not None and not ( + SourceImageCollection.objects.filter(pk=model.source_image_collection_id, project=project).exists() + ): + raise serializers.ValidationError( + { + "params": { + "config": [ + f"source_image_collection_id: Capture set {model.source_image_collection_id} " + "was not found in this project." + ] + } + } + ) + return {"task": task_key, "config": model.dict()} + + class JobListSerializer(DefaultSerializer): delay = serializers.IntegerField() project = JobProjectNestedSerializer(read_only=True) @@ -55,6 +138,8 @@ class JobListSerializer(DefaultSerializer): # All jobs created from the Jobs UI are ML jobs (datasync, etc. are created for the user) # @TODO Remove this when the UI is updated pass a job type. This should be a required field. job_type_key = serializers.SlugField(write_only=True, default=MLJob.key) + # Read by post-processing jobs only: {"task": , "config": {...}}. + params = serializers.JSONField(required=False, allow_null=True) project_id = serializers.PrimaryKeyRelatedField( label="Project", @@ -129,6 +214,7 @@ class Meta: "logs", "job_type", "job_type_key", + "params", "data_export", "dispatch_mode", # "duration", @@ -148,6 +234,31 @@ class Meta: "dispatch_mode", ] + def validate(self, attrs: dict) -> dict: + attrs = super().validate(attrs) + if attrs.get("job_type_key") == PostProcessingJob.key: + project = attrs.get("project") or getattr(self.instance, "project", None) + request = self.context.get("request") + user = getattr(request, "user", None) + attrs["params"] = validate_post_processing_params(project, attrs.get("params"), user) + self._check_may_run_post_processing(project, attrs["params"]) + else: + # Other job types do not read params, so none are stored for them. + attrs.pop("params", None) + return attrs + + def _check_may_run_post_processing(self, project: Project | None, params: dict) -> None: + # Creating a post-processing job takes the permission to run it, so a role that + # cannot start one is refused before a job it could never run is stored. + request = self.context.get("request") + if request is None: + return + job = Job(job_type_key=PostProcessingJob.key, project=project, params=params) + if not job.check_custom_permission(request.user, "run"): + raise exceptions.PermissionDenied( + "You do not have permission to run post-processing jobs in this project." + ) + @extend_schema_field( { "type": "object", diff --git a/ami/jobs/tests/test_jobs.py b/ami/jobs/tests/test_jobs.py index 00b7934a7..9efff9b1b 100644 --- a/ami/jobs/tests/test_jobs.py +++ b/ami/jobs/tests/test_jobs.py @@ -1,6 +1,7 @@ # from rich import print import logging from typing import Any +from unittest.mock import patch from django.test import TestCase from guardian.shortcuts import assign_perm @@ -23,8 +24,9 @@ from ami.ml.models import Pipeline from ami.ml.models.processing_service import ProcessingService from ami.ml.orchestration.jobs import queue_images_to_nats -from ami.tests.fixtures.main import create_captures +from ami.tests.fixtures.main import create_captures, setup_test_project from ami.users.models import User +from ami.users.roles import BasicMember, MLDataManager, ProjectManager logger = logging.getLogger(__name__) @@ -274,6 +276,23 @@ def test_create_job_unauthenticated(self): # Accept either 401 (TokenAuthentication) or 403 (SessionAuthentication with AnonymousUser) self.assertIn(resp.status_code, [status.HTTP_401_UNAUTHORIZED, status.HTTP_403_FORBIDDEN]) + def test_anonymous_writes_are_refused_before_the_body_is_validated(self): + """An unauthenticated caller learns nothing from validation errors, while reads stay public.""" + self.client.force_authenticate(user=None) + jobs_create_url = reverse_with_params("api:job-list", params={"project_id": self.project.pk}) + bodies = { + "empty": {}, + "unknown job type": {"project_id": self.project.pk, "name": "x", "job_type_key": "no-such-type"}, + "post-processing": {"project_id": self.project.pk, "job_type_key": "post_processing", "params": {}}, + } + for label, body in bodies.items(): + with self.subTest(label): + resp = self.client.post(jobs_create_url, body, format="json") + self.assertEqual(resp.status_code, status.HTTP_401_UNAUTHORIZED, resp.data) + self.assertEqual(self.client.get(jobs_create_url).status_code, status.HTTP_200_OK) + detail_url = reverse_with_params("api:job-detail", args=[self.job.pk], params={"project_id": self.project.pk}) + self.assertEqual(self.client.get(detail_url).status_code, status.HTTP_200_OK) + def _create_job(self, name: str, start_now: bool = True): jobs_create_url = reverse_with_params("api:job-list", params={"project_id": self.project.pk}) @@ -1744,3 +1763,239 @@ def test_browsable_page_renders_number_input(self): html = response.content.decode() self.assertNotIn('