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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 20 additions & 4 deletions ami/base/permissions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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
Expand All @@ -96,15 +103,24 @@ 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

if user.is_superuser:
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]
Expand Down
30 changes: 24 additions & 6 deletions ami/main/api/serializers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Comment thread
mihow marked this conversation as resolved.

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):
Expand Down
4 changes: 4 additions & 0 deletions ami/main/api/views.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
129 changes: 129 additions & 0 deletions ami/main/tests.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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).
Expand Down Expand Up @@ -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."""

Expand Down
Loading