Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions ami/jobs/serializers.py
Original file line number Diff line number Diff line change
Expand Up @@ -192,6 +192,23 @@ class Meta:
fields = ["id", "pipeline_slug"]


class JobChoiceSerializer(DefaultSerializer):
"""What a job dropdown needs to name a job.

The job counterpart of ``SourceImageCollectionNestedSerializer``, which serves the capture
set choices: no counts or nested objects, so the choices query stays cheap.
"""

def get_permissions(self, instance, instance_data):
# A picker needs no per-job permissions, and resolving them costs several queries per row.
instance_data["user_permissions"] = []
return instance_data

class Meta:
model = Job
fields = ["id", "name", "details", "job_type_key", "created_at"]


class MLJobTasksRequestSerializer(serializers.Serializer):
"""POST /jobs/{id}/tasks/ — request body sent by a processing service to fetch work.

Expand Down
90 changes: 90 additions & 0 deletions ami/jobs/tests/test_jobs.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,13 +9,15 @@

from ami.base.serializers import reverse_with_params
from ami.jobs.models import (
DataExportJob,
DataStorageSyncJob,
Job,
JobDispatchMode,
JobLog,
JobProgress,
JobState,
MLJob,
PostProcessingJob,
RegroupEventsJob,
SourceImageCollectionPopulateJob,
)
Expand Down Expand Up @@ -1744,3 +1746,91 @@ def test_browsable_page_renders_number_input(self):
html = response.content.decode()
self.assertNotIn('<select name="source_image_single"', html)
self.assertIn('<input type="number" name="source_image_single"', html)


class TestJobChoices(APITestCase):
"""Choices endpoint for the occurrence job filter.

Pins what the dropdown relies on: only jobs that can write detections or
classifications, most recently created first, one capped response, and failed
jobs included because they may have written results before failing.
"""

def setUp(self) -> None:
import datetime

from django.utils import timezone

self.user = User.objects.create_user(email="job-picker@insectai.org", is_staff=False)
self.project = Project.objects.create(name="Job picker project", owner=self.user)
other_project = Project.objects.create(name="Other job picker project", owner=self.user)
self.oldest = Job.objects.create(project=self.project, name="Oldest", job_type_key=MLJob.key)
self.failed = Job.objects.create(
project=self.project, name="Failed", job_type_key=PostProcessingJob.key, status=JobState.FAILURE
)
self.newest = Job.objects.create(project=self.project, name="Newest", job_type_key=MLJob.key)
for days_ago, job in enumerate([self.newest, self.failed, self.oldest]):
Job.objects.filter(pk=job.pk).update(created_at=timezone.now() - datetime.timedelta(days=days_ago * 10))
Job.objects.create(project=self.project, name="Export", job_type_key=DataExportJob.key)
Job.objects.create(project=other_project, name="Elsewhere", job_type_key=MLJob.key)
self.url = f"/api/v2/jobs/choices/?project_id={self.project.pk}"

def test_result_writing_jobs_most_recent_first(self):
response = self.client.get(self.url)
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual([row["name"] for row in response.json()["results"]], ["Newest", "Failed", "Oldest"])

def test_a_project_is_required(self):
response = self.client.get("/api/v2/jobs/choices/")
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)

def test_jobs_in_a_draft_project_are_listed_for_members_only(self):
draft_project = Project.objects.create(name="Draft job picker project", owner=self.user, draft=True)
Job.objects.create(project=draft_project, name="Secret", job_type_key=MLJob.key)
url = f"/api/v2/jobs/choices/?project_id={draft_project.pk}"
self.assertEqual(self.client.get(url).status_code, status.HTTP_404_NOT_FOUND)

self.client.force_authenticate(self.user) # The owner is a member of their project.
response = self.client.get(url)
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual([row["name"] for row in response.json()["results"]], ["Secret"])

def test_query_count_does_not_grow_with_the_number_of_jobs(self):
"""No per-job work: listing eight jobs runs the same queries as listing three, for a member
whose permissions would otherwise be resolved on every row. Cachalot is off so every query
counts."""
from cachalot.api import cachalot_disabled
from django.db import connection
from django.test.utils import CaptureQueriesContext

self.client.force_authenticate(self.user) # The owner is a member of their project.

def list_choices() -> tuple[int, int]:
with CaptureQueriesContext(connection) as queries:
response = self.client.get(self.url)
self.assertEqual(response.status_code, status.HTTP_200_OK)
return len(response.json()["results"]), len(queries.captured_queries)

disabled = cachalot_disabled()
disabled.__enter__()
try:
three_rows, three_rows_queries = list_choices()
Job.objects.bulk_create(
Job(project=self.project, name=f"More {index}", job_type_key=MLJob.key) for index in range(5)
)
eight_rows, eight_rows_queries = list_choices()
finally:
# cachalot_disabled() does not restore itself when the block raises.
disabled.__exit__(None, None, None)
self.assertEqual((three_rows, eight_rows), (3, 8))
self.assertEqual(eight_rows_queries, three_rows_queries)

def test_a_dropdown_gets_one_capped_response_instead_of_pages(self):
Job.objects.bulk_create(
Job(project=self.project, name=f"Bulk {index}", job_type_key=MLJob.key) for index in range(120)
)
for query in ("", "&limit=500"):
with self.subTest(query=query):
body = self.client.get(f"{self.url}{query}").json()
self.assertEqual(len(body["results"]), 100)
self.assertEqual(body["count"], 123)
34 changes: 30 additions & 4 deletions ami/jobs/views.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
from drf_spectacular.utils import extend_schema, extend_schema_view
from rest_framework import serializers
from rest_framework.decorators import action
from rest_framework.exceptions import PermissionDenied, ValidationError
from rest_framework.exceptions import NotFound, PermissionDenied, ValidationError
from rest_framework.filters import BaseFilterBackend
from rest_framework.response import Response

Expand All @@ -37,11 +37,12 @@
update_pipeline_pull_services_seen,
)
from ami.main.api.schemas import project_id_doc_param
from ami.main.api.views import DefaultViewSet
from ami.main.api.views import ChoicesPagination, DefaultViewSet
from ami.main.models import Project
from ami.utils.fields import url_boolean_param

from .models import Job, JobDispatchMode, JobState
from .serializers import JobListSerializer, JobSerializer, MinimalJobSerializer
from .models import Job, JobDispatchMode, JobState, MLJob, PostProcessingJob
from .serializers import JobChoiceSerializer, JobListSerializer, JobSerializer, MinimalJobSerializer

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -257,6 +258,31 @@ def get_serializer_context(self):
)
return context

@extend_schema(parameters=[project_id_doc_param], responses=JobChoiceSerializer(many=True))
@action(detail=False, methods=["get"], name="choices")
def choices(self, request):
"""Choices for the occurrence job filter: jobs that can write detections or classifications.

