diff --git a/ami/base/permissions.py b/ami/base/permissions.py index 247d5e5d4..328aa8384 100644 --- a/ami/base/permissions.py +++ b/ami/base/permissions.py @@ -77,7 +77,9 @@ def add_collection_level_permissions(user: User | None, response_data: dict, mod return response_data -def add_m2m_object_permissions(user, instance, project, response_data: dict) -> dict: +def add_m2m_object_permissions( + user, instance, project, response_data: dict, project_perms: set[str] | None = None +) -> dict: """ Add object-level permissions for models with an M2M relationship to Project. @@ -87,7 +89,12 @@ def add_m2m_object_permissions(user, instance, project, response_data: dict) -> against a specific project from the request context instead. Validates that the instance actually belongs to the given project before - granting any permissions (prevents cross-project permission leaks). + granting any permissions (prevents cross-project permission leaks); uses + `instance.projects`'s prefetch cache when the caller populated one. + + `project_perms` lets a caller resolving many instances for the same + (user, project) pass in `guardian.get_perms(user, project)` once instead + of once per instance; pass None to look it up here as before. 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 @@ -96,7 +103,16 @@ class (Pattern B: Bare M2M) that handles TaxaList, Taxon, ProcessingService, """ perms = set(response_data.get("user_permissions", [])) - if not project or not instance.projects.filter(pk=project.pk).exists(): + if not project: + response_data["user_permissions"] = list(perms) + return response_data + + if "projects" in getattr(instance, "_prefetched_objects_cache", {}): + is_member = any(p.pk == project.pk for p in instance.projects.all()) + else: + is_member = instance.projects.filter(pk=project.pk).exists() + + if not is_member: response_data["user_permissions"] = list(perms) return response_data @@ -104,7 +120,7 @@ class (Pattern B: Bare M2M) that handles TaxaList, Taxon, ProcessingService, perms.update(["update", "delete"]) else: model_name = instance._meta.model_name - all_perms = get_perms(user, project) + all_perms = project_perms if project_perms is not None else get_perms(user, project) for perm in all_perms: if perm.endswith(f"_{model_name}"): action = perm.split("_", 1)[0] diff --git a/ami/main/api/serializers.py b/ami/main/api/serializers.py index 952f936f0..962c802af 100644 --- a/ami/main/api/serializers.py +++ b/ami/main/api/serializers.py @@ -736,20 +736,38 @@ def get_taxa_count(self, obj): """ Return the number of taxa in this list. Uses annotated_taxa_count if available (from ViewSet) for performance. + + `getattr(obj, name, obj.taxa.count())` would evaluate `obj.taxa.count()` + as a default argument on every call regardless of whether the attribute + is present, running a COUNT query per row even when annotated. """ - return getattr(obj, "annotated_taxa_count", obj.taxa.count()) + annotated_count = getattr(obj, "annotated_taxa_count", None) + return annotated_count if annotated_count is not None else obj.taxa.count() def get_permissions(self, instance, instance_data): + # DRF's ListSerializer reuses one child instance across every row, and a + # fresh serializer is built per request, so caching on `self` resolves + # the project and the member's permissions once per request, not per row. request = self.context["request"] - project = get_active_project(request=request) - return add_m2m_object_permissions(request.user, instance, project, instance_data) + if not hasattr(self, "_active_project"): + self._active_project = get_active_project(request=request) + project = self._active_project + + project_perms = None + if project and not request.user.is_superuser: + if not hasattr(self, "_project_perms"): + self._project_perms = set(get_perms(request.user, project)) + project_perms = self._project_perms + + return add_m2m_object_permissions(request.user, instance, project, instance_data, project_perms=project_perms) def get_projects(self, obj): """ - Return list of project IDs this taxa list belongs to. - This is read-only and managed by the server. + Return list of project IDs this taxa list belongs to, sorted for a + deterministic response. Reads the `projects` prefetched by + TaxaListViewSet.get_queryset instead of querying per row. """ - return list(obj.projects.values_list("id", flat=True)) + return sorted(project.pk for project in obj.projects.all()) class TaxaListTaxonInputSerializer(serializers.Serializer): diff --git a/ami/main/api/views.py b/ami/main/api/views.py index 5fe6cf863..fdc72c963 100644 --- a/ami/main/api/views.py +++ b/ami/main/api/views.py @@ -2216,6 +2216,10 @@ def get_queryset(self): qs = super().get_queryset() # Annotate with taxa count for better performance qs = qs.annotate(annotated_taxa_count=models.Count("taxa")) + # Prefetched once here so the serializer's get_projects() and the + # membership check in add_m2m_object_permissions() don't hit the + # database once per row. + qs = qs.prefetch_related("projects") project = self.get_active_project() if project: return qs.filter(projects=project) diff --git a/ami/main/tests.py b/ami/main/tests.py index 3c9a1be17..aab6d943d 100644 --- a/ami/main/tests.py +++ b/ami/main/tests.py @@ -18,6 +18,7 @@ from rest_framework.test import APIClient, APIRequestFactory, APITestCase from rich import print +from ami.base.permissions import add_m2m_object_permissions from ami.exports.models import DataExport from ami.jobs.models import VALID_JOB_TYPES, Job from ami.main.api.serializers import MAX_BULK_IDENTIFICATIONS @@ -4418,6 +4419,65 @@ def test_example_subqueries_stripped_from_pagination_count(self): self.assertNotIn("main_identification", sql.lower(), "example subqueries leaked into the COUNT") +@override_settings(CACHALOT_ENABLED=False) +class TestTaxaListListQueryCount(APITestCase): + """Guard against N+1 regressions in TaxaListViewSet.list. + + TaxaListSerializer resolved the active project and the requesting member's + project permissions once per row instead of once per request + (get_projects, get_permissions -> add_m2m_object_permissions). Query count + must stay flat as the number of taxa lists returned grows. Uses a project + member (not a superuser) because the superuser branch of + add_m2m_object_permissions skips the guardian lookup entirely. + """ + + def setUp(self): + self.owner = User.objects.create_user(email="taxalist-qcount-owner@insectai.org") + self.member = User.objects.create_user(email="taxalist-qcount-member@insectai.org") + self.taxon = Taxon.objects.create(name="Query Count Taxon", rank=TaxonRank.SPECIES.name) + self.client.force_authenticate(self.member) + + def _make_project_with_lists(self, name: str, count: int) -> Project: + project = Project.objects.create(name=name, owner=self.owner) + project.members.add(self.member) + for i in range(count): + taxa_list = TaxaList.objects.create(name=f"{name} List {i}") + taxa_list.projects.add(project) + taxa_list.taxa.add(self.taxon) + return project + + def _list_query_count(self, project: Project, expected_rows: int) -> int: + from django.core.cache import caches + + url = f"/api/v2/taxa/lists/?project_id={project.pk}&limit=25" + caches["default"].clear() + with CaptureQueriesContext(connection) as ctx: + res = self.client.get(url) + self.assertEqual(res.status_code, status.HTTP_200_OK) + self.assertEqual(len(res.data["results"]), expected_rows) + return len(ctx.captured_queries) + + def test_list_query_count_does_not_scale_with_row_count(self): + # Each measurement targets its own project, so cachalot's per-query + # cache (keyed on SQL text + params, not row counts) can't serve one + # measurement's result to the other and mask a real regression: every + # project_id is queried exactly once across the whole test. + small_project = self._make_project_with_lists("Small", 3) + large_project = self._make_project_with_lists("Large", 10) + + # Warm up process-global caches (ContentType, guardian's content-type + # lookups) on a throwaway project first, so neither measurement below + # pays a one-time setup cost that the other doesn't. + warmup_project = self._make_project_with_lists("Warmup", 1) + self._list_query_count(warmup_project, expected_rows=1) + + small = self._list_query_count(small_project, expected_rows=3) + large = self._list_query_count(large_project, expected_rows=10) + + print(f"\n[AUDIT] TaxaList list: 3 rows -> {small}q, 10 rows -> {large}q") + self.assertEqual(small, large, f"Query count scaled with row count: {small} -> {large} (N+1 regression)") + + class TestProjectDefaultTaxaFilter(APITestCase): """ Tests for project default taxa filtering (include/exclude lists). @@ -5229,6 +5289,75 @@ def test_non_member_cannot_update_taxa_list(self): self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) +class TaxaListPermissionScopingTestCase(TestCase): + """Guard the membership check in add_m2m_object_permissions against the + prefetch-based fast path added to fix per-row queries (see #1120). + + The member holds the ProjectManager role, the only role granting + update_taxalist/delete_taxalist (ami/users/roles.py) — a plain BasicMember + never has those permissions on any project, so the negative-only guards + below would pass even if the scoping check were removed entirely. + """ + + def setUp(self): + self.owner = User.objects.create_user(email="scoping-owner@example.com", password="testpass") + self.member = User.objects.create_user(email="scoping-member@example.com", password="testpass") + self.project = Project.objects.create(name="Scoping Project", owner=self.owner) + self.project.members.add(self.member) + ProjectManager.assign_user(self.member, self.project) + self.client = APIClient() + + def test_prefetched_non_member_list_gets_no_write_permissions(self): + """A taxa list outside the active project must still report no update/delete + permissions when `instance.projects` was prefetched by the caller, not queried, + even though the member holds update/delete_taxalist on the active project.""" + other_project = Project.objects.create(name="Other Project", owner=self.owner) + outside_list = TaxaList.objects.create(name="Outside List") + outside_list.projects.add(other_project) + + instance = TaxaList.objects.filter(pk=outside_list.pk).prefetch_related("projects").get() + self.assertIn("projects", instance._prefetched_objects_cache) + + data = add_m2m_object_permissions(self.member, instance, self.project, {}) + self.assertNotIn("update", data["user_permissions"]) + self.assertNotIn("delete", data["user_permissions"]) + + def test_prefetched_member_list_gets_write_permissions(self): + """Positive counterpart to the guard above: a list that does belong to + the active project reports update/delete for a ProjectManager member, + so the negative case isn't just a permission set that's always empty.""" + member_list = TaxaList.objects.create(name="Member List") + member_list.projects.add(self.project) + + instance = TaxaList.objects.filter(pk=member_list.pk).prefetch_related("projects").get() + data = add_m2m_object_permissions(self.member, instance, self.project, {}) + self.assertIn("update", data["user_permissions"]) + self.assertIn("delete", data["user_permissions"]) + + def test_permissions_scoped_to_requested_project_not_other_memberships(self): + """A member with write permissions on project A must not see those + permissions on a taxa list shared with project B when the request is + scoped to project B, where they hold no role, even though the same + list belongs to both.""" + project_b = Project.objects.create(name="Project B", owner=self.owner) + shared_list = TaxaList.objects.create(name="Shared List") + shared_list.projects.add(self.project, project_b) + + self.client.force_authenticate(self.member) + + response_a = self.client.get(f"/api/v2/taxa/lists/?project_id={self.project.pk}") + self.assertEqual(response_a.status_code, status.HTTP_200_OK) + row_a = next(r for r in response_a.data["results"] if r["id"] == shared_list.pk) + self.assertIn("update", row_a["user_permissions"]) + self.assertIn("delete", row_a["user_permissions"]) + + response_b = self.client.get(f"/api/v2/taxa/lists/?project_id={project_b.pk}") + self.assertEqual(response_b.status_code, status.HTTP_200_OK) + row_b = next(r for r in response_b.data["results"] if r["id"] == shared_list.pk) + self.assertNotIn("update", row_b["user_permissions"]) + self.assertNotIn("delete", row_b["user_permissions"]) + + class TaxaListTaxonAPITestCase(TestCase): """Test TaxaList taxa management operations via API."""