diff --git a/ami/base/models.py b/ami/base/models.py index 2f245b745..b68e3ad49 100644 --- a/ami/base/models.py +++ b/ami/base/models.py @@ -1,6 +1,7 @@ from django.contrib.auth.models import AbstractUser, AnonymousUser +from django.core.exceptions import FieldDoesNotExist from django.db import models -from django.db.models import Q, QuerySet +from django.db.models import Exists, OuterRef, Q, QuerySet from guardian.shortcuts import get_perms import ami.tasks @@ -39,6 +40,23 @@ def has_many_to_many_project_relation(model: type[models.Model]) -> bool: return False +class PublicScopedModel(models.Model): + """ + Base for M2M-to-project models that can be marked public. for_project() and + visible_for_user() check issubclass(model, PublicScopedModel) for their + public-row bypass, instead of duck-typing an `is_public` attribute that any + model could grow by coincidence and silently pick up that behavior. + """ + + is_public = models.BooleanField( + default=False, + help_text="Public rows are shown to every project, not just the ones linked via 'projects'.", + ) + + class Meta: + abstract = True + + class BaseQuerySet(QuerySet): def visible_for_user(self, user: User | AnonymousUser) -> QuerySet: """ @@ -85,8 +103,37 @@ def visible_for_user(self, user: User | AnonymousUser) -> QuerySet: if not is_anonymous: filter_condition |= Q(**{f"{project_field}owner": user}) | Q(**{f"{project_field}members": user}) + # Public rows (e.g. public TaxaLists) are visible to everyone, draft or not. + if issubclass(model, PublicScopedModel): + filter_condition |= Q(is_public=True) + return self.filter(filter_condition).distinct() + def for_project(self, project: models.Model, include_public: bool = True) -> QuerySet: + """ + Filter to rows in the model's M2M ``projects`` field for the given project, + plus every public row when the model defines ``is_public`` and ``include_public`` + is set. + + Membership is checked with an ``Exists`` subquery against the M2M through table + instead of filtering on ``projects=project`` directly, so a row linked to the + project through multiple paths cannot appear twice and no ``.distinct()`` is + needed downstream. + """ + model = self.model + try: + field = model._meta.get_field("projects") + except FieldDoesNotExist: + field = None + if not isinstance(field, models.ManyToManyField): + raise TypeError(f"{model.__name__} has no ManyToMany 'projects' field; for_project() is not applicable.") + + condition = Q(Exists(model._default_manager.filter(pk=OuterRef("pk"), projects=project))) + if include_public and issubclass(model, PublicScopedModel): + condition |= Q(is_public=True) + + return self.filter(condition) + class BaseModel(models.Model): """ """ diff --git a/ami/base/permissions.py b/ami/base/permissions.py index 328aa8384..a611f2dc8 100644 --- a/ami/base/permissions.py +++ b/ami/base/permissions.py @@ -77,6 +77,67 @@ def add_collection_level_permissions(user: User | None, response_data: dict, mod return response_data +def user_can_manage_public(user: AbstractBaseUser | AnonymousUser, model_or_instance) -> bool: + """ + A superuser, or a user holding .manage_public_ for + the given model (or an instance of it) — the platform permission gating + write access to a public row, in place of project membership or staff status. + """ + if not user or not user.is_authenticated: + return False + if user.is_superuser: # type: ignore[union-attr] + return True + meta = model_or_instance._meta + return user.has_perm(f"{meta.app_label}.manage_public_{meta.model_name}") # type: ignore[union-attr] + + +def check_public_scoped_write_permission(user, instance, non_public_fallback) -> bool: + """ + True if `user` may write to `instance`: a public row needs the platform + manage_public_ permission (superusers always pass, via + user_can_manage_public); a non-public row falls back to `non_public_fallback()`, + the model's own write rule (project membership, active-staff status, ...). + """ + if getattr(instance, "is_public", False): + return user_can_manage_public(user, instance) + return non_public_fallback() + + +def check_taxalist_write_permission(user, taxa_list, project) -> bool: + """Thin alias kept so branches stacked on this one still import this name.""" + return check_public_scoped_write_permission( + user, taxa_list, lambda: bool(user.is_superuser or (project and project.members.filter(pk=user.pk).exists())) + ) + + +def check_processingservice_write_permission(user, processing_service) -> bool: + """Thin alias kept so branches stacked on this one still import this name.""" + return check_public_scoped_write_permission( + user, processing_service, lambda: bool(user.is_superuser or is_active_staff(user)) + ) + + +def add_processingservice_permissions(user, instance, response_data: dict) -> dict: + """ + Add update/delete to user_permissions for a ProcessingService. + + Unlike add_m2m_object_permissions, this skips the M2M membership check and + the guardian lookup entirely: no per-project *_processingservice guardian + permission exists anywhere in this codebase (see Project.Permissions and + ami/users/roles.py), so that branch could only ever fire for a superuser, + which a plain attribute check already covers for free. A public instance + still checks the platform manage_public_processingservice permission. + """ + perms = set(response_data.get("user_permissions", [])) + if getattr(instance, "is_public", False): + if user_can_manage_public(user, instance): + perms.update(["update", "delete"]) + elif user.is_superuser: + perms.update(["update", "delete"]) + response_data["user_permissions"] = list(perms) + return response_data + + def add_m2m_object_permissions( user, instance, project, response_data: dict, project_perms: set[str] | None = None ) -> dict: @@ -96,6 +157,9 @@ def add_m2m_object_permissions( (user, project) pass in `guardian.get_perms(user, project)` once instead of once per instance; pass None to look it up here as before. + A public instance is the one exception: its update/delete permissions come + from the model's manage_public_ permission, not project membership. + This is a temporary approach for the M2M permission gap described in #1120. Once that issue is resolved, this should be replaced by a generic permission class (Pattern B: Bare M2M) that handles TaxaList, Taxon, ProcessingService, @@ -103,6 +167,12 @@ class (Pattern B: Bare M2M) that handles TaxaList, Taxon, ProcessingService, """ perms = set(response_data.get("user_permissions", [])) + if getattr(instance, "is_public", False): + if user_can_manage_public(user, instance): + perms.update(["update", "delete"]) + response_data["user_permissions"] = list(perms) + return response_data + if not project: response_data["user_permissions"] = list(perms) return response_data @@ -160,6 +230,92 @@ def has_permission(self, request, view): return project.members.filter(pk=request.user.pk).exists() +class _BaseGateOrPublicManager(permissions.BasePermission): + """ + Shared shape for M2M-to-project models with a public flag: safe methods are + open to everyone; unsafe methods need the model's own base gate (project + membership, active-staff status, ...) or the manage_public_ platform + permission. `exclude_create_from_bypass` forces a plain "create a new row" + action through the base gate only, since there's no object yet to tell + whether it will be public. Subclasses implement get_model() and + get_base_gate() — get_model() does a local import to avoid a module-level + circular import between this file and the app that owns the model. + """ + + exclude_create_from_bypass = False + + def get_model(self): + raise NotImplementedError + + def get_base_gate(self, request, view) -> bool: + raise NotImplementedError + + def has_permission(self, request, view): + if request.method in permissions.SAFE_METHODS: + return True + + if not request.user or not request.user.is_authenticated: + return False + + if request.user.is_superuser: # type: ignore[union-attr] + return True + + if self.exclude_create_from_bypass and getattr(view, "action", None) == "create": + return self.get_base_gate(request, view) + + if user_can_manage_public(request.user, self.get_model()): + return True + + return self.get_base_gate(request, view) + + def has_object_permission(self, request, view, obj): + if request.method in permissions.SAFE_METHODS: + return True + return check_public_scoped_write_permission(request.user, obj, lambda: self.get_base_gate(request, view)) + + +class IsProjectMemberOrPublicListManager(_BaseGateOrPublicManager): + """ + Used by the nested add/remove-taxon route: serves both public and + project-scoped lists, and has no object to check yet at has_permission() + time. + """ + + def get_model(self): + from ami.main.models import TaxaList + + return TaxaList + + def get_base_gate(self, request, view): + get_active_project = getattr(view, "get_active_project", None) + project = get_active_project() if get_active_project else None + return bool(project and project.members.filter(pk=request.user.pk).exists()) + + +class IsProjectMemberOrPublicListManagerOrReadOnly(IsProjectMemberOrPublicListManager): + """For TaxaListViewSet: creating a brand-new list always needs real project membership.""" + + exclude_create_from_bypass = True + + +class IsActiveStaffOrPublicManager(_BaseGateOrPublicManager): + """Used by ProcessingServiceViewSet's non-create actions and any future nested route.""" + + def get_model(self): + from ami.ml.models.processing_service import ProcessingService + + return ProcessingService + + def get_base_gate(self, request, view): + return is_active_staff(request.user) + + +class IsActiveStaffOrPublicManagerOrReadOnly(IsActiveStaffOrPublicManager): + """For ProcessingServiceViewSet: creating a brand-new service always needs active-staff status.""" + + exclude_create_from_bypass = True + + class ObjectPermission(permissions.BasePermission): """ Generic permission class that delegates to the model's `check_permission(user, action)` method. diff --git a/ami/base/views.py b/ami/base/views.py index 482496b38..79631136a 100644 --- a/ami/base/views.py +++ b/ami/base/views.py @@ -94,3 +94,17 @@ def get_active_project(self) -> Project | None: raise Http404("Project not found.") return project + + def get_include_public(self) -> bool: + """ + The ?include_public query param: whether to include public rows alongside + a model's own project-scoped ones. Defaults to true; an invalid value + raises ValidationError (400) via SingleParamSerializer. + """ + from ami.base.serializers import SingleParamSerializer + + return SingleParamSerializer[bool].clean( + param_name="include_public", + field=serializers.BooleanField(required=False, default=True), + data=self.request.query_params, + ) diff --git a/ami/main/admin.py b/ami/main/admin.py index 55c84701f..88bf8169c 100644 --- a/ami/main/admin.py +++ b/ami/main/admin.py @@ -736,7 +736,7 @@ def parent_names(self, obj) -> str: class TaxaListAdmin(admin.ModelAdmin[TaxaList]): """Admin panel example for ``TaxaList`` model.""" - list_display = ("name", "taxa_count", "created_at", "updated_at") + list_display = ("name", "is_public", "taxa_count", "created_at", "updated_at") def taxa_count(self, obj) -> int: return obj.taxa.count() @@ -746,7 +746,7 @@ def taxa_count(self, obj) -> int: "projects", ) - list_filter = ("projects",) + list_filter = ("is_public", "projects") @admin.register(Device) diff --git a/ami/main/api/schemas.py b/ami/main/api/schemas.py index c9835b2fb..5b18dfb4b 100644 --- a/ami/main/api/schemas.py +++ b/ami/main/api/schemas.py @@ -13,3 +13,11 @@ required=False, type=int, ) + +include_public_doc_param = OpenApiParameter( + name="include_public", + description="Include rows that are public (available to every project), not just the ones " + "belonging to this project. Defaults to true.", + required=False, + type=bool, +) diff --git a/ami/main/api/serializers.py b/ami/main/api/serializers.py index 1ca588980..2bb655b8f 100644 --- a/ami/main/api/serializers.py +++ b/ami/main/api/serializers.py @@ -708,6 +708,7 @@ class TaxaListSerializer(DefaultSerializer): taxa = serializers.SerializerMethodField() taxa_count = serializers.SerializerMethodField() projects = serializers.SerializerMethodField() + is_public = serializers.BooleanField(read_only=True) class Meta: model = TaxaList @@ -718,6 +719,7 @@ class Meta: "taxa", "taxa_count", "projects", + "is_public", "created_at", "updated_at", ] @@ -763,11 +765,19 @@ def get_permissions(self, instance, instance_data): def get_projects(self, obj): """ - Return list of project IDs this taxa list belongs to, sorted for a - deterministic response. Reads the `projects` prefetched by + Return the ids of this list's linked projects that are visible to the + requester, sorted for a deterministic response. A public list can be linked + to a draft project it's otherwise not visible in; without this filter, an + outsider retrieving the public list would learn that draft project's id even + though they can't see the project itself. Reads the `projects` prefetched by TaxaListViewSet.get_queryset instead of querying per row. """ - return sorted(project.pk for project in obj.projects.all()) + request = self.context["request"] + if not hasattr(self, "_visible_project_ids"): + self._visible_project_ids = set( + Project.objects.visible_for_user(request.user).values_list("id", flat=True) + ) + return sorted(project.pk for project in obj.projects.all() if project.pk in self._visible_project_ids) class TaxaListTaxonInputSerializer(serializers.Serializer): diff --git a/ami/main/api/views.py b/ami/main/api/views.py index 99709eb29..e303f09f7 100644 --- a/ami/main/api/views.py +++ b/ami/main/api/views.py @@ -29,10 +29,16 @@ from ami.base.metadata import ResponseSchemaMetadata from ami.base.models import BaseQuerySet from ami.base.pagination import LimitOffsetPaginationWithPermissions -from ami.base.permissions import IsActiveStaffOrReadOnly, IsProjectMemberOrReadOnly, ObjectPermission +from ami.base.permissions import ( + IsActiveStaffOrReadOnly, + IsProjectMemberOrPublicListManager, + IsProjectMemberOrPublicListManagerOrReadOnly, + ObjectPermission, + check_taxalist_write_permission, +) from ami.base.serializers import FilterParamsSerializer, SingleParamSerializer from ami.base.views import ProjectMixin -from ami.main.api.schemas import limit_doc_param, project_id_doc_param +from ami.main.api.schemas import include_public_doc_param, limit_doc_param, project_id_doc_param from ami.main.api.serializers import TagSerializer from ami.main.models_future.identifications import create_identifications_batch, resolve_occurrences from ami.main.models_future.occurrence import model_agreement_for_project, top_identifiers_for_project @@ -2235,7 +2241,7 @@ class TaxaListViewSet(DefaultViewSet, ProjectMixin): "created_at", "updated_at", ] - permission_classes = [IsProjectMemberOrReadOnly] + permission_classes = [IsProjectMemberOrPublicListManagerOrReadOnly] require_project = True def get_queryset(self): @@ -2247,9 +2253,17 @@ def get_queryset(self): # database once per row. qs = qs.prefetch_related("projects") project = self.get_active_project() - if project: - return qs.filter(projects=project) - return qs + if not project: + return qs + # include_public governs the list action's default scope, not whether a + # specific row is reachable: a detail/update/delete on a public list must + # still resolve it even under ?include_public=false. + include_public = self.get_include_public() if self.action == "list" else True + return qs.for_project(project, include_public=include_public) + + @extend_schema(parameters=[project_id_doc_param, include_public_doc_param]) + def list(self, request, *args, **kwargs): + return super().list(request, *args, **kwargs) def perform_create(self, serializer): """ @@ -2274,18 +2288,32 @@ class TaxaListTaxonViewSet(viewsets.GenericViewSet, ProjectMixin): """ serializer_class = TaxaListTaxonSerializer - permission_classes = [IsProjectMemberOrReadOnly] + permission_classes = [IsProjectMemberOrPublicListManager] require_project = True def get_taxa_list(self): - """Get the parent taxa list from URL parameters, scoped to the active project.""" + """Get the parent taxa list, scoped to the active project or public.""" taxa_list_id = self.kwargs.get("taxalist_pk") project = self.get_active_project() try: - return TaxaList.objects.get(pk=taxa_list_id, projects=project) + return ( + TaxaList.objects.visible_for_user(self.request.user) + .for_project(project, include_public=True) + .get(pk=taxa_list_id) + ) except TaxaList.DoesNotExist: raise api_exceptions.NotFound("Taxa list not found.") from None + def check_write_permission(self, taxa_list): + """ + Re-check against the actual list: IsProjectMemberOrPublicListManager only + gates coarsely at has_permission() time, before the target list (and its + is_public flag) is known. + """ + project = self.get_active_project() + if not check_taxalist_write_permission(self.request.user, taxa_list, project): + raise api_exceptions.PermissionDenied("You do not have permission to modify this taxa list.") + def get_queryset(self): """Return taxa in the specified taxa list.""" taxa_list = self.get_taxa_list() @@ -2294,6 +2322,7 @@ def get_queryset(self): def create(self, request, taxalist_pk=None): """Add a taxon to the taxa list.""" taxa_list = self.get_taxa_list() + self.check_write_permission(taxa_list) # Validate input input_serializer = TaxaListTaxonInputSerializer(data=request.data) @@ -2322,6 +2351,7 @@ def delete_by_taxon(self, request, taxalist_pk=None, taxon_id=None): DELETE /taxa/lists/{taxa_list_id}/taxa/{taxon_id}/ """ taxa_list = self.get_taxa_list() + self.check_write_permission(taxa_list) # Check if taxon exists in list if not taxa_list.taxa.filter(pk=taxon_id).exists(): diff --git a/ami/main/management/commands/import_taxa.py b/ami/main/management/commands/import_taxa.py index a1f9cf49b..31f393677 100644 --- a/ami/main/management/commands/import_taxa.py +++ b/ami/main/management/commands/import_taxa.py @@ -236,8 +236,9 @@ def handle(self, *args, **options): else: list_name = pathlib.Path(fname).stem - # Uses get_or_create_for_project with project=None to create a global list - taxalist, created = TaxaList.objects.get_or_create_for_project(name=list_name, project=None) + # Uses get_or_create_for_project with project=None to create a public list + # available to every project. + taxalist, created = TaxaList.objects.get_or_create_for_project(name=list_name, project=None, is_public=True) if created: self.stdout.write(self.style.SUCCESS('Successfully created taxa list "%s"' % taxalist)) diff --git a/ami/main/management/commands/update_taxa.py b/ami/main/management/commands/update_taxa.py index 92779fcbc..6564f21e0 100644 --- a/ami/main/management/commands/update_taxa.py +++ b/ami/main/management/commands/update_taxa.py @@ -127,11 +127,14 @@ def handle(self, *args, **options): incoming_taxa = read_csv(fname) # Get or create taxa list if specified - # Uses get_or_create_for_project with project=None to create a global list + # Uses get_or_create_for_project with project=None to create a public list + # available to every project. taxalist = None if options["list"]: list_name = options["list"] - taxalist, created = TaxaList.objects.get_or_create_for_project(name=list_name, project=None) + taxalist, created = TaxaList.objects.get_or_create_for_project( + name=list_name, project=None, is_public=True + ) if created: self.stdout.write(self.style.SUCCESS(f"Created new taxa list '{list_name}'")) else: diff --git a/ami/main/migrations/0098_taxalist_is_public.py b/ami/main/migrations/0098_taxalist_is_public.py new file mode 100644 index 000000000..8140accde --- /dev/null +++ b/ami/main/migrations/0098_taxalist_is_public.py @@ -0,0 +1,54 @@ +""" +``TaxaList.is_public`` marks a list as available to every project, not just the +ones in its ``projects`` M2M. This backfills it: every existing list with no +project becomes public, except a "Taxa returned by " list — the +running set of taxa that algorithm has returned as a top prediction — which +stays non-public, since what a project should see or do with that list is +still undecided. ``projects`` becomes ``blank=True`` since a public list no +longer needs one. +""" + +from django.db import migrations, models +from django.db.models import Count + + +def backfill_is_public(apps, schema_editor): + TaxaList = apps.get_model("main", "TaxaList") + zero_project_lists = TaxaList.objects.annotate(project_count=Count("projects")).filter(project_count=0) + zero_project_lists.exclude(name__startswith="Taxa returned by").update(is_public=True) + + +def reverse_noop(apps, schema_editor): + # Not reversible: we don't track which rows were public before this migration. + pass + + +class Migration(migrations.Migration): + dependencies = [ + ("main", "0097_detection_and_classification_job_indexes"), + ] + + operations = [ + migrations.AlterModelOptions( + name="taxalist", + options={ + "ordering": ["-created_at"], + "permissions": [("manage_public_taxalist", "Can create, edit and delete public taxa lists")], + "verbose_name_plural": "Taxa Lists", + }, + ), + migrations.AddField( + model_name="taxalist", + name="is_public", + field=models.BooleanField( + default=False, + help_text="Public rows are shown to every project, not just the ones linked via 'projects'.", + ), + ), + migrations.AlterField( + model_name="taxalist", + name="projects", + field=models.ManyToManyField(blank=True, related_name="taxa_lists", to="main.project"), + ), + migrations.RunPython(backfill_is_public, reverse_noop), + ] diff --git a/ami/main/models.py b/ami/main/models.py index a4c0c82f1..78212a680 100644 --- a/ami/main/models.py +++ b/ami/main/models.py @@ -34,7 +34,7 @@ import ami.tasks import ami.utils from ami.base.fields import DateStringField -from ami.base.models import BaseModel, BaseQuerySet +from ami.base.models import BaseModel, BaseQuerySet, PublicScopedModel from ami.main import charts from ami.main.models_future.filters import ( build_occurrence_default_filters_q, @@ -4769,16 +4769,20 @@ def save(self, update_calculated_fields=True, *args, **kwargs): class TaxaListQuerySet(BaseQuerySet): def get_or_create_for_project( - self, name: str, project: "Project | None" = None, **defaults + self, name: str, project: "Project | None" = None, is_public: bool = False, **defaults ) -> tuple["TaxaList", bool]: """ Get or create a TaxaList with uniqueness scoped to project. - - If project is None: looks for/creates a global list (no project associations) - - If project is provided: looks for/creates a list associated with that project + - project is given: looks for/creates a list associated with that project + - project is None, is_public=True: looks for/creates the public list with this name + - project is None, is_public=False: looks for/creates a hidden list with no project :param name: Name of the taxa list. - :param project: Project to scope the list to, or None for a global list. + :param project: Project to scope the list to, or None for a list with no project. + :param is_public: When ``project`` is None, whether the list is the public list of this + name or a hidden one (e.g. a per-algorithm category-map list). Raises when combined + with a ``project``, since a project-scoped list can't also be a public one. :param defaults: Extra field values applied only when creating a new list (ignored on the get path, matching Django's ``get_or_create`` semantics). @@ -4789,18 +4793,28 @@ def get_or_create_for_project( Returns: Tuple of (TaxaList, created: bool) """ - if project is None: - # Global list: find list with this name that has no project associations - qs = self.filter(name=name).annotate(project_count=models.Count("projects")).filter(project_count=0) - else: + if project is not None and is_public: + raise ValueError("get_or_create_for_project() cannot create a public list scoped to a single project.") + + if project is not None: # Project-specific: find list with this name in this project qs = self.filter(name=name, projects=project) + else: + # No project: find a list with this name and no project associations, + # matching the requested public/hidden state. + qs = ( + self.filter(name=name, is_public=is_public) + .annotate(project_count=models.Count("projects")) + .filter(project_count=0) + ) try: return qs.get(), False except self.model.DoesNotExist: with transaction.atomic(): - taxa_list = self.create(name=name, **defaults) + # is_public is guaranteed False here whenever project is set — the guard above + # already rejects the combination that would make this ambiguous. + taxa_list = self.create(name=name, is_public=is_public, **defaults) if project: taxa_list.projects.add(project) return taxa_list, True @@ -4816,20 +4830,21 @@ class TaxaListManager(models.Manager.from_queryset(TaxaListQuerySet)): @final -class TaxaList(BaseModel): +class TaxaList(BaseModel, PublicScopedModel): """A checklist of taxa""" name = models.CharField(max_length=255) description = models.TextField(blank=True) taxa = models.ManyToManyField(Taxon, related_name="lists") - projects = models.ManyToManyField("Project", related_name="taxa_lists") + projects = models.ManyToManyField("Project", related_name="taxa_lists", blank=True) objects: TaxaListManager = TaxaListManager() class Meta: ordering = ["-created_at"] verbose_name_plural = "Taxa Lists" + permissions = [("manage_public_taxalist", "Can create, edit and delete public taxa lists")] @final diff --git a/ami/main/tests.py b/ami/main/tests.py index 27785fe7f..b1b3c46f9 100644 --- a/ami/main/tests.py +++ b/ami/main/tests.py @@ -6,7 +6,7 @@ from unittest import mock from django.conf import settings -from django.contrib.auth.models import AnonymousUser +from django.contrib.auth.models import AnonymousUser, Permission from django.core.files.uploadedfile import SimpleUploadedFile from django.db import IntegrityError, connection, models from django.test import TestCase, override_settings @@ -5545,6 +5545,537 @@ def test_defaults_applied_on_create_only(self): self.assertEqual(taxa_list_again.pk, taxa_list.pk) self.assertEqual(taxa_list_again.description, "Initial description") + def test_creates_new_public_list(self): + taxa_list, created = TaxaList.objects.get_or_create_for_project( + name="Public Moths", project=None, is_public=True + ) + self.assertTrue(created) + self.assertTrue(taxa_list.is_public) + self.assertEqual(taxa_list.projects.count(), 0) + + def test_retrieves_existing_public_list(self): + existing = TaxaList.objects.create(name="Public Moths", is_public=True) + taxa_list, created = TaxaList.objects.get_or_create_for_project( + name="Public Moths", project=None, is_public=True + ) + self.assertFalse(created) + self.assertEqual(taxa_list.pk, existing.pk) + + def test_public_and_non_public_no_project_lists_with_same_name_do_not_collide(self): + """A public list and a hidden (non-public, no-project) list can share a name.""" + public_list, public_created = TaxaList.objects.get_or_create_for_project( + name="Moths No Project", project=None, is_public=True + ) + hidden_list, hidden_created = TaxaList.objects.get_or_create_for_project( + name="Moths No Project", project=None, is_public=False + ) + self.assertTrue(public_created) + self.assertTrue(hidden_created) + self.assertNotEqual(public_list.pk, hidden_list.pk) + self.assertTrue(public_list.is_public) + self.assertFalse(hidden_list.is_public) + + def test_project_and_is_public_together_raises(self): + """ + project and is_public=True are contradictory: a list scoped to one project + cannot also be the platform's public list of that name. Silently ignoring + is_public here (as the create path already does) would hide the caller's + mistake instead of surfacing it. + """ + with self.assertRaises(ValueError): + TaxaList.objects.get_or_create_for_project(name="Moths", project=self.project_a, is_public=True) + + +class TaxaListForProjectQuerySetTestCase(TestCase): + """Direct unit tests for BaseQuerySet.for_project(), independent of the API layer.""" + + def setUp(self): + self.project_a = Project.objects.create(name="FP Project A") + self.project_b = Project.objects.create(name="FP Project B") + self.scoped = TaxaList.objects.create(name="FP Scoped") + self.scoped.projects.add(self.project_a) + self.public = TaxaList.objects.create(name="FP Public", is_public=True) + self.hidden = TaxaList.objects.create(name="FP Hidden") # is_public=False, no projects + + def test_returns_project_scoped_and_public_by_default(self): + ids = set(TaxaList.objects.for_project(self.project_a).values_list("pk", flat=True)) + self.assertEqual(ids, {self.scoped.pk, self.public.pk}) + + def test_excludes_public_when_include_public_false(self): + ids = set(TaxaList.objects.for_project(self.project_a, include_public=False).values_list("pk", flat=True)) + self.assertEqual(ids, {self.scoped.pk}) + + def test_excludes_lists_scoped_to_other_projects(self): + ids = set(TaxaList.objects.for_project(self.project_b).values_list("pk", flat=True)) + self.assertNotIn(self.scoped.pk, ids) + + def test_public_list_linked_to_project_is_not_duplicated(self): + """ + A public list linked to three projects appears once when queried for one of + them. A join-based filter (Q(projects=project) | Q(is_public=True)) would + produce one joined row per linked project, and is_public=True is true on + every one of those rows regardless of which project_id it carries, so all + three would pass the WHERE clause without .distinct() — three projects is + the minimum that demonstrates this, since is_public alone can't distinguish + the row actually matching `project` from the other two. + """ + project_c = Project.objects.create(name="FP Project C") + self.public.projects.add(self.project_a, self.project_b, project_c) + rows = list(TaxaList.objects.for_project(self.project_a).filter(pk=self.public.pk)) + self.assertEqual(len(rows), 1) + + def test_raises_on_model_without_m2m_projects_field(self): + with self.assertRaises(TypeError): + list(Event.objects.for_project(self.project_a)) + + +class TaxaListIsPublicBackfillTestCase(TestCase): + """Unit test for the 0096 migration's is_public backfill rule.""" + + def _backfill(self): + from importlib import import_module + + from django.apps import apps as real_apps + + mod = import_module("ami.main.migrations.0098_taxalist_is_public") + mod.backfill_is_public(real_apps, None) + + def test_zero_project_list_becomes_public(self): + taxa_list = TaxaList.objects.create(name="No Project List") + self._backfill() + taxa_list.refresh_from_db() + self.assertTrue(taxa_list.is_public) + + def test_project_scoped_list_stays_non_public(self): + project = Project.objects.create(name="Backfill Project") + taxa_list = TaxaList.objects.create(name="Scoped List") + taxa_list.projects.add(project) + self._backfill() + taxa_list.refresh_from_db() + self.assertFalse(taxa_list.is_public) + + def test_algorithm_category_map_list_stays_non_public(self): + taxa_list = TaxaList.objects.create(name="Taxa returned by Some Algorithm") + self._backfill() + taxa_list.refresh_from_db() + self.assertFalse(taxa_list.is_public) + + +class TaxaListPublicPermissionsTestCase(TestCase): + """Permission matrix for public vs. project-scoped TaxaLists. + + The key regression this protects against: a project member passing their own + project_id must not be able to modify, delete, or change the taxa of a public + list. Only the manage_public_taxalist platform permission (or a superuser) can. + A project-scoped list keeps its existing member-can-write behavior. + """ + + def setUp(self): + self.owner = User.objects.create_user(email="owner-pub@example.com", password="testpass") + self.member = User.objects.create_user(email="member-pub@example.com", password="testpass") + self.non_member = User.objects.create_user(email="nonmember-pub@example.com", password="testpass") + self.superuser = User.objects.create_superuser(email="super-pub@example.com", password="testpass") + self.public_manager = User.objects.create_user(email="manager-pub@example.com", password="testpass") + perm = Permission.objects.get(codename="manage_public_taxalist", content_type__app_label="main") + self.public_manager.user_permissions.add(perm) + + self.project = Project.objects.create(name="Test Project", owner=self.owner) + self.project.members.add(self.member) + + self.public_list = TaxaList.objects.create(name="Public List", is_public=True) + self.scoped_list = TaxaList.objects.create(name="Scoped List") + self.scoped_list.projects.add(self.project) + + self.taxon = Taxon.objects.create(name="Test Taxon", rank="SPECIES") + + self.client = APIClient() + + def _detail_url(self, taxa_list): + return f"/api/v2/taxa/lists/{taxa_list.pk}/?project_id={self.project.pk}" + + def _taxa_url(self, taxa_list): + return f"/api/v2/taxa/lists/{taxa_list.pk}/taxa/?project_id={self.project.pk}" + + def _taxon_detail_url(self, taxa_list, taxon): + return f"/api/v2/taxa/lists/{taxa_list.pk}/taxa/{taxon.pk}/?project_id={self.project.pk}" + + # -- Update -- + + def test_member_cannot_update_public_list(self): + self.client.force_authenticate(self.member) + response = self.client.patch(self._detail_url(self.public_list), {"name": "Hacked"}) + self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) + + def test_non_member_cannot_update_public_list(self): + self.client.force_authenticate(self.non_member) + response = self.client.patch(self._detail_url(self.public_list), {"name": "Hacked"}) + self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) + + def test_public_list_manager_can_update_public_list(self): + self.client.force_authenticate(self.public_manager) + response = self.client.patch(self._detail_url(self.public_list), {"name": "Renamed"}) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + def test_superuser_can_update_public_list(self): + self.client.force_authenticate(self.superuser) + response = self.client.patch(self._detail_url(self.public_list), {"name": "Renamed by super"}) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + def test_member_can_still_update_scoped_list(self): + self.client.force_authenticate(self.member) + response = self.client.patch(self._detail_url(self.scoped_list), {"name": "Renamed"}) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + def test_member_can_update_own_list_while_user_permissions_reports_none(self): + """ + Known gap between the write gate and the reported permissions (#1120): the + member write gate is project membership, checked directly by + IsProjectMemberOrPublicListManagerOrReadOnly. The user_permissions field + instead comes from add_m2m_object_permissions's guardian lookup, and no + per-project update_taxalist/delete_taxalist guardian grant exists for plain + membership — so a member who can successfully PATCH this list is also told, + in the same response, that they have no update permission on it. This pins + that mismatch as current, known behavior rather than a change to notice by + surprise. + """ + self.client.force_authenticate(self.member) + response = self.client.patch(self._detail_url(self.scoped_list), {"name": "Renamed Again"}) + self.assertEqual(response.status_code, status.HTTP_200_OK) + self.assertNotIn("update", response.json()["user_permissions"]) + + def test_public_list_manager_without_membership_cannot_update_scoped_list(self): + """Holding manage_public_taxalist does not grant control over project-scoped lists.""" + self.client.force_authenticate(self.public_manager) + response = self.client.patch(self._detail_url(self.scoped_list), {"name": "Hacked"}) + self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) + + # -- Delete -- + + def test_member_cannot_delete_public_list(self): + self.client.force_authenticate(self.member) + response = self.client.delete(self._detail_url(self.public_list)) + self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) + + def test_public_list_manager_can_delete_public_list(self): + self.client.force_authenticate(self.public_manager) + response = self.client.delete(self._detail_url(self.public_list)) + self.assertEqual(response.status_code, status.HTTP_204_NO_CONTENT) + + def test_public_list_manager_can_delete_public_list_with_include_public_false(self): + """ + include_public=false must not make get_queryset() 404 the very row being + deleted: it governs the list action's default scope, not whether a + public row can be looked up for a detail action. + """ + self.client.force_authenticate(self.public_manager) + url = f"{self._detail_url(self.public_list)}&include_public=false" + response = self.client.delete(url) + self.assertEqual(response.status_code, status.HTTP_204_NO_CONTENT) + + def test_member_can_still_delete_scoped_list(self): + self.client.force_authenticate(self.member) + response = self.client.delete(self._detail_url(self.scoped_list)) + self.assertEqual(response.status_code, status.HTTP_204_NO_CONTENT) + + # -- Retrieve / list visibility -- + + def test_anonymous_can_retrieve_public_list(self): + response = self.client.get(self._detail_url(self.public_list)) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + def test_non_member_can_see_public_list(self): + """A public list is visible even to a user with no connection to the active project. + + The project here is non-draft, so its own scoped_list is also visible to + anyone by the platform's general draft-visibility rule — that is separate + from is_public and is covered by TaxaListDraftProjectVisibilityTestCase. + """ + self.client.force_authenticate(self.non_member) + response = self.client.get(f"/api/v2/taxa/lists/?project_id={self.project.pk}") + self.assertEqual(response.status_code, status.HTTP_200_OK) + ids = {row["id"] for row in response.json()["results"]} + self.assertIn(self.public_list.pk, ids) + + def test_is_public_is_reported_in_the_response(self): + self.client.force_authenticate(self.member) + response = self.client.get(self._detail_url(self.public_list)) + self.assertTrue(response.json()["is_public"]) + + def test_is_public_cannot_be_set_through_the_api(self): + self.client.force_authenticate(self.member) + self.client.patch(self._detail_url(self.scoped_list), {"is_public": True}) + self.scoped_list.refresh_from_db() + self.assertFalse(self.scoped_list.is_public) + + def test_user_permissions_include_update_delete_for_public_manager_only(self): + self.client.force_authenticate(self.member) + response = self.client.get(self._detail_url(self.public_list)) + self.assertNotIn("update", response.json()["user_permissions"]) + + self.client.force_authenticate(self.public_manager) + response = self.client.get(self._detail_url(self.public_list)) + perms = response.json()["user_permissions"] + self.assertIn("update", perms) + self.assertIn("delete", perms) + + # -- Add / remove taxon (nested route) -- + + def test_member_cannot_add_taxon_to_public_list(self): + self.client.force_authenticate(self.member) + response = self.client.post(self._taxa_url(self.public_list), {"taxon_id": self.taxon.pk}) + self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) + self.assertFalse(self.public_list.taxa.filter(pk=self.taxon.pk).exists()) + + def test_public_list_manager_can_add_taxon_to_public_list(self): + self.client.force_authenticate(self.public_manager) + response = self.client.post(self._taxa_url(self.public_list), {"taxon_id": self.taxon.pk}) + self.assertEqual(response.status_code, status.HTTP_201_CREATED) + self.assertTrue(self.public_list.taxa.filter(pk=self.taxon.pk).exists()) + + def test_member_cannot_remove_taxon_from_public_list(self): + self.public_list.taxa.add(self.taxon) + self.client.force_authenticate(self.member) + response = self.client.delete(self._taxon_detail_url(self.public_list, self.taxon)) + self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) + self.assertTrue(self.public_list.taxa.filter(pk=self.taxon.pk).exists()) + + def test_public_list_manager_can_remove_taxon_from_public_list(self): + self.public_list.taxa.add(self.taxon) + self.client.force_authenticate(self.public_manager) + response = self.client.delete(self._taxon_detail_url(self.public_list, self.taxon)) + self.assertEqual(response.status_code, status.HTTP_204_NO_CONTENT) + self.assertFalse(self.public_list.taxa.filter(pk=self.taxon.pk).exists()) + + def test_member_can_still_add_taxon_to_scoped_list(self): + self.client.force_authenticate(self.member) + response = self.client.post(self._taxa_url(self.scoped_list), {"taxon_id": self.taxon.pk}) + self.assertEqual(response.status_code, status.HTTP_201_CREATED) + + def test_public_list_manager_without_membership_cannot_add_taxon_to_scoped_list(self): + self.client.force_authenticate(self.public_manager) + response = self.client.post(self._taxa_url(self.scoped_list), {"taxon_id": self.taxon.pk}) + self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) + + +class TaxaListIncludePublicParamTestCase(TestCase): + """?include_public toggles whether public lists appear alongside a project's own lists.""" + + def setUp(self): + self.user = User.objects.create_user(email="scope-user@example.com", password="testpass") + self.project = Project.objects.create(name="Scope Project", owner=self.user) + self.other_project = Project.objects.create(name="Other Scope Project") + + self.scoped_list = TaxaList.objects.create(name="Scoped") + self.scoped_list.projects.add(self.project) + self.public_list = TaxaList.objects.create(name="Public", is_public=True) + # A public list can also be linked to an unrelated project without appearing twice. + self.public_list.projects.add(self.other_project) + + self.hidden_list = TaxaList.objects.create(name="Hidden with no project") + + self.client = APIClient() + self.client.force_authenticate(self.user) + + def _list_ids(self, **params): + params["project_id"] = self.project.pk + query = "&".join(f"{k}={v}" for k, v in params.items()) + response = self.client.get(f"/api/v2/taxa/lists/?{query}") + self.assertEqual(response.status_code, status.HTTP_200_OK) + return response.json()["results"] + + def test_public_lists_included_by_default(self): + ids = {row["id"] for row in self._list_ids()} + self.assertEqual(ids, {self.scoped_list.pk, self.public_list.pk}) + + def test_include_public_false_hides_public_lists(self): + ids = {row["id"] for row in self._list_ids(include_public="false")} + self.assertEqual(ids, {self.scoped_list.pk}) + + def test_include_public_invalid_value_returns_400(self): + response = self.client.get(f"/api/v2/taxa/lists/?project_id={self.project.pk}&include_public=notabool") + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + + def test_public_list_linked_to_multiple_projects_appears_once(self): + """ + A public list attached to several projects must not be duplicated by the M2M + join. Runs as a superuser: visible_for_user() returns the queryset unchanged + for a superuser (no .distinct() applied there), so this exercises for_project() + on its own instead of relying on visible_for_user()'s .distinct() to mask a + duplication bug in for_project() itself. + """ + self.public_list.projects.add(self.project) + superuser = User.objects.create_superuser(email="include-public-super@example.com", password="testpass") + self.client.force_authenticate(superuser) + rows = self._list_ids() + matches = [row for row in rows if row["id"] == self.public_list.pk] + self.assertEqual(len(matches), 1) + + def test_hidden_zero_project_list_is_invisible(self): + ids = {row["id"] for row in self._list_ids()} + self.assertNotIn(self.hidden_list.pk, ids) + + def test_include_public_false_does_not_hide_a_public_list_from_retrieve(self): + """ + include_public governs the list action's default scope, not whether a + specific public row exists. ?include_public=false on a detail URL must not + 404 a public list the caller is otherwise allowed to see. + """ + response = self.client.get( + f"/api/v2/taxa/lists/{self.public_list.pk}/?project_id={self.project.pk}&include_public=false" + ) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + +class TaxaListProjectsFieldVisibilityTestCase(TestCase): + """ + A public list bypasses the draft-project visibility filter that would otherwise + hide it, so its own `projects` field must not become a side channel for + disclosing a draft project's id to someone who can't see that project. + """ + + def setUp(self): + self.owner = User.objects.create_user(email="projfield-owner@example.com", password="testpass") + self.member = User.objects.create_user(email="projfield-member@example.com", password="testpass") + self.draft_project = Project.objects.create(name="Projfield Draft Project", owner=self.owner, draft=True) + self.draft_project.members.add(self.member) + self.public_project = Project.objects.create(name="Projfield Public Project", create_defaults=False) + + self.public_list = TaxaList.objects.create(name="Cross-Project Public List", is_public=True) + self.public_list.projects.add(self.draft_project, self.public_project) + + self.client = APIClient() + + def _detail_url(self): + return f"/api/v2/taxa/lists/{self.public_list.pk}/?project_id={self.public_project.pk}" + + def test_anonymous_sees_only_the_non_draft_project_id(self): + response = self.client.get(self._detail_url()) + self.assertEqual(response.status_code, status.HTTP_200_OK) + self.assertEqual(response.json()["projects"], [self.public_project.pk]) + + def test_draft_project_member_sees_both_project_ids(self): + self.client.force_authenticate(self.member) + response = self.client.get(self._detail_url()) + self.assertEqual(response.status_code, status.HTTP_200_OK) + self.assertEqual(set(response.json()["projects"]), {self.draft_project.pk, self.public_project.pk}) + + +class TaxaListDraftProjectVisibilityTestCase(TestCase): + """A non-public list in a draft project follows the same visibility rule as any + other draft-project object: only members, owners and superusers can see it.""" + + def setUp(self): + self.owner = User.objects.create_user(email="draft-owner@example.com", password="testpass") + self.member = User.objects.create_user(email="draft-member@example.com", password="testpass") + self.outsider = User.objects.create_user(email="draft-outsider@example.com", password="testpass") + self.draft_project = Project.objects.create(name="Draft Project", owner=self.owner, draft=True) + self.draft_project.members.add(self.member) + self.scoped_list = TaxaList.objects.create(name="Draft List") + self.scoped_list.projects.add(self.draft_project) + self.client = APIClient() + + def _list_url(self): + return f"/api/v2/taxa/lists/?project_id={self.draft_project.pk}" + + def test_outsider_cannot_see_list_in_draft_project(self): + self.client.force_authenticate(self.outsider) + response = self.client.get(self._list_url()) + ids = {row["id"] for row in response.json()["results"]} + self.assertNotIn(self.scoped_list.pk, ids) + + def test_member_can_see_list_in_draft_project(self): + self.client.force_authenticate(self.member) + response = self.client.get(self._list_url()) + ids = {row["id"] for row in response.json()["results"]} + self.assertIn(self.scoped_list.pk, ids) + + +class TaxaListTaxonDraftProjectVisibilityTestCase(TestCase): + """ + A manage_public_taxalist holder bypasses the project-membership check at + has_permission() (the target might turn out to be public), but a + project-scoped list in an unrelated draft project is neither public nor + theirs — get_taxa_list() must still hide it via visible_for_user() rather + than leaking that it exists. A genuine project member can still read/write. + + A plain non-member (no platform permission) never reaches this check at + all: IsProjectMemberOrPublicListManager.has_permission() already denies + them with 403 for not being a member of the active project, regardless of + draft status — that's a separate, pre-existing gate, not this fix. + """ + + def setUp(self): + self.owner = User.objects.create_user(email="taxon-draft-owner@example.com", password="testpass") + self.member = User.objects.create_user(email="taxon-draft-member@example.com", password="testpass") + self.public_list_manager = User.objects.create_user( + email="taxon-draft-manager@example.com", password="testpass" + ) + perm = Permission.objects.get(codename="manage_public_taxalist", content_type__app_label="main") + self.public_list_manager.user_permissions.add(perm) + self.draft_project = Project.objects.create(name="Taxon Draft Project", owner=self.owner, draft=True) + self.draft_project.members.add(self.member) + self.scoped_list = TaxaList.objects.create(name="Taxon Draft List") + self.scoped_list.projects.add(self.draft_project) + self.taxon = Taxon.objects.create(name="Taxon Draft Species", rank="SPECIES") + self.client = APIClient() + + def _taxa_url(self): + return f"/api/v2/taxa/lists/{self.scoped_list.pk}/taxa/?project_id={self.draft_project.pk}" + + def test_public_list_manager_cannot_see_scoped_list_in_unrelated_draft_project(self): + self.client.force_authenticate(self.public_list_manager) + response = self.client.post(self._taxa_url(), {"taxon_id": self.taxon.pk}) + self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND) + self.assertFalse(self.scoped_list.taxa.filter(pk=self.taxon.pk).exists()) + + def test_member_can_add_taxon_to_list_in_draft_project(self): + self.client.force_authenticate(self.member) + response = self.client.post(self._taxa_url(), {"taxon_id": self.taxon.pk}) + self.assertEqual(response.status_code, status.HTTP_201_CREATED) + + +@override_settings(CACHALOT_ENABLED=False) +class TaxaListQueryCountTestCase(APITestCase): + """ + Pins the query count for TaxaListViewSet.list on a mixed public/scoped fixture, + and checks that doubling the rows adds no queries, so a per-row lookup on either + kind of list is caught. + """ + + def setUp(self): + self.user = User.objects.create_user(email="qc-user@example.com", password="testpass") + self.project = Project.objects.create(name="QC Project", owner=self.user) + for i in range(3): + scoped = TaxaList.objects.create(name=f"Scoped {i}") + scoped.projects.add(self.project) + for i in range(2): + TaxaList.objects.create(name=f"Public {i}", is_public=True) + self.client = APIClient() + self.client.force_authenticate(self.user) + + def _count_list_queries(self, expected_rows: int) -> int: + from cachalot.api import cachalot_disabled + + url = f"/api/v2/taxa/lists/?project_id={self.project.pk}" + with cachalot_disabled(), CaptureQueriesContext(connection) as ctx: + response = self.client.get(url) + self.assertEqual(response.status_code, status.HTTP_200_OK) + self.assertEqual(len(response.json()["results"]), expected_rows) + return len(ctx.captured_queries) + + def test_list_query_count(self): + # Warm process-wide caches (content types, permissions) so the pinned + # count does not depend on which tests ran earlier in the suite. + self._count_list_queries(expected_rows=5) + five_rows = self._count_list_queries(expected_rows=5) + for i in range(3, 6): + TaxaList.objects.create(name=f"Scoped {i}").projects.add(self.project) + for i in range(2, 4): + TaxaList.objects.create(name=f"Public {i}", is_public=True) + ten_rows = self._count_list_queries(expected_rows=10) + + self.assertEqual((five_rows, ten_rows), (13, 13)) + class TaxaListDedupeMigrationTestCase(TestCase): """Exercise the 0083_dedupe_taxalist_names data migration logic against seeded duplicates.""" diff --git a/ami/ml/admin.py b/ami/ml/admin.py index 008b20e84..1a8a78a89 100644 --- a/ami/ml/admin.py +++ b/ami/ml/admin.py @@ -70,8 +70,10 @@ class ProcessingServiceAdmin(AdminBase): "id", "name", "endpoint_url", + "is_public", "created_at", ] + list_filter = ["is_public"] @admin.register(AlgorithmCategoryMap) diff --git a/ami/ml/migrations/0029_processingservice_is_public.py b/ami/ml/migrations/0029_processingservice_is_public.py new file mode 100644 index 000000000..fd623cccb --- /dev/null +++ b/ami/ml/migrations/0029_processingservice_is_public.py @@ -0,0 +1,35 @@ +""" +``ProcessingService.is_public`` marks a service as available to every project, +not just the ones in its ``projects`` M2M. Schema only: unlike TaxaList, an +existing zero-project service does not become public here — it stays visible +to superusers only until someone explicitly opts it in. +""" + +from django.db import migrations, models + + +class Migration(migrations.Migration): + dependencies = [ + ("ml", "0028_normalize_empty_endpoint_url_to_null"), + ] + + operations = [ + migrations.AlterModelOptions( + name="processingservice", + options={ + "permissions": [ + ("manage_public_processingservice", "Can create, edit and delete public processing services") + ], + "verbose_name": "Processing Service", + "verbose_name_plural": "Processing Services", + }, + ), + migrations.AddField( + model_name="processingservice", + name="is_public", + field=models.BooleanField( + default=False, + help_text="Public services are shown to every project, not just the ones linked via 'projects'; running a job still requires the service to be linked to the project.", + ), + ), + ] diff --git a/ami/ml/models/pipeline.py b/ami/ml/models/pipeline.py index 3d5b77d58..359d69759 100644 --- a/ami/ml/models/pipeline.py +++ b/ami/ml/models/pipeline.py @@ -727,7 +727,7 @@ def get_or_create_taxon_for_classification( """ taxa_list, created = TaxaList.objects.get_or_create_for_project( name=f"Taxa returned by {algorithm.name}", - project=None, # Algorithm taxa lists are global + project=None, # Algorithm taxa lists have no project and stay hidden, not public ) if created: logger.info(f"Created new taxa list {taxa_list}") diff --git a/ami/ml/models/processing_service.py b/ami/ml/models/processing_service.py index fce1aefc5..b463459e8 100644 --- a/ami/ml/models/processing_service.py +++ b/ami/ml/models/processing_service.py @@ -8,7 +8,7 @@ from django.conf import settings from django.db import models -from ami.base.models import BaseQuerySet +from ami.base.models import BaseQuerySet, PublicScopedModel from ami.main.models import BaseModel, Project from ami.ml.models.pipeline import Pipeline, get_or_create_algorithm_and_category_map from ami.ml.models.project_pipeline_config import ProjectPipelineConfig @@ -59,7 +59,7 @@ def create(self, **kwargs) -> "ProcessingService": @typing.final -class ProcessingService(BaseModel): +class ProcessingService(BaseModel, PublicScopedModel): """An ML processing service""" name = models.CharField(max_length=255) @@ -70,6 +70,17 @@ class ProcessingService(BaseModel): last_seen = models.DateTimeField(null=True) last_seen_live = models.BooleanField(null=True) last_seen_latency = models.FloatField(null=True) + # Overrides PublicScopedModel.is_public with wording specific to this model: + # unlike TaxaList, job dispatch and the async heartbeat still filter processing + # services by project link regardless of is_public, so "shown" is the accurate + # claim here, not "available". + is_public = models.BooleanField( + default=False, + help_text=( + "Public services are shown to every project, not just the ones linked via 'projects'; " + "running a job still requires the service to be linked to the project." + ), + ) objects = ProcessingServiceManager() @@ -89,6 +100,7 @@ def __str__(self): class Meta: verbose_name = "Processing Service" verbose_name_plural = "Processing Services" + permissions = [("manage_public_processingservice", "Can create, edit and delete public processing services")] def create_pipelines( self, diff --git a/ami/ml/serializers.py b/ami/ml/serializers.py index e7e9e6aaf..57876a211 100644 --- a/ami/ml/serializers.py +++ b/ami/ml/serializers.py @@ -1,7 +1,9 @@ from django_pydantic_field.rest_framework import SchemaField from rest_framework import serializers +from ami.base.permissions import add_processingservice_permissions from ami.main.api.serializers import DefaultSerializer, MinimalNestedModelSerializer +from ami.main.models import Project from .models.algorithm import Algorithm, AlgorithmCategoryMap from .models.pipeline import Pipeline, PipelineStage @@ -138,6 +140,7 @@ class ProcessingServiceSerializer(DefaultSerializer): pipelines = PipelineNestedSerializer(many=True, read_only=True) projects = serializers.SerializerMethodField() is_async = serializers.BooleanField(read_only=True) + is_public = serializers.BooleanField(read_only=True) endpoint_url = serializers.CharField(required=False, allow_null=True, allow_blank=False, max_length=1024) class Meta: @@ -150,6 +153,7 @@ class Meta: "projects", "endpoint_url", "is_async", + "is_public", "pipelines", "created_at", "updated_at", @@ -159,10 +163,21 @@ class Meta: def get_projects(self, obj): """ - Return list of project IDs this processing service belongs to. - This is read-only and managed by the server. + Return the ids of this service's linked projects that are visible to the + requester. A public service can be linked to a draft project it's + otherwise not visible in; without this filter, an outsider retrieving the + public service would learn that draft project's id even though they + can't see the project itself. """ - return list(obj.projects.values_list("id", flat=True)) + request = self.context["request"] + if not hasattr(self, "_visible_project_ids"): + self._visible_project_ids = set( + Project.objects.visible_for_user(request.user).values_list("id", flat=True) + ) + return [pid for pid in obj.projects.values_list("id", flat=True) if pid in self._visible_project_ids] + + def get_permissions(self, instance, instance_data): + return add_processingservice_permissions(self.context["request"].user, instance, instance_data) class PipelineRegistrationSerializer(serializers.Serializer): diff --git a/ami/ml/tests.py b/ami/ml/tests.py index b5f99a66e..5612b1928 100644 --- a/ami/ml/tests.py +++ b/ami/ml/tests.py @@ -4,8 +4,12 @@ import unittest import uuid -from django.test import TestCase -from rest_framework.test import APIRequestFactory, APITestCase +from django.contrib.auth.models import Permission +from django.db import connection +from django.test import TestCase, override_settings +from django.test.utils import CaptureQueriesContext +from rest_framework import status +from rest_framework.test import APIClient, APIRequestFactory, APITestCase from ami.base.serializers import reverse_with_params from ami.main.models import ( @@ -82,9 +86,14 @@ def _delete_processing_service(self, processing_service_id: int): self.assertEqual(resp.status_code, 204) return resp - def _register_pipelines(self, processing_service_id): + def _register_pipelines(self, processing_service_id, with_project_id=False): + """ + Pins the frontend's call shape: usePopulateProcessingService.ts POSTs this + endpoint with no project_id. with_project_id=True exercises the other shape. + """ + params = {"project_id": self.project.pk} if with_project_id else {} processing_services_register_pipelines_url = reverse_with_params( - "api:processingservice-register-pipelines", args=[processing_service_id] + "api:processingservice-register-pipelines", args=[processing_service_id], params=params ) self.client.force_authenticate(user=self.user) resp = self.client.post(processing_services_register_pipelines_url) @@ -108,6 +117,8 @@ def test_project_was_added(self): self.assertIn(self.project, processing_service.projects.all()) def test_processing_service_pipeline_registration(self): + """Pins the frontend's call shape: usePopulateProcessingService.ts POSTs register_pipelines + with no project_id, relying on the endpoint resolving the user-visible set instead.""" # register a processing service response = self._create_processing_service( name="Processing Service Test", @@ -122,6 +133,30 @@ def test_processing_service_pipeline_registration(self): self.assertEqual(pipelines_queryset.count(), len(response["pipelines"])) + def test_processing_service_pipeline_registration_with_project_id(self): + """The other call shape: register_pipelines also works when project_id is supplied.""" + response = self._create_processing_service( + name="Processing Service Test With Project", endpoint_url="http://processing_service:2000" + ) + processing_service_id = response["id"] + + response = self._register_pipelines(processing_service_id, with_project_id=True) + processing_service = ProcessingService.objects.get(pk=processing_service_id) + + self.assertEqual(processing_service.pipelines.count(), len(response["pipelines"])) + + def test_check_status_without_project_id(self): + """Pins the frontend's call shape: useTestProcessingServiceConnection.ts GETs status + with no project_id.""" + service = ProcessingService.objects.create(name="Status Check Service", endpoint_url=None) + service.projects.add(self.project) + url = reverse_with_params("api:processingservice-status", args=[service.pk]) + + self.client.force_authenticate(user=self.user) + response = self.client.get(url) + + self.assertEqual(response.status_code, 200) + def test_create_processing_service_without_endpoint_url(self): """Test creating a ProcessingService without endpoint_url (pull mode)""" processing_services_create_url = reverse_with_params( @@ -227,6 +262,329 @@ def test_model_has_last_seen_fields(self): self.assertFalse(hasattr(service, "last_checked_latency")) +class ProcessingServicePublicPermissionsTestCase(TestCase): + """ + Permission matrix for public vs. project-scoped ProcessingServices. + + A project-scoped service keeps its existing staff-only write rule (any + active staff member, project membership not required). A public service + instead requires the manage_public_processingservice platform permission + (or a superuser) — plain staff status is not enough. All services use + endpoint_url=None (pull-mode) so get_status()/create_pipelines() never + make a real network call. + """ + + def setUp(self): + self.staff = User.objects.create_user(email="staff-ps@example.com", password="testpass", is_staff=True) + self.member = User.objects.create_user(email="member-ps@example.com", password="testpass") + self.non_member = User.objects.create_user(email="nonmember-ps@example.com", password="testpass") + self.superuser = User.objects.create_superuser(email="super-ps@example.com", password="testpass") + self.public_manager = User.objects.create_user( + email="manager-ps@example.com", password="testpass", is_staff=True + ) + perm = Permission.objects.get(codename="manage_public_processingservice", content_type__app_label="ml") + self.public_manager.user_permissions.add(perm) + + self.project = Project.objects.create(name="PS Test Project", create_defaults=False) + self.project.members.add(self.member) + + self.public_service = ProcessingService.objects.create( + name="Public Service", endpoint_url=None, is_public=True + ) + self.scoped_service = ProcessingService.objects.create(name="Scoped Service", endpoint_url=None) + self.scoped_service.projects.add(self.project) + + self.client = APIClient() + + def _detail_url(self, service): + return f"/api/v2/ml/processing_services/{service.pk}/?project_id={self.project.pk}" + + def _status_url(self, service): + return f"/api/v2/ml/processing_services/{service.pk}/status/?project_id={self.project.pk}" + + def _register_url(self, service): + return f"/api/v2/ml/processing_services/{service.pk}/register_pipelines/?project_id={self.project.pk}" + + def _register_url_no_project(self, service): + """Pins the frontend's call shape: usePopulateProcessingService.ts POSTs with no project_id.""" + return f"/api/v2/ml/processing_services/{service.pk}/register_pipelines/" + + # -- Update -- + + def test_staff_can_update_scoped_service(self): + self.client.force_authenticate(self.staff) + response = self.client.patch(self._detail_url(self.scoped_service), {"name": "Renamed"}) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + def test_staff_cannot_update_public_service(self): + """Plain staff status is not enough for a public service.""" + self.client.force_authenticate(self.staff) + response = self.client.patch(self._detail_url(self.public_service), {"name": "Hacked"}) + self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) + + def test_member_cannot_update_scoped_service(self): + """Project membership alone is not the write gate here — staff status is.""" + self.client.force_authenticate(self.member) + response = self.client.patch(self._detail_url(self.scoped_service), {"name": "Hacked"}) + self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) + + def test_superuser_can_update_public_service(self): + self.client.force_authenticate(self.superuser) + response = self.client.patch(self._detail_url(self.public_service), {"name": "Renamed by super"}) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + def test_public_manager_can_update_public_service(self): + self.client.force_authenticate(self.public_manager) + response = self.client.patch(self._detail_url(self.public_service), {"name": "Renamed"}) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + def test_public_manager_can_still_update_scoped_service(self): + """The public manager is also staff, so the existing staff-only rule still applies.""" + self.client.force_authenticate(self.public_manager) + response = self.client.patch(self._detail_url(self.scoped_service), {"name": "Renamed"}) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + # -- Delete -- + + def test_staff_cannot_delete_public_service(self): + self.client.force_authenticate(self.staff) + response = self.client.delete(self._detail_url(self.public_service)) + self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) + + def test_staff_can_delete_scoped_service(self): + self.client.force_authenticate(self.staff) + response = self.client.delete(self._detail_url(self.scoped_service)) + self.assertEqual(response.status_code, status.HTTP_204_NO_CONTENT) + + def test_public_manager_can_delete_public_service(self): + self.client.force_authenticate(self.public_manager) + response = self.client.delete(self._detail_url(self.public_service)) + self.assertEqual(response.status_code, status.HTTP_204_NO_CONTENT) + + def test_public_manager_can_delete_public_service_with_include_public_false(self): + """ + include_public=false must not make get_queryset() 404 the very row being + deleted: it governs the list action's default scope, not whether a public + row can be looked up for a detail action. + """ + self.client.force_authenticate(self.public_manager) + url = f"{self._detail_url(self.public_service)}&include_public=false" + response = self.client.delete(url) + self.assertEqual(response.status_code, status.HTTP_204_NO_CONTENT) + + # -- Retrieve / list visibility -- + + def test_anonymous_can_retrieve_public_service(self): + response = self.client.get(self._detail_url(self.public_service)) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + def test_non_member_can_see_public_service(self): + self.client.force_authenticate(self.non_member) + response = self.client.get(f"/api/v2/ml/processing_services/?project_id={self.project.pk}") + self.assertEqual(response.status_code, status.HTTP_200_OK) + ids = {row["id"] for row in response.json()["results"]} + self.assertIn(self.public_service.pk, ids) + + def test_is_public_is_reported_in_the_response(self): + self.client.force_authenticate(self.staff) + response = self.client.get(self._detail_url(self.public_service)) + self.assertTrue(response.json()["is_public"]) + + def test_is_public_cannot_be_set_through_the_api(self): + self.client.force_authenticate(self.staff) + self.client.patch(self._detail_url(self.scoped_service), {"is_public": True}) + self.scoped_service.refresh_from_db() + self.assertFalse(self.scoped_service.is_public) + + def test_user_permissions_include_update_delete_for_public_manager_only(self): + self.client.force_authenticate(self.staff) + response = self.client.get(self._detail_url(self.public_service)) + self.assertNotIn("update", response.json()["user_permissions"]) + + self.client.force_authenticate(self.public_manager) + response = self.client.get(self._detail_url(self.public_service)) + perms = response.json()["user_permissions"] + self.assertIn("update", perms) + self.assertIn("delete", perms) + + # -- status (a read-type action; open to everyone like any other safe method) -- + + def test_anonymous_can_check_status_of_public_service(self): + response = self.client.get(self._status_url(self.public_service)) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + def test_non_member_can_check_status_of_scoped_service(self): + self.client.force_authenticate(self.non_member) + response = self.client.get(self._status_url(self.scoped_service)) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + # -- register_pipelines -- + + def test_staff_cannot_register_pipelines_on_public_service(self): + self.client.force_authenticate(self.staff) + response = self.client.post(self._register_url(self.public_service)) + self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) + + def test_staff_can_register_pipelines_on_scoped_service(self): + self.client.force_authenticate(self.staff) + response = self.client.post(self._register_url(self.scoped_service)) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + def test_public_manager_can_register_pipelines_on_public_service(self): + self.client.force_authenticate(self.public_manager) + response = self.client.post(self._register_url(self.public_service)) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + def test_member_cannot_register_pipelines_on_scoped_service(self): + self.client.force_authenticate(self.member) + response = self.client.post(self._register_url(self.scoped_service)) + self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) + + def test_public_manager_can_register_pipelines_on_public_service_without_project_id(self): + """The platform-permission bypass works the same whether or not project_id is supplied.""" + self.client.force_authenticate(self.public_manager) + response = self.client.post(self._register_url_no_project(self.public_service)) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + def test_include_public_false_does_not_hide_a_public_service_from_retrieve(self): + """ + include_public governs the list action's default scope, not whether a + specific public row exists. ?include_public=false on a detail URL must not + 404 a public service the caller is otherwise allowed to see. + """ + url = f"{self._detail_url(self.public_service)}&include_public=false" + response = self.client.get(url) + self.assertEqual(response.status_code, status.HTTP_200_OK) + + +class ProcessingServiceIncludePublicParamTestCase(TestCase): + """?include_public toggles whether public services appear alongside a project's own.""" + + def setUp(self): + self.user = User.objects.create_user(email="ps-scope-user@example.com", password="testpass") + self.project = Project.objects.create(name="PS Scope Project", create_defaults=False) + self.project.members.add(self.user) + self.other_project = Project.objects.create(name="PS Other Scope Project", create_defaults=False) + + self.scoped_service = ProcessingService.objects.create(name="PS Scoped", endpoint_url=None) + self.scoped_service.projects.add(self.project) + self.public_service = ProcessingService.objects.create(name="PS Public", endpoint_url=None, is_public=True) + # A public service can also be linked to an unrelated project without appearing twice. + self.public_service.projects.add(self.other_project) + + self.hidden_service = ProcessingService.objects.create(name="PS Hidden with no project", endpoint_url=None) + + self.client = APIClient() + self.client.force_authenticate(self.user) + + def _list_ids(self, **params): + params["project_id"] = self.project.pk + query = "&".join(f"{k}={v}" for k, v in params.items()) + response = self.client.get(f"/api/v2/ml/processing_services/?{query}") + self.assertEqual(response.status_code, status.HTTP_200_OK) + return response.json()["results"] + + def test_public_services_included_by_default(self): + ids = {row["id"] for row in self._list_ids()} + self.assertEqual(ids, {self.scoped_service.pk, self.public_service.pk}) + + def test_include_public_false_hides_public_services(self): + ids = {row["id"] for row in self._list_ids(include_public="false")} + self.assertEqual(ids, {self.scoped_service.pk}) + + def test_include_public_invalid_value_returns_400(self): + response = self.client.get( + f"/api/v2/ml/processing_services/?project_id={self.project.pk}&include_public=notabool" + ) + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + + # The no-duplication guarantee for a public row linked to several projects is a + # property of the shared for_project()/visible_for_user() code, tested once at + # the querySet level (TaxaListForProjectQuerySetTestCase) and once at the API + # level, as a superuser, in TaxaListIncludePublicParamTestCase — no need to + # repeat it here for ProcessingService. + + def test_hidden_zero_project_service_is_invisible_to_non_superuser(self): + ids = {row["id"] for row in self._list_ids()} + self.assertNotIn(self.hidden_service.pk, ids) + + +class ProcessingServiceProjectsFieldVisibilityTestCase(TestCase): + """ + A public service bypasses the draft-project visibility filter that would + otherwise hide it, so its own `projects` field must not become a side channel + for disclosing a draft project's id to someone who can't see that project. + """ + + def setUp(self): + self.owner = User.objects.create_user(email="ps-projfield-owner@example.com", password="testpass") + self.member = User.objects.create_user(email="ps-projfield-member@example.com", password="testpass") + self.draft_project = Project.objects.create( + name="PS Projfield Draft Project", owner=self.owner, draft=True, create_defaults=False + ) + self.draft_project.members.add(self.member) + self.public_project = Project.objects.create(name="PS Projfield Public Project", create_defaults=False) + + self.public_service = ProcessingService.objects.create( + name="PS Cross-Project Public Service", endpoint_url=None, is_public=True + ) + self.public_service.projects.add(self.draft_project, self.public_project) + + self.client = APIClient() + + def _detail_url(self): + return f"/api/v2/ml/processing_services/{self.public_service.pk}/?project_id={self.public_project.pk}" + + def test_anonymous_sees_only_the_non_draft_project_id(self): + response = self.client.get(self._detail_url()) + self.assertEqual(response.status_code, status.HTTP_200_OK) + self.assertEqual(response.json()["projects"], [self.public_project.pk]) + + def test_draft_project_member_sees_both_project_ids(self): + self.client.force_authenticate(self.member) + response = self.client.get(self._detail_url()) + self.assertEqual(response.status_code, status.HTTP_200_OK) + self.assertEqual(set(response.json()["projects"]), {self.draft_project.pk, self.public_project.pk}) + + +@override_settings(CACHALOT_ENABLED=False) +class ProcessingServiceQueryCountTestCase(APITestCase): + """ + Pins the current query count for ProcessingServiceViewSet.list on a mixed + public/scoped, multi-row fixture, so a regression that adds queries is + noticed. This does not certify the absence of per-row queries — it only + catches a further increase from where things stand today. + """ + + def setUp(self): + self.user = User.objects.create_user(email="ps-qc-user@example.com", password="testpass") + self.project = Project.objects.create(name="PS QC Project", create_defaults=False) + self.project.members.add(self.user) + for i in range(3): + scoped = ProcessingService.objects.create(name=f"PS Scoped {i}", endpoint_url=None) + scoped.projects.add(self.project) + for i in range(2): + ProcessingService.objects.create(name=f"PS Public {i}", endpoint_url=None, is_public=True) + self.client = APIClient() + self.client.force_authenticate(self.user) + + def test_list_query_count(self): + from cachalot.api import cachalot_disabled + + url = f"/api/v2/ml/processing_services/?project_id={self.project.pk}" + with cachalot_disabled(), CaptureQueriesContext(connection) as ctx: + response = self.client.get(url) + self.assertEqual(response.status_code, status.HTTP_200_OK) + self.assertEqual(len(response.json()["results"]), 5) + # 34 (previous baseline) down to 21: add_processingservice_permissions() no + # longer runs the M2M membership check or the guardian get_perms() lookup + # per non-public row (no per-project *_processingservice guardian permission + # exists, so that branch could only ever fire for a superuser, which a plain + # attribute check covers for free); get_projects() adds back one query per + # request (not per row) for the draft-project-id visibility filter. + self.assertEqual(len(ctx.captured_queries), 21) + + class TestProjectPipelineRegistrationUpdatesLastSeen(APITestCase): """Test that async pipeline registration updates last_seen on the processing service.""" diff --git a/ami/ml/views.py b/ami/ml/views.py index 63e460af6..cb3997a58 100644 --- a/ami/ml/views.py +++ b/ami/ml/views.py @@ -11,9 +11,9 @@ from rest_framework.request import Request from rest_framework.response import Response -from ami.base.permissions import ProjectPipelineConfigPermission -from ami.base.views import ProjectMixin -from ami.main.api.schemas import project_id_doc_param +from ami.base.permissions import IsActiveStaffOrPublicManagerOrReadOnly, ProjectPipelineConfigPermission +from ami.base.views import ProjectMixin, get_active_project +from ami.main.api.schemas import include_public_doc_param, project_id_doc_param from ami.main.api.views import DefaultViewSet from ami.main.models import Project, SourceImage from ami.ml.schemas import PipelineRegistrationResponse @@ -168,16 +168,30 @@ class ProcessingServiceViewSet(DefaultViewSet, ProjectMixin): serializer_class = ProcessingServiceSerializer filterset_fields = ["projects"] ordering_fields = ["id", "created_at", "updated_at"] + permission_classes = [IsActiveStaffOrPublicManagerOrReadOnly] require_project = True + def get_active_project(self) -> Project | None: + """ + project_id is optional for status/register_pipelines — the frontend calls both + without one — but stays required for every other action, as declared above. + """ + if self.action in ("status", "register_pipelines"): + return get_active_project(request=self.request, kwargs=self.kwargs, required=False) + return super().get_active_project() + def get_queryset(self) -> QuerySet: qs: QuerySet = super().get_queryset() project = self.get_active_project() - if project: - qs = qs.filter(projects=project) - return qs - - @extend_schema(parameters=[project_id_doc_param]) + if not project: + return qs + # include_public governs the list action's default scope, not whether a + # specific row is reachable: a detail/update/delete/status/register_pipelines + # on a public service must still resolve it even under ?include_public=false. + include_public = self.get_include_public() if self.action == "list" else True + return qs.for_project(project, include_public=include_public) + + @extend_schema(parameters=[project_id_doc_param, include_public_doc_param]) def list(self, request, *args, **kwargs): return super().list(request, *args, **kwargs) @@ -214,13 +228,13 @@ def status(self, request: Request, pk=None) -> Response: """ Test the connection to the processing service. """ - processing_service = ProcessingService.objects.get(pk=pk) + processing_service = self.get_object() response = processing_service.get_status() return Response(response.dict()) @action(detail=True, methods=["post"]) def register_pipelines(self, request: Request, pk=None) -> Response: - processing_service = ProcessingService.objects.get(pk=pk) + processing_service = self.get_object() response = processing_service.create_pipelines() processing_service.save() return Response(response.dict())