Follows the capture set choices pattern (``SourceImageCollectionViewSet.choices`` in
ami/main/api/views.py, #1381): a slim serializer, most recent first, and one response
capped by ``ChoicesPagination``. Failed and older jobs are included, since they may
have written results.
"""
project = self.get_active_project()
if project is None:
raise ValidationError({"project_id": "This parameter is required."})
if not Project.objects.visible_for_user(request.user).filter(pk=project.pk).exists():
raise NotFound("Project not found.")
queryset = (
Job.objects.filter(project=project, job_type_key__in=[MLJob.key, PostProcessingJob.key])
.only("id", "name", "job_type_key", "created_at", "project_id")
.order_by("-created_at", "-pk")
)
paginator = ChoicesPagination()
page = paginator.paginate_queryset(queryset, request, view=self)
serializer = JobChoiceSerializer(page, many=True, context=self.get_serializer_context())
return paginator.get_paginated_response(serializer.data)

@action(detail=True, methods=["post"], name="run")
def run(self, request, pk=None):
"""
Expand Down
4 changes: 4 additions & 0 deletions ami/main/api/serializers.py
Original file line number Diff line number Diff line change
Expand Up @@ -1105,6 +1105,7 @@ class ClassificationSerializer(DefaultSerializer):
algorithm = AlgorithmSerializer(read_only=True)
top_n = ClassificationPredictionItemSerializer(many=True, read_only=True)
applied_to = ClassificationAppliedToSerializer(read_only=True)
job = serializers.PrimaryKeyRelatedField(read_only=True, help_text="The job that wrote this classification.")

class Meta:
model = Classification
Expand All @@ -1118,6 +1119,7 @@ class Meta:
"logits",
"top_n",
"applied_to",
"job",
"created_at",
"updated_at",
]
Expand Down Expand Up @@ -1263,6 +1265,7 @@ class DetectionSerializer(DefaultSerializer):
queryset=Algorithm.objects.all(), source="detection_algorithm", write_only=True
)
classifications = ClassificationNestedSerializer(many=True, read_only=True)
job = serializers.PrimaryKeyRelatedField(read_only=True, help_text="The job that wrote this detection.")

class Meta:
model = Detection
Expand All @@ -1271,6 +1274,7 @@ class Meta:
"detection_algorithm",
"detection_algorithm_id",
"classifications",
"job",
]


Expand Down
32 changes: 29 additions & 3 deletions ami/main/api/views.py
Original file line number Diff line number Diff line change
Expand Up @@ -942,8 +942,8 @@ def retrieve(self, request, pk=None):
return response


class CaptureSetChoicesPagination(LimitOffsetPaginationWithPermissions):
"""Sends a whole project's capture set choices in one response.
class ChoicesPagination(LimitOffsetPaginationWithPermissions):
"""Sends a dropdown's choices in one response. Used by the capture set and job choices.

