diff --git a/src/documents/serialisers.py b/src/documents/serialisers.py index aaded56fe..2d6ff43a3 100644 --- a/src/documents/serialisers.py +++ b/src/documents/serialisers.py @@ -39,7 +39,6 @@ from drf_spectacular.utils import extend_schema_field from drf_spectacular.utils import extend_schema_serializer from drf_writable_nested.serializers import NestedUpdateMixin from guardian.core import ObjectPermissionChecker -from guardian.shortcuts import get_objects_for_user from guardian.shortcuts import get_users_with_perms from guardian.utils import get_group_obj_perms_model from guardian.utils import get_user_obj_perms_model @@ -80,8 +79,8 @@ from documents.models import WorkflowTrigger from documents.parsers import is_mime_type_supported from documents.permissions import get_document_count_filter_for_user from documents.permissions import get_groups_with_only_permission -from documents.permissions import get_objects_for_user_owner_aware from documents.permissions import has_perms_owner_aware +from documents.permissions import permitted_document_ids from documents.permissions import set_permissions_for_object from documents.regex import validate_regex_pattern from documents.templating.filepath import validate_filepath_template_and_render @@ -1011,13 +1010,8 @@ def _get_viewable_duplicates( ).exclude(pk=document.pk) duplicates = duplicates.filter(root_document__isnull=True) duplicates = duplicates.order_by("-created") - allowed = get_objects_for_user_owner_aware( - user, - "documents.view_document", - Document, - include_deleted=True, - ) - return duplicates.filter(id__in=allowed) + allowed_ids = permitted_document_ids(user, include_deleted=True) + return duplicates.filter(id__in=allowed_ids) class DuplicateDocumentSummarySerializer(serializers.Serializer[dict[str, Any]]): @@ -2672,13 +2666,8 @@ class TaskSerializerV9(serializers.ModelSerializer[PaperlessTask]): user = request.user qs = Document.global_objects.filter(pk=dup_of) if not user.is_staff: - with_perms = get_objects_for_user( - user, - "documents.view_document", - qs, - accept_global_perms=False, - ) - qs = with_perms | qs.filter(owner=user) | qs.filter(owner__isnull=True) + allowed_ids = permitted_document_ids(user, include_deleted=True) + qs = qs.filter(pk__in=allowed_ids) return list(qs.values("id", "title", "deleted_at")) @@ -3528,8 +3517,6 @@ class StoragePathTestSerializer(SerializerWithPerms): document_field = self.fields.get("document") if not isinstance(document_field, serializers.PrimaryKeyRelatedField): return - document_field.queryset = get_objects_for_user_owner_aware( - user, - "documents.view_document", - Document, + document_field.queryset = Document.objects.filter( + id__in=permitted_document_ids(user), ) diff --git a/src/documents/tests/test_permission_filtering_security.py b/src/documents/tests/test_permission_filtering_security.py index 0db6a692a..ba5522cb4 100644 --- a/src/documents/tests/test_permission_filtering_security.py +++ b/src/documents/tests/test_permission_filtering_security.py @@ -1,12 +1,18 @@ from __future__ import annotations +from unittest.mock import patch + import pytest from django.contrib.auth.models import AnonymousUser from django.contrib.auth.models import Group +from django.contrib.auth.models import Permission from django.contrib.auth.models import User +from django.test import override_settings from guardian.shortcuts import assign_perm +from rest_framework.test import APIClient from documents.permissions import permitted_document_ids +from documents.serialisers import _get_viewable_duplicates from documents.tests.factories import DocumentFactory @@ -151,3 +157,63 @@ class TestPermittedDocumentIdsIncludeDeleted: expected_visible=[], expected_hidden=[doc.pk], ) + + +@pytest.mark.django_db +class TestAiChatAllDocumentsPermissionBoundary: + """ + Regression test pinning the "ask across all documents" AI chat behavior + (ChatStreamingView.post, no document_id) to the same owner/permission + boundary enforced by permitted_document_ids(). This call site was + migrated from get_objects_for_user_owner_aware() to + permitted_document_ids(); this test must stay green across that swap. + """ + + ENDPOINT = "/api/documents/chat/" + + @override_settings(AI_ENABLED=True) + @patch("documents.views.stream_chat_with_documents") + def test_chat_all_documents_excludes_unshared_document(self, mock_stream_chat): + mock_stream_chat.return_value = iter([b"data"]) + + owner = User.objects.create_user(username="owner") + asker = User.objects.create_user(username="asker") + asker.user_permissions.add( + *Permission.objects.filter(codename="view_document"), + ) + shared = DocumentFactory(owner=owner) + not_shared = DocumentFactory(owner=owner) + assign_perm("view_document", asker, shared) + + client = APIClient() + client.force_authenticate(user=asker) + response = client.post( + self.ENDPOINT, + data={"q": "question"}, + format="json", + ) + + assert response.status_code == 200 + mock_stream_chat.assert_called_once() + _, kwargs = mock_stream_chat.call_args + visible_ids = {doc.pk for doc in kwargs["documents"]} + assert shared.pk in visible_ids + assert not_shared.pk not in visible_ids + + +@pytest.mark.django_db +class TestDuplicateDocumentsPermissionBoundary: + def test_get_viewable_duplicates_includes_soft_deleted_but_respects_perms(self): + owner = User.objects.create_user(username="owner") + stranger = User.objects.create_user(username="mallory") + original = DocumentFactory(owner=owner, checksum="dupe-checksum") + dup_visible = DocumentFactory(owner=owner, checksum="dupe-checksum") + dup_hidden = DocumentFactory(owner=owner, checksum="dupe-checksum") + dup_hidden.delete() # soft delete, should still be found (include_deleted=True) + assign_perm("view_document", stranger, dup_visible) + + result_owner = _get_viewable_duplicates(original, owner) + assert {d.pk for d in result_owner} == {dup_visible.pk, dup_hidden.pk} + + result_stranger = _get_viewable_duplicates(original, stranger) + assert {d.pk for d in result_stranger} == {dup_visible.pk} diff --git a/src/documents/tests/test_views.py b/src/documents/tests/test_views.py index 0cf1bcdae..c385e9b57 100644 --- a/src/documents/tests/test_views.py +++ b/src/documents/tests/test_views.py @@ -648,11 +648,11 @@ class TestAIChatStreamingView(DirectoriesMixin, TestCase): self.assertIn(b"AI is required for this feature", response.content) @patch("documents.views.stream_chat_with_documents") - @patch("documents.views.get_objects_for_user_owner_aware") + @patch("documents.views.permitted_document_ids") @override_settings(AI_ENABLED=True) - def test_post_no_document_id(self, mock_get_objects, mock_stream_chat) -> None: + def test_post_no_document_id(self, mock_permitted_ids, mock_stream_chat) -> None: self.grant_view_document_permission() - mock_get_objects.return_value = [self.document] + mock_permitted_ids.return_value = [self.document.pk] mock_stream_chat.return_value = iter([b"data"]) response = self.client.post( self.ENDPOINT, @@ -661,23 +661,23 @@ class TestAIChatStreamingView(DirectoriesMixin, TestCase): ) self.assertEqual(response.status_code, 200) self.assertEqual(response["Content-Type"], "text/event-stream") - mock_stream_chat.assert_called_once_with( - query_str="question", - documents=[self.document], - output_language=None, - ) + mock_stream_chat.assert_called_once() + call_kwargs = mock_stream_chat.call_args.kwargs + self.assertEqual(call_kwargs["query_str"], "question") + self.assertEqual(list(call_kwargs["documents"]), [self.document]) + self.assertIsNone(call_kwargs["output_language"]) @patch("documents.views.stream_chat_with_documents") - @patch("documents.views.get_objects_for_user_owner_aware") + @patch("documents.views.permitted_document_ids") @override_settings(AI_ENABLED=True) def test_post_uses_user_display_language( self, - mock_get_objects, + mock_permitted_ids, mock_stream_chat, ) -> None: UiSettings.objects.create(user=self.user, settings={"language": "de-de"}) self.grant_view_document_permission() - mock_get_objects.return_value = [self.document] + mock_permitted_ids.return_value = [self.document.pk] mock_stream_chat.return_value = iter([b"data"]) response = self.client.post( @@ -687,11 +687,11 @@ class TestAIChatStreamingView(DirectoriesMixin, TestCase): ) self.assertEqual(response.status_code, 200) - mock_stream_chat.assert_called_once_with( - query_str="question", - documents=[self.document], - output_language="de-de", - ) + mock_stream_chat.assert_called_once() + call_kwargs = mock_stream_chat.call_args.kwargs + self.assertEqual(call_kwargs["query_str"], "question") + self.assertEqual(list(call_kwargs["documents"]), [self.document]) + self.assertEqual(call_kwargs["output_language"], "de-de") @patch("documents.views.stream_chat_with_documents") @override_settings(AI_ENABLED=True) diff --git a/src/documents/views.py b/src/documents/views.py index 405913f95..65a92c0e9 100644 --- a/src/documents/views.py +++ b/src/documents/views.py @@ -177,6 +177,7 @@ from documents.permissions import get_objects_for_user_owner_aware from documents.permissions import has_global_statistics_permission from documents.permissions import has_perms_owner_aware from documents.permissions import has_system_status_permission +from documents.permissions import permitted_document_ids from documents.permissions import set_permissions_for_object from documents.plugins.date_parsing import get_date_parser from documents.schema import generate_object_with_permissions_schema @@ -2270,10 +2271,8 @@ class ChatStreamingView(GenericAPIView[Any]): documents = [document] else: - documents = get_objects_for_user_owner_aware( - request.user, - "view_document", - Document, + documents = Document.objects.filter( + id__in=permitted_document_ids(request.user), ) output_language = _get_llm_output_language(ai_config=ai_config, request=request) @@ -2728,7 +2727,6 @@ class DocumentSelectionMixin: *, user: User, validated_data: dict[str, Any], - permission_codename: str = "view_document", ) -> list[int]: if not validated_data.get("all", False): # if all is not true, just pass through the provided document ids @@ -2741,10 +2739,8 @@ class DocumentSelectionMixin: for key, value in filters.items() if key not in _TANTIVY_SEARCH_PARAM_NAMES } - permitted_documents = get_objects_for_user_owner_aware( - user, - permission_codename, - Document, + permitted_documents = Document.objects.filter( + id__in=permitted_document_ids(user), ) # orm-filtered docs filtered_documents = DocumentFilterSet( @@ -3352,10 +3348,8 @@ class SelectionDataView(GenericAPIView[Any]): serializer.is_valid(raise_exception=True) ids = serializer.validated_data.get("documents") - permitted_documents = get_objects_for_user_owner_aware( - request.user, - "documents.view_document", - Document, + permitted_documents = Document.objects.filter( + id__in=permitted_document_ids(request.user), ) if permitted_documents.filter(pk__in=ids).count() != len(ids): return HttpResponseForbidden("Insufficient permissions") @@ -3527,10 +3521,8 @@ class GlobalSearchView(PassUserMixin): OBJECT_LIMIT = 3 docs = [] if request.user.has_perm("documents.view_document"): - all_docs = get_objects_for_user_owner_aware( - request.user, - "view_document", - Document, + all_docs = Document.objects.filter( + id__in=permitted_document_ids(request.user), ) if db_only: docs = all_docs.filter(title__icontains=query)[:OBJECT_LIMIT] @@ -3734,11 +3726,7 @@ class StatisticsView(GenericAPIView[Any]): documents = ( Document.objects.all() if can_view_global_stats - else get_objects_for_user_owner_aware( - user, - "documents.view_document", - Document, - ) + else Document.objects.filter(id__in=permitted_document_ids(user)) ).filter(root_document__isnull=True) tags = ( Tag.objects.all()