diff --git a/ami/jobs/models.py b/ami/jobs/models.py index ff65f31f2..cac728838 100644 --- a/ami/jobs/models.py +++ b/ami/jobs/models.py @@ -1,4 +1,5 @@ import datetime +import inspect import logging import random import time @@ -11,15 +12,27 @@ from django.conf import settings from django.db import models, transaction from django.utils.text import slugify +from django.utils.translation import gettext_lazy as _ from django_pydantic_field import SchemaField from guardian.shortcuts import get_perms +from rest_framework import serializers from ami.base.models import BaseModel from ami.base.schemas import ConfigurableStage, ConfigurableStageParam +from ami.jobs.schemas import ( + JOB_GROUP_LABELS, + CaptureSetJobConfig, + JobGroup, + JobGroupDescription, + JobTypeDescription, + JobTypeVariantDescription, + MLJobConfig, + StationJobConfig, +) 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 POSTPROCESSING_TASKS, get_postprocessing_task from ami.utils.schemas import OrderedEnum logger = logging.getLogger(__name__) @@ -433,6 +446,88 @@ def emit(self, record: logging.LogRecord): logger.error(f"Failed to save log for job #{self.job.pk}: {e}") +# Config fields that are also stored on a Job column, so the jobs list can filter and join on them. +JOB_COLUMNS = ("pipeline_id", "source_image_collection_id", "source_image_single_id", "deployment_id") + + +def entity_fields(model: type[pydantic.BaseModel]) -> dict[str, str]: + """Map each config field carrying an ``ami_entity`` hint to that entity (an API route).""" + return { + name: prop["ami_entity"] + for name, prop in model.schema().get("properties", {}).items() + if isinstance(prop, dict) and "ami_entity" in prop + } + + +def pydantic_messages(exc: pydantic.ValidationError) -> list[str]: + """Flatten a pydantic error into ``"field: message"`` lines for a 400 response.""" + 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 _entity_queryset(entity: str, project: Project | None): + """The rows of ``entity`` a job in ``project`` may refer to, or None when not project-scoped.""" + from django.db.models import Q + + from ami.main.models import Event, Occurrence, TaxaList + from ami.ml.models import Algorithm + + scoped = { + "captures/collections": lambda: SourceImageCollection.objects.filter(project=project), + "deployments": lambda: Deployment.objects.filter(project=project), + "captures": lambda: SourceImage.objects.filter(project=project), + # Algorithms are shared catalogue rows: any existing one may be named. + "ml/algorithms": lambda: Algorithm.objects.all(), + # A pipeline is shared; a job may use the ones its project has enabled. + "ml/pipelines": lambda: Pipeline.objects.filter( + project_pipeline_configs__project=project, project_pipeline_configs__enabled=True + ), + "events": lambda: Event.objects.filter(project=project), + "occurrences": lambda: Occurrence.objects.filter(project=project), + # Public lists belong to no project and may be used by any. + "taxa/lists": lambda: TaxaList.objects.filter(Q(projects=project) | Q(projects__isnull=True)), + } + factory = scoped.get(entity) + return factory() if factory else None + + +def check_entities_in_project(values: dict, entities: dict[str, str], project: Project | None) -> None: + """Refuse ids in ``values`` that point outside ``project``. + + ``entities`` maps a field name to its API entity. The schema can only say an id is an + integer; this is the check that it names a row the job's project owns. + """ + errors = [] + for field, entity in entities.items(): + value = values.get(field) + if value in (None, [], ""): + continue + ids = set(value) if isinstance(value, (list, tuple)) else {value} + queryset = _entity_queryset(entity, project) + if queryset is None: + continue + found = set(queryset.filter(pk__in=ids).values_list("pk", flat=True).distinct()) + missing = sorted(ids - found) + if missing: + errors.append(f"{field}: {missing} not found in this project.") + if errors: + raise serializers.ValidationError({"params": {"config": errors}}) + + +def _validate_config(model_cls: type[pydantic.BaseModel], config, project: Project | None) -> pydantic.BaseModel: + if not isinstance(config, dict): + raise serializers.ValidationError({"params": {"config": "Must be an object."}}) + try: + model = model_cls(**config) + except pydantic.ValidationError as exc: + raise serializers.ValidationError({"params": {"config": pydantic_messages(exc)}}) + check_entities_in_project(model.dict(), entity_fields(model_cls), project) + return model + + @dataclass class JobType: """ @@ -441,8 +536,83 @@ class JobType: Job types must be defined as classes because they define code, not just configuration. """ + # A fixed internal name, used in logs and stage names. Users see ``label``. name: str key: str + # What users see for this job type in the Create Job picker and the jobs list, as an action + # ("Process captures"). Wrap it in gettext_lazy to translate it; left empty, ``name`` is used. + label: str = "" + # The Create Job picker heading this type is listed under. A type with variants leaves it + # empty and each variant declares its own. + group: JobGroup | None = None + # Help text under the job type select. Wrap it in gettext_lazy to translate it; left empty, + # the first paragraph of the class docstring is used. + description: str = "" + + # Whether a person can start one from the Create Job dialog. The rest are created by + # the platform for the user: an export from the exports page, for example. + user_creatable: bool = False + + # Everything a new job of this type takes, as a pydantic model: the Create Job dialog renders + # it and the API validates against it. See ami/jobs/schemas.py. + config_schema: type[pydantic.BaseModel] | None = None + + # A job type whose work is chosen from a registry (post-processing tasks) names the + # ``params`` key that holds the choice, and lists the choices as variants. + variant_key: str | None = None + + @classmethod + def help_text(cls) -> str: + """``description``, else the first paragraph of the class docstring.""" + if cls.description: + return str(cls.description) + return (inspect.getdoc(cls) or "").split("\n\n")[0].replace("\n", " ").strip() + + @classmethod + def label_for(cls, params: dict | None) -> str: + """The label users see for one job of this type.""" + return str(cls.label or cls.name) + + @classmethod + def variants(cls, project: Project) -> list[JobTypeVariantDescription]: + return [] + + @classmethod + def describe(cls, project: Project, allowed: bool) -> JobTypeDescription | None: + """What the Create Job dialog shows for this type, or None when it has nothing to offer.""" + variants = cls.variants(project) + if cls.variant_key and not variants: + return None # e.g. post-processing with no method turned on for this project + return JobTypeDescription( + key=cls.key, + name=str(cls.label or cls.name), + description=cls.help_text(), + group=cls.group, + allowed=allowed, + config_schema=cls.config_schema.schema() if cls.config_schema else None, + variant_key=cls.variant_key, + variants=variants, + ) + + @classmethod + def column_ids(cls, params: dict) -> dict[str, int]: + """The Job column ids (``pipeline_id``, ...) carried in validated params.""" + config = params.get("config") or {} + return {field: config[field] for field in JOB_COLUMNS if config.get(field) is not None} + + @classmethod + def validate_params(cls, project: Project | None, user, params) -> dict: + """Check a new job's ``params`` before it is saved and return what should be stored. + + Raises ``serializers.ValidationError`` (a 400) for a bad value. The config is validated + against ``config_schema`` and every id in it must belong to the project. + """ + if not isinstance(params, dict): + raise serializers.ValidationError({"params": "Must be an object."}) + if cls.config_schema is None: + return {} + model = _validate_config(cls.config_schema, params.get("config") or {}, project) + return {"config": model.dict()} # @TODO Consider adding custom vocabulary for job types to be used in the UI # verb: str = "Sync" @@ -460,6 +630,11 @@ def run(cls, job: "Job"): class MLJob(JobType): name = "ML pipeline" key = "ml" + label = _("Process captures") + group = JobGroup.PROCESS_IMAGES + description = _("Detects insects and predicts their species with a pipeline.") + user_creatable = True + config_schema = MLJobConfig @classmethod def run(cls, job: "Job"): @@ -710,6 +885,11 @@ class DataStorageSyncJob(JobType): name = "Data storage sync" key = "data_storage_sync" + label = _("Sync captures from storage") + group = JobGroup.ORGANIZE_CAPTURES + description = _("Add new captures from a station's data storage, then regroup them into sessions.") + user_creatable = True + config_schema = StationJobConfig regroup_stage_key = "regroup_sessions" regroup_stage_name = "Regroup sessions" @@ -804,6 +984,11 @@ def run(cls, job: "Job"): class SourceImageCollectionPopulateJob(JobType): name = "Populate capture set" key = "populate_captures_collection" + label = _("Fill a capture set") + group = JobGroup.ORGANIZE_CAPTURES + description = _("Fill a capture set with the captures its sampling method selects.") + user_creatable = True + config_schema = CaptureSetJobConfig @classmethod def run(cls, job: "Job"): @@ -891,6 +1076,61 @@ def run(cls, job: "Job"): class PostProcessingJob(JobType): name = "Post Processing" key = "post_processing" + description = _( + "Revise existing results with a post-processing method, such as masking classes or " + "filtering out detections too small to identify." + ) + user_creatable = True + variant_key = "task" + + @classmethod + def enabled_tasks(cls, project: Project | None) -> dict: + """The registered tasks whose feature flag is on for ``project``; the others are hidden.""" + flags = project.feature_flags if project else None + return {key: task for key, task in POSTPROCESSING_TASKS.items() if flags and getattr(flags, task.feature_flag)} + + @classmethod + def label_for(cls, params: dict | None) -> str: + """A post-processing job is named by its method, e.g. "Mark detections too small to identify".""" + task_key = (params or {}).get("task") + task_cls = get_postprocessing_task(task_key) if isinstance(task_key, str) else None + return str(task_cls.label) if task_cls else super().label_for(params) + + @classmethod + def variants(cls, project: Project) -> list[JobTypeVariantDescription]: + return [ + JobTypeVariantDescription( + key=key, + name=str(task_cls.label), + description=str(task_cls.description), + group=task_cls.group, + config_schema=task_cls.config_schema.schema(), + ) + for key, task_cls in cls.enabled_tasks(project).items() + ] + + @classmethod + def validate_params(cls, project: Project | None, user, params) -> dict: + """Check a post-processing job's ``{"task": ..., "config": {...}}`` before it is saved. + + The task must be turned on for the project (its feature flag), the config must pass the + task's schema, and every id in it must belong to the project. Returns the params with the + config normalized by the schema, so the stored job carries every default the worker uses. + """ + 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 cls.enabled_tasks(project): + raise serializers.ValidationError( + {"params": {"task": f"{task_cls.name} is not turned on for this project."}} + ) + model = _validate_config(task_cls.config_schema, params.get("config") or {}, project) + return {"task": task_key, "config": model.dict()} @classmethod def run(cls, job: "Job"): @@ -940,6 +1180,11 @@ class RegroupEventsJob(JobType): name = "Regroup sessions" key = "regroup_events" + label = _("Regroup captures into sessions") + group = JobGroup.ORGANIZE_CAPTURES + description = _("Regroup a station's captures into sessions using the project's session time gap.") + user_creatable = True + config_schema = StationJobConfig @classmethod def run(cls, job: "Job"): @@ -988,6 +1233,27 @@ def run(cls, job: "Job"): ] +def describe_job_groups(job_types: list[JobTypeDescription]) -> list[JobGroupDescription]: + """The picker headings ``job_types`` use, in ``JobGroup`` order.""" + used = {jt.group for jt in job_types} | {v.group for jt in job_types for v in jt.variants} + return [JobGroupDescription(key=group, label=str(JOB_GROUP_LABELS[group])) for group in JobGroup if group in used] + + +def describe_job_types(project: Project, user) -> list[JobTypeDescription]: + """The job types ``user`` may pick in the Create Job dialog for ``project``. + + A type the user may not run is still listed with ``allowed=False``, so the dialog can show it + disabled. Permissions are read once for the whole list. + """ + perms = set(get_perms(user, project)) + described = ( + job_type.describe(project, allowed=user.is_superuser or f"run_{job_type.key}_job" in perms) + for job_type in VALID_JOB_TYPES + if job_type.user_creatable + ) + return [description for description in described if description] + + def get_job_type_by_key(key: str) -> type[JobType] | None: for job_type in VALID_JOB_TYPES: if job_type.key == key: @@ -1395,6 +1661,11 @@ 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: + # Turning a method's feature flag off also stops its existing jobs being re-run. + task_key = (self.params or {}).get("task") + if task_key not in PostProcessingJob.enabled_tasks(project): + return False return user.has_perm(permission_codename, project) def get_custom_user_permissions(self, user) -> list[str]: diff --git a/ami/jobs/schemas.py b/ami/jobs/schemas.py index c4b37b92a..11413db5d 100644 --- a/ami/jobs/schemas.py +++ b/ami/jobs/schemas.py @@ -1,4 +1,7 @@ +import enum + import pydantic +from django.utils.translation import gettext_lazy as _ from drf_spectacular.utils import OpenApiParameter @@ -10,6 +13,119 @@ class QueuedTaskAcknowledgment(pydantic.BaseModel): task_id: str +# What each job type takes. The Create Job dialog renders these models' JSON Schema as its form and +# the API validates a new job against them. Fields named in ``JOB_COLUMNS`` (ami/jobs/models.py) are +# also stored on that Job column. Extra ``Field`` keywords are UI hints that pydantic copies into the +# schema: ``ami_widget="entity"`` + ``ami_entity=""`` renders a picker (and the server +# checks the id belongs to the project), ``ami_widget="hidden"`` keeps a field out of the form, and +# ``ami_advanced=True`` puts it under "More settings". Use ``typing.Literal`` for choices. + + +def _pipeline(): + return pydantic.Field(..., title=_("Pipeline"), ami_widget="entity", ami_entity="ml/pipelines") + + +def _capture_set(): + return pydantic.Field( + ..., + title=_("Capture set"), + ami_widget="entity", + ami_entity="captures/collections", + ) + + +def _station(): + return pydantic.Field(..., title=_("Station"), ami_widget="entity", ami_entity="deployments") + + +class MLJobConfig(pydantic.BaseModel): + pipeline_id: int = _pipeline() + source_image_collection_id: int | None = pydantic.Field( + None, + title=_("Capture set"), + description=_("The captures to process."), + ami_widget="entity", + ami_entity="captures/collections", + ) + # Set by other entry points (a single capture's "Process now", a station), not by the dialog. + source_image_single_id: int | None = pydantic.Field(None, ami_widget="hidden", ami_entity="captures") + deployment_id: int | None = pydantic.Field(None, ami_widget="hidden", ami_entity="deployments") + + @pydantic.root_validator(skip_on_failure=True) + def _needs_captures(cls, values: dict) -> dict: + if not any(values.get(f) for f in ("source_image_collection_id", "source_image_single_id", "deployment_id")): + raise ValueError("Choose a capture set to process.") + return values + + class Config: + extra = "forbid" + + +class CaptureSetJobConfig(pydantic.BaseModel): + source_image_collection_id: int = _capture_set() + + class Config: + extra = "forbid" + + +class StationJobConfig(pydantic.BaseModel): + deployment_id: int = _station() + + class Config: + extra = "forbid" + + +class JobGroup(str, enum.Enum): + """Where a job is listed in the Create Job picker, by what the user wants to do. + + Declared by each job type and post-processing task as ``group``; the picker shows the groups in + this order. + """ + + PROCESS_IMAGES = "process_images" # sends images to a processing service + REFINE_RESULTS = "refine_results" # works on results already in Antenna + ORGANIZE_CAPTURES = "organize_captures" + MODELS = "models" + + +JOB_GROUP_LABELS = { + JobGroup.PROCESS_IMAGES: _("Process images"), + JobGroup.REFINE_RESULTS: _("Refine results"), + JobGroup.ORGANIZE_CAPTURES: _("Organize captures"), + JobGroup.MODELS: _("Models"), +} + + +class JobGroupDescription(pydantic.BaseModel): + """A heading in the Create Job picker.""" + + key: JobGroup + label: str + + +class JobTypeVariantDescription(pydantic.BaseModel): + """One method a job type can run, such as a post-processing task.""" + + key: str + name: str + description: str + group: JobGroup + config_schema: dict + + +class JobTypeDescription(pydantic.BaseModel): + """A job type the Create Job dialog can offer, as served by ``GET /jobs/types/``.""" + + key: str + name: str + description: str + group: JobGroup | None # None when each variant is listed under its own group instead + allowed: bool # whether the requesting user may run it in this project + config_schema: dict | None + variant_key: str | None # the params key holding the chosen variant, e.g. "task" + variants: list[JobTypeVariantDescription] + + ids_only_param = OpenApiParameter( name="ids_only", description="Return only job IDs instead of full objects", diff --git a/ami/jobs/serializers.py b/ami/jobs/serializers.py index f53199e73..78a125db0 100644 --- a/ami/jobs/serializers.py +++ b/ami/jobs/serializers.py @@ -1,6 +1,6 @@ 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.exports.models import DataExport from ami.main.api.serializers import ( @@ -9,13 +9,21 @@ SourceImageCollectionNestedSerializer, SourceImageNestedSerializer, ) -from ami.main.models import Deployment, Project, SourceImage, SourceImageCollection -from ami.ml.models import Pipeline +from ami.main.models import Project 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 .schemas import QueuedTaskAcknowledgment +from .models import ( + JOB_LOGS_DEFAULT_LIMIT, + VALID_JOB_TYPES, + Job, + JobProgress, + JobType, + _legacy_logs_shape, + get_job_type_by_key, + serialize_job_logs, +) +from .schemas import JobGroupDescription, JobTypeDescription, QueuedTaskAcknowledgment class JobProjectNestedSerializer(DefaultSerializer): @@ -37,8 +45,14 @@ class Meta: class JobTypeSerializer(serializers.Serializer): - name = serializers.CharField(read_only=True) + """A job's type, named as users see it. A post-processing job is named by its method.""" + key = serializers.SlugField(read_only=True) + name = serializers.CharField(read_only=True) + + def to_representation(self, job) -> dict: + job_type = job.job_type() + return {"key": job_type.key, "name": job_type.label_for(job.params)} class JobListSerializer(DefaultSerializer): @@ -51,10 +65,8 @@ class JobListSerializer(DefaultSerializer): data_export = DataExportNestedSerializer(read_only=True) progress = SchemaField(schema=JobProgress, read_only=True) logs = serializers.SerializerMethodField() - job_type = JobTypeSerializer(read_only=True) - # 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) + job_type = JobTypeSerializer(source="*", read_only=True) + job_type_key = serializers.SlugField(write_only=True) project_id = serializers.PrimaryKeyRelatedField( label="Project", @@ -63,42 +75,6 @@ class JobListSerializer(DefaultSerializer): queryset=Project.objects.all(), source="project", ) - deployment_id = serializers.PrimaryKeyRelatedField( - label="Deployment", - write_only=True, - required=False, - allow_null=True, - # @TODO should this be filtered by project (from URL for new job?) - queryset=Deployment.objects.all(), - source="deployment", - ) - source_image_single_id = serializers.PrimaryKeyRelatedField( - label="Source Image", - write_only=True, - required=False, - allow_null=True, - # @TODO should this be filtered by project (from URL for new job?) - queryset=SourceImage.objects.all(), - source="source_image_single", - ) - source_image_collection_id = serializers.PrimaryKeyRelatedField( - label="Capture Set", - write_only=True, - required=False, - allow_null=True, - # @TODO should this be filtered by project (from URL for new job?) - queryset=SourceImageCollection.objects.all(), - source="source_image_collection", - ) - pipeline_id = serializers.PrimaryKeyRelatedField( - label="Pipeline", - write_only=True, - required=False, - allow_null=True, - # @TODO should this be filtered by project (from URL for new job?) - queryset=Pipeline.objects.all(), - source="pipeline", - ) class Meta: model = Job @@ -112,13 +88,9 @@ class Meta: "project", "project_id", "deployment", - "deployment_id", "source_image_collection", - "source_image_collection_id", "source_image_single", - "source_image_single_id", "pipeline", - "pipeline_id", "status", "created_at", "updated_at", @@ -176,11 +148,61 @@ def get_logs(self, obj: Job) -> dict[str, list[str]]: class JobSerializer(JobListSerializer): # progress = serializers.JSONField(initial=Job.default_progress(), allow_null=False, required=False) + # A job's settings, checked by its job type when the job is created and fixed after that: + # a later update cannot swap in settings that were never validated. See JobType.validate_params. + params = serializers.JSONField(required=False, allow_null=True) + class Meta(JobListSerializer.Meta): fields = JobListSerializer.Meta.fields + [ "result", + "params", ] + def validate_job_type_key(self, value: str) -> str: + job_type = get_job_type_by_key(value) + if not job_type: + known = sorted(t.key for t in VALID_JOB_TYPES if t.user_creatable) + raise serializers.ValidationError(f"Unknown job type '{value}'. Known types: {known}") + if self.instance is None and not job_type.user_creatable: + raise serializers.ValidationError( + f"{job_type.name} jobs are created by the platform, not through this API." + ) + return value + + def validate(self, attrs: dict) -> dict: + attrs = super().validate(attrs) + if self.instance is not None: + # A job's project, type and settings are fixed once it exists: they were validated + # together, and the worker reads the settings against that project. + for field in ("project", "job_type_key"): + if field in attrs and attrs[field] != getattr(self.instance, field): + raise serializers.ValidationError({field: "Cannot be changed after the job is created."}) + attrs.pop("params", None) + return attrs + + # A new job's inputs all arrive in params["config"] and are checked by its job type's + # model; the ids among them that are Job columns are stored there too. + job_type = get_job_type_by_key(attrs["job_type_key"]) + project = attrs["project"] + user = getattr(self.context.get("request"), "user", None) + attrs["params"] = job_type.validate_params(project, user, attrs.get("params") or {}) + attrs.update(job_type.column_ids(attrs["params"])) + if job_type.variant_key: + self._check_may_run(job_type, project, attrs["params"], user) + return attrs + + def _check_may_run(self, job_type: type[JobType], project: Project | None, params: dict, user) -> None: + # Creating a job whose type runs registered methods (post-processing) takes the + # permission to run it, so a role that cannot start one is refused before a job it + # could never run is stored. + if user is None: + return + job = Job(job_type_key=job_type.key, project=project, params=params) + if not job.check_custom_permission(user, "run"): + raise exceptions.PermissionDenied( + f"You do not have permission to run {job_type.name} jobs in this project." + ) + class MinimalJobSerializer(DefaultSerializer): """Minimal serializer returning only essential job fields.""" @@ -192,6 +214,13 @@ class Meta: fields = ["id", "pipeline_slug"] +class JobTypesResponseSerializer(serializers.Serializer): + """GET /jobs/types/ — the job types the Create Job dialog may offer for a project.""" + + groups = SchemaField(schema=list[JobGroupDescription]) + results = SchemaField(schema=list[JobTypeDescription]) + + class MLJobTasksRequestSerializer(serializers.Serializer): """POST /jobs/{id}/tasks/ — request body sent by a processing service to fetch work. diff --git a/ami/jobs/tests/test_job_types.py b/ami/jobs/tests/test_job_types.py new file mode 100644 index 000000000..94f2fb12a --- /dev/null +++ b/ami/jobs/tests/test_job_types.py @@ -0,0 +1,243 @@ +from cachalot.api import cachalot_disabled +from rest_framework import status +from rest_framework.test import APITestCase + +from ami.base.serializers import reverse_with_params +from ami.jobs.models import Job, PostProcessingJob +from ami.main.models import Project, ProjectFeatureFlags, SourceImageCollection, TaxaList +from ami.ml.models import Algorithm, Pipeline +from ami.ml.models.project_pipeline_config import ProjectPipelineConfig +from ami.ml.post_processing.registry import POSTPROCESSING_TASKS +from ami.users.models import User +from ami.users.roles import BasicMember, MLDataManager + + +def types_url(project_id=None): + params = {"project_id": project_id} if project_id is not None else {} + return reverse_with_params("api:job-types", params=params) + + +def enable(project: Project, *flags: str) -> None: + for flag in flags: + setattr(project.feature_flags, flag, True) + project.save() + + +class JobTypesTestBase(APITestCase): + @classmethod + def setUpTestData(cls): + owner = User.objects.create_user(email="owner@insectai.org") + cls.project = Project.objects.create(name="Job types project", owner=owner) + cls.other_project = Project.objects.create(name="Other project", owner=owner) + cls.basic = User.objects.create_user(email="basic@insectai.org") + BasicMember.assign_user(cls.basic, cls.project) + cls.ml_manager = User.objects.create_user(email="ml@insectai.org") + MLDataManager.assign_user(cls.ml_manager, cls.project) + cls.superuser = User.objects.create_user(email="super@insectai.org", is_staff=True, is_superuser=True) + + def get_types(self, user, project_id=None) -> dict: + self.client.force_authenticate(user=user) + response = self.client.get(types_url(self.project.pk if project_id is None else project_id)) + self.assertEqual(response.status_code, status.HTTP_200_OK) + return {t["key"]: t for t in response.json()["results"]} + + +class TestJobTypesEndpoint(JobTypesTestBase): + """GET /jobs/types/ lists what a project member may create, and only members may read it.""" + + def test_only_project_members_may_read_it(self): + self.client.force_authenticate(user=None) + self.assertIn( + self.client.get(types_url(self.project.pk)).status_code, + (status.HTTP_401_UNAUTHORIZED, status.HTTP_403_FORBIDDEN), + ) + self.client.force_authenticate(user=User.objects.create_user(email="outsider@insectai.org")) + self.assertEqual(self.client.get(types_url(self.project.pk)).status_code, status.HTTP_403_FORBIDDEN) + + def test_project_id_is_required_and_validated(self): + self.client.force_authenticate(user=self.basic) + self.assertEqual(self.client.get(types_url()).status_code, status.HTTP_400_BAD_REQUEST) + self.assertEqual(self.client.get(types_url("abc")).status_code, status.HTTP_400_BAD_REQUEST) + self.assertEqual(self.client.get(types_url(999999)).status_code, status.HTTP_404_NOT_FOUND) + + def test_lists_creatable_types_with_the_users_permission(self): + types = self.get_types(self.basic) + # Exports are made from the exports page; post-processing has no method turned on yet. + self.assertEqual(set(types), {"ml", "data_storage_sync", "populate_captures_collection", "regroup_events"}) + self.assertFalse(types["ml"]["allowed"]) + ml_schema = types["ml"]["config_schema"] + self.assertEqual(ml_schema["required"], ["pipeline_id"]) + self.assertEqual(ml_schema["properties"]["source_image_collection_id"]["ami_entity"], "captures/collections") + self.assertTrue(self.get_types(self.ml_manager)["ml"]["allowed"]) + + def test_only_methods_turned_on_for_the_project_are_offered(self): + enable(self.project, "class_masking") + post_processing = self.get_types(self.ml_manager)["post_processing"] + self.assertTrue(post_processing["allowed"]) + self.assertEqual([v["key"] for v in post_processing["variants"]], ["class_masking"]) + properties = post_processing["variants"][0]["config_schema"]["properties"] + self.assertEqual(properties["taxa_list_id"]["title"], "Taxa list to keep") + self.assertEqual(properties["source_image_collection_id"]["ami_entity"], "captures/collections") + self.assertEqual(properties["occurrence_id"]["ami_widget"], "hidden") + + def test_choices_are_named_as_actions_under_picker_headings(self): + """Each choice carries a user-facing label and a heading; only headings in use are listed.""" + enable(self.project, "class_masking") + self.client.force_authenticate(user=self.ml_manager) + data = self.client.get(types_url(self.project.pk)).json() + self.assertEqual([g["key"] for g in data["groups"]], ["process_images", "refine_results", "organize_captures"]) + types = {t["key"]: t for t in data["results"]} + self.assertEqual((types["ml"]["name"], types["ml"]["group"]), ("Process captures", "process_images")) + # A type with variants is listed through them, each under its own heading. + self.assertIsNone(types["post_processing"]["group"]) + masking = types["post_processing"]["variants"][0] + self.assertEqual( + (masking["name"], masking["group"]), ("Limit predictions to a species list", "refine_results") + ) + + def test_query_count_does_not_grow_with_job_types(self): + enable(self.project, "class_masking", "small_size_filter") + self.client.force_authenticate(user=self.ml_manager) + with cachalot_disabled(): + # Project, membership, user and group permissions, plus the request's savepoint pair. + with self.assertNumQueries(6): + response = self.client.get(types_url(self.project.pk)) + self.assertEqual(len(response.json()["results"]), 5) + + def test_every_task_names_a_real_feature_flag(self): + for task in POSTPROCESSING_TASKS.values(): + self.assertIn(task.feature_flag, ProjectFeatureFlags.__fields__, task.key) + + +class TestCreateJobWithParams(JobTypesTestBase): + """POST /jobs/ checks a job's settings against its job type before the job is stored.""" + + @classmethod + def setUpTestData(cls): + super().setUpTestData() + cls.collection = SourceImageCollection.objects.create(name="Mine", project=cls.project) + cls.other_collection = SourceImageCollection.objects.create(name="Theirs", project=cls.other_project) + cls.taxa_list = TaxaList.objects.create(name="Keep") + cls.taxa_list.projects.add(cls.project) + cls.algorithm = Algorithm.objects.create(name="Classifier", key="classifier") + enable(cls.project, "class_masking") + + def post_job(self, user, **body): + self.client.force_authenticate(user=user) + payload = {"name": "Job", "delay": 0, "project_id": self.project.pk, **body} + return self.client.post(reverse_with_params("api:job-list"), payload, format="json") + + def post_masking(self, user, **config): + config = { + "source_image_collection_id": self.collection.pk, + "taxa_list_id": self.taxa_list.pk, + "algorithm_id": self.algorithm.pk, + **config, + } + return self.post_job(user, job_type_key="post_processing", params={"task": "class_masking", "config": config}) + + def test_ml_data_manager_starts_an_enabled_method_with_any_settings(self): + response = self.post_masking(self.ml_manager, reweight=False) + self.assertEqual(response.status_code, status.HTTP_201_CREATED, response.json()) + params = Job.objects.get(pk=response.json()["id"]).params + self.assertEqual(params["task"], "class_masking") + self.assertFalse(params["config"]["reweight"]) + self.assertIsNone(params["config"]["occurrence_id"]) # defaults are stored + + def test_basic_member_cannot_start_post_processing(self): + self.assertEqual(self.post_masking(self.basic).status_code, status.HTTP_403_FORBIDDEN) + + def test_a_method_turned_off_for_the_project_is_refused(self): + self.project.feature_flags.class_masking = False + self.project.save() + response = self.post_masking(self.superuser) + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + self.assertIn("not turned on", str(response.json())) + + def test_ids_from_another_project_are_refused(self): + response = self.post_masking(self.ml_manager, source_image_collection_id=self.other_collection.pk) + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + self.assertIn("source_image_collection_id", str(response.json())) + # Shared rows such as algorithms must at least exist. + response = self.post_masking(self.ml_manager, algorithm_id=999999) + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + self.assertIn("algorithm_id", str(response.json())) + response = self.post_job( + self.superuser, + job_type_key="populate_captures_collection", + params={"config": {"source_image_collection_id": self.other_collection.pk}}, + ) + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + # A pipeline the project has not enabled (or has switched off) counts as another project's. + pipeline = Pipeline.objects.create(name="Not enabled here") + ProjectPipelineConfig.objects.create(project=self.project, pipeline=pipeline, enabled=False) + response = self.post_job( + self.superuser, + job_type_key="ml", + params={"config": {"pipeline_id": pipeline.pk, "source_image_collection_id": self.collection.pk}}, + ) + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + self.assertIn("pipeline_id", str(response.json())) + + def test_schema_errors_come_back_per_field(self): + response = self.post_masking(self.ml_manager, taxa_list_id=None) + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + self.assertIn("taxa_list_id: none is not an allowed value", response.json()["params"]["config"]) + + def test_platform_job_types_cannot_be_created_through_the_api(self): + response = self.post_job(self.superuser, job_type_key="data_export") + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + self.assertIn("job_type_key", response.json()) + + def test_inputs_are_stored_in_params_and_on_job_columns(self): + response = self.post_job( + self.superuser, + job_type_key="populate_captures_collection", + params={"config": {"source_image_collection_id": self.collection.pk}}, + ) + self.assertEqual(response.status_code, status.HTTP_201_CREATED, response.json()) + job = Job.objects.get(pk=response.json()["id"]) + self.assertEqual(job.params, {"config": {"source_image_collection_id": self.collection.pk}}) + self.assertEqual(job.source_image_collection_id, self.collection.pk) + + def test_jobs_list_names_a_post_processing_job_by_its_method(self): + response = self.post_masking(self.ml_manager) + self.assertEqual(response.status_code, status.HTTP_201_CREATED, response.json()) + job_type = self.client.get(reverse_with_params("api:job-detail", args=[response.json()["id"]])).json()[ + "job_type" + ] + self.assertEqual(job_type, {"key": "post_processing", "name": "Limit predictions to a species list"}) + + def test_an_ml_job_needs_a_pipeline(self): + response = self.post_job(self.superuser, job_type_key="ml", params={"config": {}}) + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + self.assertIn("pipeline_id: field required", response.json()["params"]["config"]) + + def test_an_ml_job_needs_something_to_process(self): + pipeline = Pipeline.objects.create(name="Enabled here") + pipeline.projects.add(self.project) + response = self.post_job(self.superuser, job_type_key="ml", params={"config": {"pipeline_id": pipeline.pk}}) + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + self.assertIn("Choose a capture set to process.", response.json()["params"]["config"]) + + def test_project_type_and_params_are_fixed_after_creation(self): + job_id = self.post_masking(self.superuser).json()["id"] + detail = reverse_with_params("api:job-detail", args=[job_id]) + self.client.patch(detail, {"params": {}}, format="json") + self.assertEqual(Job.objects.get(pk=job_id).params["task"], "class_masking") + # The viewset also scopes its lookup by a posted project_id, so a move may 404 first. + response = self.client.patch(detail, {"project_id": self.other_project.pk}, format="json") + self.assertIn(response.status_code, (status.HTTP_400_BAD_REQUEST, status.HTTP_404_NOT_FOUND)) + self.assertEqual(Job.objects.get(pk=job_id).project_id, self.project.pk) + response = self.client.patch(detail, {"job_type_key": "ml"}, format="json") + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + + def test_turning_a_method_off_stops_members_re_running_its_jobs(self): + job = Job.objects.get(pk=self.post_masking(self.ml_manager).json()["id"]) + self.assertTrue(job.check_custom_permission(self.ml_manager, "retry")) + self.project.feature_flags.class_masking = False + self.project.save() + job.refresh_from_db() + self.assertFalse(job.check_custom_permission(self.ml_manager, "retry")) + self.assertTrue(job.check_custom_permission(self.superuser, "retry")) + self.assertEqual(PostProcessingJob.enabled_tasks(self.project), {}) diff --git a/ami/jobs/tests/test_jobs.py b/ami/jobs/tests/test_jobs.py index 00b7934a7..60d81c9b9 100644 --- a/ami/jobs/tests/test_jobs.py +++ b/ami/jobs/tests/test_jobs.py @@ -264,13 +264,13 @@ def test_create_job_unauthenticated(self): jobs_create_url = reverse_with_params("api:job-list", params={"project_id": self.project.pk}) job_data = { "project_id": self.project.pk, - "source_image_collection_id": self.source_image_collection.pk, + "params": {"config": {"source_image_collection_id": self.source_image_collection.pk}}, "name": "Test job unauthenticated", "delay": 0, "job_type_key": SourceImageCollectionPopulateJob.key, } self.client.force_authenticate(user=None) - resp = self.client.post(jobs_create_url, job_data) + resp = self.client.post(jobs_create_url, job_data, format="json") # Accept either 401 (TokenAuthentication) or 403 (SessionAuthentication with AnonymousUser) self.assertIn(resp.status_code, [status.HTTP_401_UNAUTHORIZED, status.HTTP_403_FORBIDDEN]) @@ -281,12 +281,12 @@ def _create_job(self, name: str, start_now: bool = True): job_data = { "project_id": self.job.project.pk, "name": name, - "source_image_collection_id": self.source_image_collection.pk, + "params": {"config": {"source_image_collection_id": self.source_image_collection.pk}}, "delay": 0, "start_now": start_now, "job_type_key": SourceImageCollectionPopulateJob.key, } - resp = self.client.post(jobs_create_url, job_data) + resp = self.client.post(jobs_create_url, job_data, format="json") self.client.force_authenticate(user=None) self.assertEqual(resp.status_code, 201) return resp.json() diff --git a/ami/jobs/views.py b/ami/jobs/views.py index 380d9f959..5c34a995d 100644 --- a/ami/jobs/views.py +++ b/ami/jobs/views.py @@ -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 NotAuthenticated, PermissionDenied, ValidationError from rest_framework.filters import BaseFilterBackend from rest_framework.response import Response @@ -40,8 +40,8 @@ from ami.main.api.views import DefaultViewSet 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, describe_job_groups, describe_job_types +from .serializers import JobListSerializer, JobSerializer, JobTypesResponseSerializer, MinimalJobSerializer logger = logging.getLogger(__name__) @@ -257,6 +257,34 @@ def get_serializer_context(self): ) return context + @extend_schema(parameters=[project_id_doc_param], responses=JobTypesResponseSerializer) + @action(detail=False, methods=["get"], name="types") + def types(self, request): + """ + List the job types the Create Job dialog can offer for a project, each with a JSON Schema + of its settings (``config_schema``), and the picker headings (``groups``) they are listed under. + + Only members of the project may read it. See docs/claude/reference/jobs-panel.md. + """ + # ObjectPermission.has_permission allows every request, and a list-style action never + # reaches the object check, so this action gates itself. + if not request.user.is_authenticated: + raise NotAuthenticated() + self.require_project = True + project = self.get_active_project() + if project is None: # get_active_project already raises 400/404 when required + raise ValidationError({"project_id": "This parameter is required."}) + user = request.user + if not (user.is_superuser or project.owner_id == user.pk or project.members.filter(pk=user.pk).exists()): + raise PermissionDenied("Only members of this project can list its job types.") + job_types = describe_job_types(project, user) + return Response( + { + "groups": [group.dict() for group in describe_job_groups(job_types)], + "results": [job_type.dict() for job_type in job_types], + } + ) + @action(detail=True, methods=["post"], name="run") def run(self, request, pk=None): """ diff --git a/ami/main/migrations/0096_grant_run_post_processing_to_ml_data_manager.py b/ami/main/migrations/0096_grant_run_post_processing_to_ml_data_manager.py new file mode 100644 index 000000000..edecc614d --- /dev/null +++ b/ami/main/migrations/0096_grant_run_post_processing_to_ml_data_manager.py @@ -0,0 +1,84 @@ +""" +Grant ``run_post_processing_job`` to ``MLDataManager`` and ``ProjectManager`` role groups on +existing projects, so the roles that run ML jobs can also start the post-processing methods a +project has enabled, from the Create Job dialog or the jobs API. + +The grant is a guardian object-level permission per project, which is what +``user.has_perm(codename, project)`` reads; see 0095 for the same pattern. The +``post_migrate`` role sync also applies it, and this migration keeps the grant explicit. +""" + +from django.db import migrations +from django.db.models import Q + +CODENAME = "run_post_processing_job" +# ProjectManager inherits MLDataManager's permissions, so both groups hold the grant. +ROLE_SUFFIXES = ("_MLDataManager", "_ProjectManager") + + +def _role_groups() -> Q: + return Q(name__endswith=ROLE_SUFFIXES[0]) | Q(name__endswith=ROLE_SUFFIXES[1]) + + +def _permission(apps): + Permission = apps.get_model("auth", "Permission") + ContentType = apps.get_model("contenttypes", "ContentType") + try: + project_ct = ContentType.objects.get(app_label="main", model="project") + except ContentType.DoesNotExist: + return None, None + try: + return Permission.objects.get(codename=CODENAME, content_type=project_ct), project_ct + except Permission.DoesNotExist: + return None, None + + +def _project_pk_from_group(group): + # Group names are "{project_pk}_{project_name}_{RoleName}"; the pk is the leading segment. + try: + return int(group.name.split("_", 1)[0]) + except (ValueError, IndexError): + return None + + +def grant(apps, schema_editor): + Group = apps.get_model("auth", "Group") + GroupObjectPermission = apps.get_model("guardian", "GroupObjectPermission") + + perm, project_ct = _permission(apps) + if perm is None: + return + for group in Group.objects.filter(_role_groups()): + project_pk = _project_pk_from_group(group) + if project_pk is None: + continue + group.permissions.add(perm) + GroupObjectPermission.objects.get_or_create( + permission=perm, content_type=project_ct, object_pk=str(project_pk), group=group + ) + + +def revoke(apps, schema_editor): + Group = apps.get_model("auth", "Group") + GroupObjectPermission = apps.get_model("guardian", "GroupObjectPermission") + + perm, project_ct = _permission(apps) + if perm is None: + return + for group in Group.objects.filter(_role_groups()): + project_pk = _project_pk_from_group(group) + if project_pk is None: + continue + group.permissions.remove(perm) + GroupObjectPermission.objects.filter( + permission=perm, content_type=project_ct, object_pk=str(project_pk), group=group + ).delete() + + +class Migration(migrations.Migration): + dependencies = [ + ("main", "0095_grant_sync_deployment_to_mldatamanager"), + ("guardian", "0002_generic_permissions_index"), + ] + + operations = [migrations.RunPython(grant, revoke)] diff --git a/ami/main/models.py b/ami/main/models.py index 9943c111a..e84bce46c 100644 --- a/ami/main/models.py +++ b/ami/main/models.py @@ -279,6 +279,10 @@ class ProjectFeatureFlags(pydantic.BaseModel): # Feature flag for jobs to reprocess all images in the project, even if already processed reprocess_all_images: bool = False async_pipeline_workers: bool = True # Whether to use async pipeline workers that pull tasks from a queue + # Post-processing methods offered in the Create Job dialog. Each one is off until a project + # turns it on; then ML data managers and project managers can run it with any settings. + class_masking: bool = False + small_size_filter: bool = False def get_default_feature_flags() -> ProjectFeatureFlags: diff --git a/ami/main/tests.py b/ami/main/tests.py index aab6d943d..a9b6de587 100644 --- a/ami/main/tests.py +++ b/ami/main/tests.py @@ -2594,6 +2594,8 @@ def _assign_roles(self): def _create_project(self, owner): self.project = Project.objects.create(name="Insect Project", description="Test Description", owner=owner) self.deployment = Deployment.objects.create(name="Test Deployment", project=self.project) + self.pipeline, _ = Pipeline.objects.get_or_create(name="Role test pipeline") + self.pipeline.projects.add(self.project) S3StorageSource.objects.create(name="New source", project=self.project, bucket="Test Bucket") create_captures(deployment=self.deployment) create_taxa(project=self.project) @@ -2682,7 +2684,13 @@ def _test_role_permissions(self, role_class, user, permissions_map): "site": {"description": "New Site", "name": "Site 1", "project": self.project.pk}, "device": {"description": "New Device", "name": "Device 1", "project": self.project.pk}, "storage": {"name": "New Storage", "project": self.project.pk, "bucket": "test-bucket"}, - "job": {"delay": "1", "name": "Test Job", "project_id": self.project.pk}, + "job": { + "delay": 1, + "name": "Test Job", + "project_id": self.project.pk, + "job_type_key": "ml", + "params": {"config": {"pipeline_id": self.pipeline.pk, "deployment_id": self.deployment.pk}}, + }, "identification": {"occurrence_id": occurrence_id, "taxon_id": taxon_id, "comment": "Identifier comment"}, "project": {"name": "New Project", "description": "This is a test project."}, } @@ -2712,6 +2720,8 @@ def _test_role_permissions(self, role_class, user, permissions_map): logger.info(f"entity endpoint : {endpoints[entity]}") if entity == "project": response = self.client.post(endpoints[entity], create_data.get(entity, {}), format="multipart") + elif entity == "job": + response = self.client.post(endpoints[entity], create_data[entity], format="json") else: response = self.client.post(endpoints[entity], create_data.get(entity, {})) expected_status = status.HTTP_201_CREATED if can_create else status.HTTP_403_FORBIDDEN @@ -2798,7 +2808,9 @@ def _test_role_permissions(self, role_class, user, permissions_map): continue if "create" in actions: logger.info(f"Testing {role_class} for create permission on {entity} after role unassignment") - response = self.client.post(endpoints[entity], create_data.get(entity, {})) + response = self.client.post( + endpoints[entity], create_data.get(entity, {}), format="json" if entity == "job" else None + ) self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) if "update" in actions: logger.info(f"Testing {role_class} for update permission on {entity} after role unassignment") @@ -3331,8 +3343,8 @@ def test_user_can_run_single_image_job_and_perm_is_reflected(self): "delay": 0, "name": f"Capture #{self.capture.pk}", "project_id": str(self.project.pk), - "pipeline_id": str(self.pipeline.pk), - "source_image_single_id": str(self.capture.pk), + "job_type_key": "ml", + "params": {"config": {"pipeline_id": self.pipeline.pk, "source_image_single_id": self.capture.pk}}, } response = self.client.post(run_url, payload, format="json") self.assertEqual( diff --git a/ami/ml/post_processing/admin/actions.py b/ami/ml/post_processing/admin/actions.py index 4e7729256..46967ac05 100644 --- a/ami/ml/post_processing/admin/actions.py +++ b/ami/ml/post_processing/admin/actions.py @@ -34,7 +34,7 @@ from django.urls import reverse from django.utils.html import format_html, format_html_join -from ami.jobs.models import Job +from ami.jobs.models import Job, PostProcessingJob from ami.ml.post_processing.admin.forms import BasePostProcessingActionForm from ami.ml.post_processing.base import BasePostProcessingTask @@ -124,11 +124,14 @@ def default_build_jobs( job_pks: list[int] = [] with transaction.atomic(): for obj, model in validated: + params = {"task": task_cls.key, "config": model.dict()} job = Job.objects.create( name=name_resolver(task_cls, obj), project=project_resolver(obj), - job_type_key="post_processing", - params={"task": task_cls.key, "config": model.dict()}, + job_type_key=PostProcessingJob.key, + params=params, + # Same columns as a job created through the API, so the jobs list filters agree. + **PostProcessingJob.column_ids(params), ) job.enqueue() job_pks.append(job.pk) diff --git a/ami/ml/post_processing/base.py b/ami/ml/post_processing/base.py index 8f197f192..b8facdc1e 100644 --- a/ami/ml/post_processing/base.py +++ b/ami/ml/post_processing/base.py @@ -6,6 +6,7 @@ import pydantic +from ami.jobs.schemas import JobGroup from ami.ml.models import Algorithm from ami.ml.models.algorithm import AlgorithmTaskType @@ -25,8 +26,18 @@ class BasePostProcessingTask(abc.ABC): # Each task must override these key: str + # A fixed internal name. It identifies the task's Algorithm row, so renaming it starts a new one. name: str + # What users see in the Create Job picker and the jobs list, as an action (wrap in gettext_lazy). + label: str + # The Create Job picker heading the task is listed under. + group: JobGroup config_schema: type[pydantic.BaseModel] + # The ProjectFeatureFlags field that offers this task in the Create Job dialog. While it is + # off the task is hidden and its jobs cannot be started or re-run, except by a superuser. + feature_flag: str + # Help text under the method select (wrap in gettext_lazy); left empty, the docstring is used. + description: str = "" def __init_subclass__(cls, **kwargs): super().__init_subclass__(**kwargs) @@ -35,7 +46,7 @@ def __init_subclass__(cls, **kwargs): # defer key/name/config_schema to its concrete subclasses. if inspect.isabstract(cls): return - required_attrs = ["key", "name", "config_schema"] + required_attrs = ["key", "name", "config_schema", "feature_flag", "label", "group"] for attr in required_attrs: if not hasattr(cls, attr) or getattr(cls, attr) is None: raise TypeError(f"{cls.__name__} must define '{attr}' class attribute") diff --git a/ami/ml/post_processing/class_masking.py b/ami/ml/post_processing/class_masking.py index 2da2001b7..1f5af8fb4 100644 --- a/ami/ml/post_processing/class_masking.py +++ b/ami/ml/post_processing/class_masking.py @@ -6,7 +6,9 @@ from django.db import transaction from django.db.models import QuerySet from django.utils import timezone +from django.utils.translation import gettext_lazy as _ +from ami.jobs.schemas import JobGroup from ami.main.models import Classification, Occurrence, SourceImageCollection, TaxaList from ami.ml.models.algorithm import Algorithm, AlgorithmTaskType from ami.ml.post_processing.base import BasePostProcessingTask @@ -19,16 +21,37 @@ class ClassMaskingConfig(pydantic.BaseModel): # capture set is the bulk path; a single occurrence is the spot/dev path (fast # feedback while tuning a taxa list). This mirrors SmallSizeFilterConfig's # discriminated-scope shape — the shared pattern for per-occurrence triggers. - source_image_collection_id: int | None = None - occurrence_id: int | None = None - # The taxa list to keep: classes whose taxon is not in this list are masked out. - taxa_list_id: int - # The source classifier whose terminal classifications are re-scored. - algorithm_id: int - # When True (default), renormalise the kept classes' scores to sum to 1 after - # masking. When False, the kept classes retain their original absolute scores and - # the excluded classes are zeroed; the chosen species is identical either way. - reweight: bool = True + source_image_collection_id: int | None = pydantic.Field( + None, + title=_("Capture set"), + ami_widget="entity", + ami_entity="captures/collections", + ) + # The admin's single-occurrence path; not offered in the Create Job dialog. + occurrence_id: int | None = pydantic.Field(None, ami_widget="hidden", ami_entity="occurrences") + taxa_list_id: int = pydantic.Field( + ..., + title=_("Taxa list to keep"), + description=_("Classes outside this list are masked out."), + ami_widget="entity", + ami_entity="taxa/lists", + ) + algorithm_id: int = pydantic.Field( + ..., + title=_("Source classifier"), + description=_("Its terminal predictions are the ones re-scored."), + ami_widget="entity", + ami_entity="ml/algorithms", + ami_entity_filters={"task_type": "classification"}, + ) + reweight: bool = pydantic.Field( + True, + title=_("Reweight scores"), + description=_( + "Renormalise the kept classes to sum to 1. Off keeps raw absolute scores; " + "the chosen species is the same either way." + ), + ) @pydantic.root_validator(skip_on_failure=True) def _exactly_one_scope(cls, values: dict) -> dict: @@ -237,6 +260,13 @@ def make_classifications_filtered_by_taxa_list( class ClassMaskingTask(BasePostProcessingTask): key = "class_masking" name = "Class masking" + label = _("Limit predictions to a species list") + group = JobGroup.REFINE_RESULTS + feature_flag = "class_masking" + description = _( + "Masks out classes whose taxon is not on the chosen list and renormalises each prediction over " + "what remains. The original classification is kept and demoted." + ) config_schema = ClassMaskingConfig def _get_or_create_masking_algorithm( diff --git a/ami/ml/post_processing/small_size_filter.py b/ami/ml/post_processing/small_size_filter.py index 6a39af780..9c5632328 100644 --- a/ami/ml/post_processing/small_size_filter.py +++ b/ami/ml/post_processing/small_size_filter.py @@ -1,7 +1,9 @@ import pydantic from django.db.models import QuerySet from django.utils import timezone +from django.utils.translation import gettext_lazy as _ +from ami.jobs.schemas import JobGroup from ami.main.models import Classification, Detection, Occurrence, SourceImageCollection, Taxon, TaxonRank from ami.ml.post_processing.base import BasePostProcessingTask from ami.ml.schemas import BoundingBox @@ -12,9 +14,21 @@ class SmallSizeFilterConfig(pydantic.BaseModel): # set is the bulk path; a single occurrence is the spot/dev path (fast feedback # while tuning a filter). This discriminated-scope shape is the pattern other # post-processing tasks copy when they gain per-occurrence / per-event triggers. - source_image_collection_id: int | None = None - occurrence_id: int | None = None - size_threshold: float = 0.0008 + source_image_collection_id: int | None = pydantic.Field( + None, + title=_("Capture set"), + ami_widget="entity", + ami_entity="captures/collections", + ) + # The admin's single-occurrence path; not offered in the Create Job dialog. + occurrence_id: int | None = pydantic.Field(None, ami_widget="hidden", ami_entity="occurrences") + size_threshold: float = pydantic.Field( + 0.0008, + title=_("Size threshold"), + description=_("Detections smaller than this fraction of the image area are marked as not identifiable."), + gt=0.0, + lt=1.0, + ) @pydantic.validator("size_threshold") def _threshold_in_unit_interval(cls, v: float) -> float: @@ -36,6 +50,10 @@ class Config: class SmallSizeFilterTask(BasePostProcessingTask): key = "small_size_filter" name = "Small size filter" + label = _("Mark detections too small to identify") + group = JobGroup.REFINE_RESULTS + feature_flag = "small_size_filter" + description = _("Marks detections that are too small to identify, so they stop counting towards species totals.") config_schema = SmallSizeFilterConfig def _scoped_detections(self, config: SmallSizeFilterConfig) -> tuple[QuerySet[Detection], str]: diff --git a/ami/ml/post_processing/tests/test_small_size_filter_admin.py b/ami/ml/post_processing/tests/test_small_size_filter_admin.py index 2759bea64..e1eac7dd1 100644 --- a/ami/ml/post_processing/tests/test_small_size_filter_admin.py +++ b/ami/ml/post_processing/tests/test_small_size_filter_admin.py @@ -94,6 +94,7 @@ def test_valid_post_creates_one_job_with_threshold_in_config(self): self.assertEqual(job.params["task"], "small_size_filter") self.assertEqual(job.params["config"]["size_threshold"], 0.001) self.assertEqual(job.params["config"]["source_image_collection_id"], self.collection.pk) + self.assertEqual(job.source_image_collection_id, self.collection.pk) # same column as an API-made job def test_success_message_links_to_the_created_job(self): """The post-run admin message links each created Job to its admin change diff --git a/ami/ml/views.py b/ami/ml/views.py index 63e460af6..e3b3afd5f 100644 --- a/ami/ml/views.py +++ b/ami/ml/views.py @@ -40,7 +40,7 @@ class AlgorithmViewSet(DefaultViewSet, ProjectMixin): queryset = Algorithm.objects.all() serializer_class = AlgorithmSerializer - filterset_fields = ["name", "version"] + filterset_fields = ["name", "version", "task_type"] ordering_fields = [ "id", "created_at", diff --git a/ami/users/roles.py b/ami/users/roles.py index 718146e41..23a0a7e28 100644 --- a/ami/users/roles.py +++ b/ami/users/roles.py @@ -146,6 +146,9 @@ class MLDataManager(Role): Project.Permissions.SYNC_DEPLOYMENT, Project.Permissions.RUN_REGROUP_EVENTS_JOB, Project.Permissions.RUN_DATA_EXPORT_JOB, + # Post-processing methods are offered only where the project turns on their + # feature flag; see BasePostProcessingTask.feature_flag. + Project.Permissions.RUN_POST_PROCESSING_JOB, Project.Permissions.DELETE_OCCURRENCES, Project.Permissions.CREATE_PROJECT_PIPELINE_CONFIG, Project.Permissions.UPDATE_PROJECT_PIPELINE_CONFIG, diff --git a/docs/claude/INDEX.md b/docs/claude/INDEX.md index 61f0980d8..89e290c18 100644 --- a/docs/claude/INDEX.md +++ b/docs/claude/INDEX.md @@ -9,6 +9,7 @@ archived. | File | Description | |---|---| +| `reference/jobs-panel.md` | How a job type or post-processing task appears in the generated Create Job dialog: one pydantic input model per job type (`ami/jobs/schemas.py`), params-only create contract, `GET /jobs/types/`, Field hints (`ami_widget`/`ami_entity`/`ami_advanced`), per-task feature flags, translation via `gettext_lazy`, gating gotcha. Indexed 2026-10-01. Keywords: jobs panel, job types, config_schema, post-processing, params | | `reference/canonical-patterns.md` | Existing helpers/patterns to reuse before writing new ones, with file:line refs (SingleParamSerializer, ProjectMixin, permissions, schemas, fixtures). Keywords: reuse, helpers, conventions, DRF | | `reference/query-patterns.md` | DB model relationship table, composite indexes, prefetch/select_related patterns, full custom QuerySet method catalog, query anti-patterns. Keywords: N+1, indexes, ORM, performance | | `reference/api-stats-pattern.md` | How to add aggregate/leaderboard/chart endpoints (`//stats//`): GenericViewSet + @action, pure querysets in models_future. Keywords: stats, charts, aggregation | diff --git a/docs/claude/reference/jobs-panel.md b/docs/claude/reference/jobs-panel.md new file mode 100644 index 000000000..32a4e5004 --- /dev/null +++ b/docs/claude/reference/jobs-panel.md @@ -0,0 +1,86 @@ +# Jobs panel: how a job type appears in the Create Job dialog + +The Create Job dialog is generated from `GET /api/v2/jobs/types/?project_id=N`. A job type or +post-processing task shows up there, with a working form, once it declares the attributes below. +No frontend change is needed. Design history: branch `feat/jobs-panel-design`, +`docs/claude/planning/2026-09-18-jobs-panel-schema-driven-design.md`. PR #1447. + +## The contract in one paragraph + +Each job type has one pydantic model describing everything a new job of that type takes +(`JobType.config_schema`; post-processing has one per task). `GET /jobs/types/` serves each model's +JSON Schema unchanged, wrapped in `JobTypeDescription` (`ami/jobs/schemas.py`), plus the picker +headings in use (`groups`, `JobGroup` order). The dialog shows one grouped picker: a type without +variants is one entry, a type with variants (post-processing) contributes one entry per variant. The dialog renders +the schema as its form and posts `{"name", "delay", "project_id", "job_type_key", "params": {"config": +{...}}}` (plus `"task"` for post-processing). `JobSerializer.validate` hands `params` to +`JobType.validate_params`, which parses it with the model and checks every id against the project; +config fields named in `JOB_COLUMNS` (`pipeline_id`, `source_image_collection_id`, ...) are also set on +the Job's columns so the jobs list can filter on them. + +## Where things live + +| What | File | +|---|---| +| Per-type input models (`MLJobConfig`, ...) and the response models | `ami/jobs/schemas.py` | +| `JobType.describe`, `validate_params`, `column_ids`, `describe_job_types`, project-scope id checks | `ami/jobs/models.py` | +| `params` validation on create | `ami/jobs/serializers.py` (`JobSerializer.validate`) | +| The gated `types` action | `ami/jobs/views.py` (`JobViewSet.types`) | +| Per-task feature flags | `ami/main/models.py` (`ProjectFeatureFlags`), `BasePostProcessingTask.feature_flag` | +| `run_post_processing_job` for ML data managers | `ami/users/roles.py`, `ami/main/migrations/0096_grant_run_post_processing_to_ml_data_manager.py` | +| Form mapping and payload | `ui/src/components/form/schema-form/` (`schema-to-fields.ts`, `build-job-payload.ts`) | + +## Making a job type creatable + +On the `JobType` subclass: `user_creatable = True`, a `label` and `description` (in `gettext_lazy`), +a `group` (`JobGroup`, the picker heading), and a `config_schema` model in `ami/jobs/schemas.py`. +`label` is what users see in the picker, the jobs list and job details ("Process captures"); `name` +stays the fixed internal name used in logs and stage names. Required inputs are pydantic-required fields; Job +columns are fields named as in `JOB_COLUMNS`. + +## Making a post-processing task appear + +Register it in `POSTPROCESSING_TASKS`, give it `feature_flag` (a `ProjectFeatureFlags` field you +add, default off), a `label`, a `group` and a `description`. Never rename a task's `name`: it is the +lookup key for the task's Algorithm row (`BasePostProcessingTask.__init__`), so a rename starts a new +algorithm. The jobs list names a post-processing job by its task's `label` (`PostProcessingJob.label_for`). When a project turns the flag on, ML data managers and +project managers can run the task with any settings; while it is off the task is hidden, refused on +create through the API, and its jobs cannot be re-run except by a superuser. The Django admin action +still lets superusers start it, so staff can try a method on a project before turning it on. + +The served schema is `config_schema.schema()` unchanged; the same rules apply to the per-type models +in `ami/jobs/schemas.py`. Write it for the form: + +```python +source_image_collection_id: int | None = pydantic.Field( + None, title=_("Capture set"), ami_widget="entity", ami_entity="captures/collections", +) +occurrence_id: int | None = pydantic.Field(None, ami_widget="hidden", ami_entity="occurrences") +cost_threshold: float = pydantic.Field(0.2, title=_("Cost threshold"), description=_("..."), ge=0) +appearance_weight: float = pydantic.Field(1.0, title=_("Appearance weight"), ami_advanced=True) +``` + +- Give every visible field a `title` (pydantic's own is title-cased from the name) and a + `description`, both in `gettext_lazy`: they are served in the request's language. +- `ami_widget="entity"` + `ami_entity=""` (+ `ami_entity_filters`) renders a picker, + and the server checks the id belongs to the job's project (`_entity_queryset`; add the route + there for a new entity, or it is not checked). Check the target viewset can filter by every + `ami_entity_filters` key; DRF ignores unknown filters silently. +- `ami_widget="hidden"` keeps a field out of the form; `ami_advanced=True` puts it under the + collapsed "More settings". +- Use `typing.Literal` for choices, not an `Enum` class (which becomes a `$ref`). +- Cross-field rules (such as "exactly one of capture set or sessions") stay in the model's + validators; their errors show in the dialog's general error block. + +## Where the settings are stored + +`Job.params`, validated and filled with every default on create, then fixed: an update cannot +change it. It is returned on the job detail response, not the list. Ids that are also Job columns +(pipeline, capture set, station, capture) are stored in both places. + +## Gotchas + +- `ObjectPermission.has_permission` returns True for every request, and a `detail=False` action + never reaches the object check. `types` gates itself (authenticated project member or superuser). +- Pydantic is v1 here (`Model.schema()`). A move to v2 changes the served shape (`$defs`, `anyOf` for optionals); update the frontend mapping with it. +- `assertNumQueries` for `types` counts the request's savepoint pair (6 total). diff --git a/ui/AGENTS.md b/ui/AGENTS.md index 461e19975..520852327 100644 --- a/ui/AGENTS.md +++ b/ui/AGENTS.md @@ -11,6 +11,7 @@ in `src/design-system/` — check there before writing a new component. - All user-facing strings go through the translation layer: `translate(STRING.KEY)` from `src/utils/language.ts`. Add new keys to the `STRING` enum and `ENGLISH_STRINGS` map. Never hardcode UI copy in components. +- Exception: labels and help text in server-provided schemas (the Create job dialog's settings fields and job type descriptions) are rendered as received. The server translates them with Django's i18n (`gettext_lazy` on the pydantic `Field`), so they need no `STRING` key. - UI copy uses sentence case: "Taxa list", not "Taxa List". ## Data services & types diff --git a/ui/src/components/form/schema-form/build-job-payload.ts b/ui/src/components/form/schema-form/build-job-payload.ts new file mode 100644 index 000000000..ad2cb6762 --- /dev/null +++ b/ui/src/components/form/schema-form/build-job-payload.ts @@ -0,0 +1,88 @@ +import { + ServerJobType, + ServerJobTypeVariant, +} from 'data-services/models/job-type' +import { parseIntegerList } from 'utils/fieldProcessors' +import { FieldDescriptor, schemaToFields } from './schema-to-fields' + +export interface CreateJobState { + projectId: string + jobType: ServerJobType + variant?: ServerJobTypeVariant + // Names of the rows chosen in pickers, used for the default job name. + pickedLabels?: { [field: string]: string } + configValues: { [field: string]: unknown } + name?: string + delay?: string | number + startNow?: boolean + today?: string +} + +export const coerceValue = (field: FieldDescriptor, raw: unknown): unknown => { + if (raw === undefined || raw === null || raw === '') { + return undefined + } + switch (field.kind) { + case 'integer': + case 'number': + case 'entity': + return Number.isNaN(Number(raw)) ? raw : Number(raw) + case 'integer-list': + return parseIntegerList(`${raw}`) ?? undefined + case 'json': + try { + return JSON.parse(`${raw}`) + } catch { + return raw + } + default: + return raw + } +} + +const collect = ( + fields: FieldDescriptor[], + values: { [field: string]: unknown } +) => { + const result: { [field: string]: unknown } = {} + fields.forEach((field) => { + const value = coerceValue(field, values[field.name]) + if (value !== undefined) { + result[field.name] = value + } + }) + return result +} + +export const buildJobPayload = (state: CreateJobState) => { + const { jobType, variant, projectId } = state + const schema = variant?.config_schema ?? jobType.config_schema + const config = collect(schemaToFields(schema), state.configValues) + + const params = + jobType.variant_key && variant + ? { [jobType.variant_key]: variant.key, config } + : { config } + + const label = variant?.name ?? jobType.name + // The capture set names a job best; otherwise use whatever was picked first. + const pickedLabel = + state.pickedLabels?.source_image_collection_id || + Object.values(state.pickedLabels ?? {}).find(Boolean) + const name = + state.name?.trim() || + `${label} – ${ + pickedLabel ?? state.today ?? new Date().toISOString().slice(0, 10) + }` + + return { + body: { + name, + delay: Number(state.delay) || 0, + project_id: projectId, + job_type_key: jobType.key, + params, + }, + startNow: !!state.startNow, + } +} diff --git a/ui/src/components/form/schema-form/entity-select.tsx b/ui/src/components/form/schema-form/entity-select.tsx new file mode 100644 index 000000000..743a8cdfb --- /dev/null +++ b/ui/src/components/form/schema-form/entity-select.tsx @@ -0,0 +1,82 @@ +import { API_URL } from 'data-services/constants' +import { useAuthorizedQuery } from 'data-services/hooks/auth/useAuthorizedQuery' +import { ServerEntityOption } from 'data-services/models/job-type' +import { Select } from 'nova-ui-kit' +import { STRING, translate } from 'utils/language' + +const PAGE_SIZE = 100 + +interface EntityOption { + id: string + name: string + label: string +} + +const getLabel = (record: ServerEntityOption): string => { + const name = record.name ?? `${record.id}` + return typeof record.source_images_count === 'number' + ? `${name} (${record.source_images_count.toLocaleString()})` + : name +} + +export const EntitySelect = ({ + entity, + entityFilters, + projectId, + value, + label, + placeholder, + onValueChange, +}: { + entity: string + entityFilters?: { [key: string]: string | number | boolean } + projectId: string + value?: string + label: string + placeholder?: string + onValueChange: (value: string | undefined, label?: string) => void +}) => { + const params = new URLSearchParams({ + project_id: projectId, + limit: `${PAGE_SIZE}`, + ...Object.fromEntries( + Object.entries(entityFilters ?? {}).map(([k, v]) => [k, `${v}`]) + ), + }) + const { data, isLoading } = useAuthorizedQuery<{ + results: ServerEntityOption[] + }>({ + queryKey: [entity, 'options', params.toString()], + url: `${API_URL}/${entity}/?${params.toString()}`, + }) + const options: EntityOption[] = (data?.results ?? []).map((record) => ({ + id: `${record.id}`, + name: record.name ?? `${record.id}`, + label: getLabel(record), + })) + const selected = options.some((option) => option.id === value) ? value : '' + + return ( + + onValueChange(id, options.find((option) => option.id === id)?.name) + } + value={selected} + > + + + + + {options.map((option) => ( + + {option.label} + + ))} + + + ) +} diff --git a/ui/src/components/form/schema-form/map-server-errors.ts b/ui/src/components/form/schema-form/map-server-errors.ts new file mode 100644 index 000000000..3254fae35 --- /dev/null +++ b/ui/src/components/form/schema-form/map-server-errors.ts @@ -0,0 +1,40 @@ +// Server validation errors arrive as DRF JSON. Problems with a setting are +// strings of the form ": message" under params.config. +export const mapServerErrors = ( + data: unknown, + { configFields }: { configFields: string[] } +) => { + const fieldErrors: { [formName: string]: string } = {} + const general: string[] = [] + + const asList = (value: unknown): string[] => + (Array.isArray(value) ? value : [value]) + .filter((item) => item !== undefined && item !== null) + .map((item) => (typeof item === 'string' ? item : JSON.stringify(item))) + + if (!data || typeof data !== 'object') { + return { fieldErrors, general } + } + + Object.entries(data as { [key: string]: unknown }).forEach(([key, value]) => { + if (key === 'params' && value && typeof value === 'object') { + Object.entries(value as { [key: string]: unknown }).forEach( + ([paramKey, paramValue]) => { + asList(paramValue).forEach((message) => { + const separator = message.indexOf(': ') + const field = separator > 0 ? message.slice(0, separator) : '' + if (paramKey === 'config' && configFields.includes(field)) { + fieldErrors[`config.${field}`] ??= message.slice(separator + 2) + } else { + general.push(message) + } + }) + } + ) + } else { + general.push(...asList(value)) + } + }) + + return { fieldErrors, general } +} diff --git a/ui/src/components/form/schema-form/schema-field.tsx b/ui/src/components/form/schema-form/schema-field.tsx new file mode 100644 index 000000000..f7f92374e --- /dev/null +++ b/ui/src/components/form/schema-form/schema-field.tsx @@ -0,0 +1,147 @@ +import { Checkbox, Input, InputContent, Select } from 'nova-ui-kit' +import { Control, Controller } from 'react-hook-form' +import { STRING, translate } from 'utils/language' +import { EntitySelect } from './entity-select' +import { FieldDescriptor, validateNumber } from './schema-to-fields' + +// Labels and help text come from the server schema and are shown as received. +export const SchemaField = ({ + control, + field, + formName, + projectId, + onLabelChange, +}: { + control: Control + field: FieldDescriptor + formName: string + projectId: string + onLabelChange?: (label?: string) => void +}) => ( + { + if (field.kind === 'integer' || field.kind === 'number') { + return validateNumber(field, value) + } + if (field.kind === 'json' && value) { + try { + JSON.parse(`${value}`) + } catch { + return translate(STRING.MESSAGE_VALUE_INVALID) + } + } + return undefined + }, + }} + render={({ field: controller, fieldState }) => { + const label = field.required ? `${field.label} *` : field.label + const notSet = field.required + ? undefined + : translate(STRING.JOB_VALUE_NOT_SET) + const error = fieldState.error?.message + + switch (field.kind) { + case 'boolean': + return ( + + + + ) + case 'entity': + return ( + + { + controller.onChange(value) + onLabelChange?.(optionLabel) + }} + /> + + ) + case 'select': + return ( + + + + + + + {field.options?.map((option) => ( + + {option} + + ))} + + + + ) + case 'json': + return ( + +