Dropdowns cannot page, so the limit is set here rather than by each caller. It is
capped as well as defaulted, so no caller can ask for a larger one.
Expand Down Expand Up @@ -1031,11 +1031,12 @@ def choices(self, request: Request) -> Response:
Returns only what those consumers need to name a capture set, most recently
updated first, and enough of them that a dropdown never has to page. A project
with more capture sets than the cap needs the search field in #1380.
``JobViewSet.choices`` (ami/jobs/views.py) follows the same pattern.
"""
# Sorting by a count would fail here, since the counts are never annotated.
self.ordering_fields = ["id", "created_at", "updated_at", "name", "method"]
queryset = self.filter_queryset(self.get_queryset())
paginator = CaptureSetChoicesPagination()
paginator = ChoicesPagination()
page = paginator.paginate_queryset(queryset, request, view=self)
serializer = SourceImageCollectionNestedSerializer(page, many=True, context=self.get_serializer_context())
return paginator.get_paginated_response(serializer.data)
Expand Down Expand Up @@ -1448,10 +1449,29 @@ def filter_queryset(self, request, queryset, view):
return queryset


class OccurrenceJobFilter(filters.BaseFilterBackend):
"""
Filter occurrences created or updated by a job.
"""

query_param = "job"

def filter_queryset(self, request, queryset, view):
job_id = SingleParamSerializer[int].clean(
param_name=self.query_param,
field=serializers.IntegerField(required=False, min_value=1),
data=request.query_params,
)
if job_id is None:
return queryset
return queryset.created_or_updated_by_job(job_id)


OCCURRENCE_FILTER_BACKENDS = (
CustomOccurrenceDeterminationFilter,
OccurrenceCollectionFilter,
OccurrenceAlgorithmFilter,
OccurrenceJobFilter,
OccurrenceDateFilter,
OccurrenceVerified,
OccurrenceVerifiedByMeFilter,
Expand Down Expand Up @@ -1558,6 +1578,12 @@ def get_queryset(self) -> QuerySet["Occurrence"]:
required=False,
type=OpenApiTypes.INT,
),
OpenApiParameter(
name="job",
description="Filter occurrences created or updated by a job.",
required=False,
type=OpenApiTypes.INT,
),
]
)
def list(self, request, *args, **kwargs):
Expand Down
43 changes: 43 additions & 0 deletions ami/main/migrations/0096_detection_and_classification_job.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,43 @@
from django.db import migrations, models
import django.db.models.deletion


class Migration(migrations.Migration):
"""Record the job that created each detection and classification.

Adding a nullable foreign key column is a catalogue change: Postgres does not scan
the table, so the lock on these two large tables is held only briefly. The column
is indexed in 0097, concurrently, so no index is built inside this transaction.
"""

dependencies = [
("jobs", "0023_alter_job_job_type_key"),
("main", "0095_grant_sync_deployment_to_mldatamanager"),
]

operations = [
migrations.AddField(
model_name="classification",
name="job",
field=models.ForeignKey(
blank=True,
db_index=False,
null=True,
on_delete=django.db.models.deletion.SET_NULL,
related_name="classifications",
to="jobs.job",
),
),
migrations.AddField(
model_name="detection",
name="job",
field=models.ForeignKey(
blank=True,
db_index=False,
null=True,
on_delete=django.db.models.deletion.SET_NULL,
related_name="detections",
to="jobs.job",
),
),
]
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
from django.contrib.postgres.operations import AddIndexConcurrently
from django.db import migrations, models


class Migration(migrations.Migration):
"""Index the job columns added in 0096 for the ``?job=`` occurrence filter.

Each index covers what the filter reads (job and occurrence, or job and detection),
so the filter is answered from the index alone. Both are partial: rows written before
jobs were recorded have no job and are left out, which keeps the build short and the
index small. Detection and classification are the largest tables, so the indexes are
built CONCURRENTLY, which needs a non-atomic migration; see 0093 for why the statement
timeout is cleared and restored around the build.
"""

atomic = False

dependencies = [
("main", "0096_detection_and_classification_job"),
]

operations = [
migrations.RunSQL(
sql="SET statement_timeout = 0;",
reverse_sql=migrations.RunSQL.noop,
),
AddIndexConcurrently(
model_name="detection",
index=models.Index(
condition=models.Q(job__isnull=False),
fields=["job", "occurrence"],
name="det_job_occurrence_idx",
),
),
AddIndexConcurrently(
model_name="classification",
index=models.Index(
condition=models.Q(job__isnull=False),
fields=["job", "detection"],
name="cls_job_detection_idx",
),
),
migrations.RunSQL(
sql="RESET statement_timeout;",
reverse_sql=migrations.RunSQL.noop,
),
]
Loading
Loading