diff --git a/ami/exports/base.py b/ami/exports/base.py index 389480d5e..f5aed4c8d 100644 --- a/ami/exports/base.py +++ b/ami/exports/base.py @@ -11,6 +11,7 @@ class BaseExporter(ABC): """Base class for all data export handlers.""" file_format = "" # To be defined in child classes + filename_label = "" # Optional slug token inserted into export filenames (e.g. "taxa_list") serializer_class = None filter_backends = [] diff --git a/ami/exports/format_types.py b/ami/exports/format_types.py index a3f4c82d0..940fb0ae9 100644 --- a/ami/exports/format_types.py +++ b/ami/exports/format_types.py @@ -8,7 +8,7 @@ from ami.exports.base import BaseExporter from ami.exports.utils import get_data_in_batches -from ami.main.models import Occurrence, SourceImage, get_media_url +from ami.main.models import Detection, Occurrence, SourceImage, get_media_url from ami.ml.schemas import BoundingBox logger = logging.getLogger(__name__) @@ -35,9 +35,10 @@ def to_representation(self, instance): return OccurrenceExportSerializer -class JSONExporter(BaseExporter): +class OccurrencesJSONExporter(BaseExporter): """Handles JSON export of occurrences.""" + filename_label = "occurrences" file_format = "json" def get_serializer_class(self): @@ -209,11 +210,34 @@ def get_best_detection_capture_url(self, obj): return None -class CSVExporter(BaseExporter): - """Handles CSV export of occurrences.""" - +class BaseCSVExporter(BaseExporter): file_format = "csv" + def export(self): + """Exports to CSV format.""" + + temp_file = tempfile.NamedTemporaryFile(delete=False, suffix=".csv", mode="w", newline="", encoding="utf-8") + + # Extract field names dynamically from the serializer + serializer = self.serializer_class() + field_names = list(serializer.fields.keys()) + records_exported = 0 + with open(temp_file.name, "w", newline="", encoding="utf-8") as csvfile: + writer = csv.DictWriter(csvfile, fieldnames=field_names) + writer.writeheader() + + for i, batch in enumerate(get_data_in_batches(self.queryset, self.serializer_class)): + writer.writerows(batch) + records_exported += len(batch) + 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 OccurrencesCSVExporter(BaseCSVExporter): + """Handles CSV export of occurrences.""" + + filename_label = "occurrences" serializer_class = OccurrenceTabularSerializer def get_queryset(self): @@ -233,22 +257,86 @@ def get_queryset(self): .with_verification_info() ) - def export(self): - """Exports occurrences to CSV format.""" - temp_file = tempfile.NamedTemporaryFile(delete=False, suffix=".csv", mode="w", newline="", encoding="utf-8") +class DetectionsTabularSerializer(serializers.ModelSerializer): + """Serializer to format occurrences for tabular data export.""" - # Extract field names dynamically from the serializer - serializer = self.serializer_class() - field_names = list(serializer.fields.keys()) - records_exported = 0 - with open(temp_file.name, "w", newline="", encoding="utf-8") as csvfile: - writer = csv.DictWriter(csvfile, fieldnames=field_names) - writer.writeheader() + event_id = serializers.IntegerField(source="source_image.event.id", allow_null=True) + event_name = serializers.CharField(source="source_image.event.name", allow_null=True) + deployment_id = serializers.IntegerField(source="source_image.deployment.id", allow_null=True) + deployment_name = serializers.CharField(source="source_image.deployment.name", allow_null=True) + project_id = serializers.IntegerField(source="source_image.project.id", allow_null=True) + project_name = serializers.CharField(source="source_image.project.name", allow_null=True) + + source_image_id = serializers.IntegerField(source="source_image.id", allow_null=True) + source_image_path = serializers.CharField(source="source_image.path", allow_null=True) + source_image_timestamp = serializers.DateTimeField(source="source_image.timestamp", allow_null=True) + + detection_bbox = serializers.CharField(source="bbox", allow_null=True) + detection_crop_url = serializers.SerializerMethodField() + detection_score = serializers.FloatField(allow_null=True) + detection_algorithm_id = serializers.IntegerField(source="detection_algorithm.id", allow_null=True) + detection_algorithm_key = serializers.CharField(source="detection_algorithm.key", allow_null=True) + detection_algorithm_name = serializers.CharField(source="detection_algorithm.name", allow_null=True) + determination_id = serializers.IntegerField(source="occurrence.determination.id", allow_null=True) + determination_name = serializers.CharField(source="occurrence.determination.name", allow_null=True) + determination_score = serializers.FloatField(source="occurrence.determination_score", allow_null=True) - for i, batch in enumerate(get_data_in_batches(self.queryset, self.serializer_class)): - writer.writerows(batch) - records_exported += len(batch) - 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 Meta: + model = Detection + fields = [ + "id", + "event_id", + "event_name", + "deployment_id", + "deployment_name", + "project_id", + "project_name", + "source_image_id", + "source_image_path", + "source_image_timestamp", + "detection_bbox", + "detection_crop_url", + "detection_score", + "detection_algorithm_id", + "detection_algorithm_key", + "detection_algorithm_name", + "determination_id", + "determination_name", + "determination_score", + ] + + def get_detection_crop_url(self, obj): + """Returns the full URL to the cropped detection image.""" + path = getattr(obj, "path", None) + return get_media_url(path) if path else None + + +class DetectionsCSVExporter(BaseCSVExporter): + """Handles CSV export of detections.""" + + filename_label = "detections" + serializer_class = DetectionsTabularSerializer + + def get_filter_backends(self): + from ami.main.api.views import OccurrenceCollectionFilter + + class DetectionCollectionFilter(OccurrenceCollectionFilter): + queryset_filter_path = "source_image__collections" + + return [DetectionCollectionFilter] + + def get_queryset(self): + return ( + Detection.objects.valid() # type: ignore[union-attr] Custom queryset method + .filter(source_image__project=self.project) + .select_related( + "occurrence", + "occurrence__determination", + "source_image", + "source_image__project", + "source_image__deployment", + "source_image__event", + "detection_algorithm", + ) + ) diff --git a/ami/exports/migrations/0002_alter_dataexport_format.py b/ami/exports/migrations/0002_alter_dataexport_format.py new file mode 100644 index 000000000..22b13f653 --- /dev/null +++ b/ami/exports/migrations/0002_alter_dataexport_format.py @@ -0,0 +1,24 @@ +# Generated by Django 4.2.10 on 2026-09-02 14:04 + +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"), + ("detections_csv", "detections_csv"), + ], + max_length=255, + ), + ), + ] diff --git a/ami/exports/models.py b/ami/exports/models.py index 073a396da..2b64a4dc0 100644 --- a/ami/exports/models.py +++ b/ami/exports/models.py @@ -66,14 +66,21 @@ def get_filters_display(self): return filters_display def generate_filename(self): - """Generates a slugified filename using project name and export ID.""" + """Generates a slugified filename using project name and export ID. + + When the exporter sets a non-empty `filename_label` (e.g. "taxa_list"), + the label is inserted between the project slug and the export id so + users can tell formats apart in their downloads folder. + """ from ami.exports.registry import ExportRegistry registry = ExportRegistry.get_exporter(self.format) assert registry, f"Export format '{self.format}' not found in registry" extension = registry.file_format + label = getattr(registry, "filename_label", "") or "" project_slug = slugify(self.project.name) # Convert project name to a slug - return f"{project_slug}_export-{self.pk}.{extension}" + label_token = f"{label}_" if label else "" + return f"{project_slug}_{label_token}export-{self.pk}.{extension}" def save_export_file(self, file_temp_path): """ diff --git a/ami/exports/registry.py b/ami/exports/registry.py index 29a4cc0e7..c5feae43e 100644 --- a/ami/exports/registry.py +++ b/ami/exports/registry.py @@ -25,5 +25,6 @@ def get_supported_formats(cls): return list(cls._registry.keys()) -ExportRegistry.register("occurrences_api_json")(format_types.JSONExporter) -ExportRegistry.register("occurrences_simple_csv")(format_types.CSVExporter) +ExportRegistry.register("occurrences_api_json")(format_types.OccurrencesJSONExporter) +ExportRegistry.register("occurrences_simple_csv")(format_types.OccurrencesCSVExporter) +ExportRegistry.register("detections_csv")(format_types.DetectionsCSVExporter) diff --git a/ami/exports/tests.py b/ami/exports/tests.py index 866b1af61..8722de5dd 100644 --- a/ami/exports/tests.py +++ b/ami/exports/tests.py @@ -8,6 +8,7 @@ from rest_framework.test import APIClient from ami.exports.models import DataExport +from ami.exports.registry import ExportRegistry from ami.main.models import Detection, Identification, Occurrence, SourceImageCollection, Taxon from ami.ml.models import Algorithm from ami.tests.fixtures.main import ( @@ -38,7 +39,11 @@ def setUp(self): # Create a collection using the provided method self.collection = self._create_collection() # Define export formats - self.export_formats = ["occurrences_simple_csv", "occurrences_api_json"] + self.export_formats = [ + "occurrences_simple_csv", + "occurrences_api_json", + "detections_csv", + ] def _create_export_with_file(self, format_type): filename = f"exports/test_export_file_{format_type}.json" @@ -113,6 +118,9 @@ def run_and_validate_export(self, format_type): self.validate_csv_records(f) elif format_type == "occurrences_api_json": self.validate_json_records(f) + elif format_type == "detections_csv": + # TODO this checks against Occurrence count not Detections, but 1:1 for now + self.validate_csv_records(f) # Clean up the exported file after the test default_storage.delete(file_path) @@ -305,8 +313,8 @@ def test_non_member_cannot_create_export(self): ) -class ExportNewFieldsTest(TestCase): - """Test the new machine prediction, verification, and detection fields in CSV exports.""" +class ExportDataTestCase(TestCase): + format_type = None def setUp(self): self.project, self.deployment = setup_test_project(reuse=False) @@ -335,6 +343,23 @@ def setUp(self): self.taxon_b = Taxon.objects.create(name="Test Taxon B") self.taxon_b.projects.add(self.project) + def _run_csv_export(self): + """Run a CSV export and return the rows as a list of dicts.""" + data_export = DataExport.objects.create( + user=self.user, + project=self.project, + format=self.format_type, + job=None, + ) + self.data_export = data_export + file_url = data_export.run_export() + self.assertIsNotNone(file_url) + file_path = file_url.replace("/media/", "") + with default_storage.open(file_path, "r") as f: + rows = list(csv.DictReader(f)) + default_storage.delete(file_path) + return rows + def _create_occurrence_with_prediction(self, taxon=None, score=0.85): """Create an occurrence with a single detection and ML classification.""" taxon = taxon or self.taxon_a @@ -355,21 +380,19 @@ def _create_occurrence_with_prediction(self, taxon=None, score=0.85): occurrence = detection.associate_new_occurrence() return occurrence, classification - def _run_csv_export(self): - """Run a CSV export and return the rows as a list of dicts.""" - data_export = DataExport.objects.create( - user=self.user, - project=self.project, - format="occurrences_simple_csv", - job=None, - ) - file_url = data_export.run_export() - self.assertIsNotNone(file_url) - file_path = file_url.replace("/media/", "") - with default_storage.open(file_path, "r") as f: - rows = list(csv.DictReader(f)) - default_storage.delete(file_path) - return rows + def test_export_filename_label(self): + if not self.format_type: + return + label = ExportRegistry.get_exporter(self.format_type).filename_label + occurrence, classification = self._create_occurrence_with_prediction() + self._run_csv_export() + self.assertIn(label, self.data_export.file_url or "") + + +class ExportNewFieldsTest(ExportDataTestCase): + """Test the new machine prediction, verification, and detection fields in CSV exports.""" + + format_type = "occurrences_simple_csv" def test_ml_prediction_only(self): """Occurrence with only ML prediction: machine prediction fields populated, verified_by null.""" @@ -551,3 +574,50 @@ def test_csv_has_all_new_fields(self): ] for field in expected_fields: self.assertIn(field, headers, f"Missing CSV field: {field}") + + +class DetectionsExportFieldsTest(ExportDataTestCase): + format_type = "detections_csv" + + def test_detection_row(self): + """Detection has expected columns""" + occurrence, classification = self._create_occurrence_with_prediction() + detection = occurrence.detections.first() + rows = self._run_csv_export() + + row = next(r for r in rows if int(r["id"]) == detection.pk) + self.assertEqual(row["determination_name"], self.taxon_a.name) + self.assertEqual(row["detection_bbox"], str(detection.bbox)) + self.assertEqual(row["detection_crop_url"], "/media/" + detection.path) + self.assertEqual(row["source_image_path"], detection.source_image.path) + self.assertAlmostEqual(float(row["determination_score"]), 0.85, places=2) + + def test_csv_has_expected_fields(self): + """fields are present as CSV column headers.""" + self._create_occurrence_with_prediction() + rows = self._run_csv_export() + self.assertGreater(len(rows), 0) + headers = rows[0].keys() + expected_fields = [ + "id", + "event_id", + "event_name", + "deployment_id", + "deployment_name", + "project_id", + "project_name", + "source_image_id", + "source_image_path", + "source_image_timestamp", + "detection_bbox", + "detection_crop_url", + "detection_score", + "detection_algorithm_id", + "detection_algorithm_key", + "detection_algorithm_name", + "determination_id", + "determination_name", + "determination_score", + ] + for field in expected_fields: + self.assertIn(field, headers, f"Missing CSV field: {field}") diff --git a/ami/main/api/views.py b/ami/main/api/views.py index 7a17d4a09..18f427b59 100644 --- a/ami/main/api/views.py +++ b/ami/main/api/views.py @@ -1251,6 +1251,7 @@ class OccurrenceCollectionFilter(filters.BaseFilterBackend): Filter occurrences by the capture set their detections' captures belong to. """ + queryset_filter_path = "detections__source_image__collections" query_params = ["collection_id", "collection"] # @TODO remove "collection" param when UI is updated def filter_queryset(self, request, queryset, view): @@ -1261,7 +1262,7 @@ def filter_queryset(self, request, queryset, view): break if collection_id: # Here the queryset is the Occurrence queryset - return queryset.filter(detections__source_image__collections=collection_id) + return queryset.filter(**{self.queryset_filter_path: collection_id}) else: return queryset diff --git a/ui/src/data-services/models/export.ts b/ui/src/data-services/models/export.ts index d96ee2e39..1816fb24a 100644 --- a/ui/src/data-services/models/export.ts +++ b/ui/src/data-services/models/export.ts @@ -6,6 +6,7 @@ import { JobDetails } from './job-details' export const SERVER_EXPORT_TYPES = [ 'occurrences_simple_csv', 'occurrences_api_json', + 'detections_csv', ] as const export type ServerExportType = (typeof SERVER_EXPORT_TYPES)[number] @@ -27,6 +28,7 @@ export class Export extends Entity { const label = { occurrences_simple_csv: 'Occurrences (simple CSV)', occurrences_api_json: 'Occurrences (API JSON)', + detections_csv: 'Detections (CSV)', }[key] return {