diff --git a/src/documents/tests/test_api_chat.py b/src/documents/tests/test_api_chat.py index 421c74558..fae18f2c0 100644 --- a/src/documents/tests/test_api_chat.py +++ b/src/documents/tests/test_api_chat.py @@ -1,11 +1,18 @@ from __future__ import annotations +from typing import TYPE_CHECKING from unittest import mock +import pytest +from django.contrib.auth.models import Permission from django.contrib.auth.models import User from rest_framework import status +from rest_framework.test import APIClient from rest_framework.test import APITestCase +if TYPE_CHECKING: + from pytest_mock import MockerFixture + class TestChatStreamingViewInputValidation(APITestCase): def setUp(self) -> None: @@ -42,3 +49,74 @@ class TestChatStreamingViewInputValidation(APITestCase): format="json", ) assert resp.status_code == status.HTTP_400_BAD_REQUEST + + +@pytest.mark.django_db +class TestChatStreamingViewUnrestrictedFlag: + """ChatStreamingView must only skip the vector store's document id + filter (``unrestricted=True``) for a caller who can see every document + -- an active superuser -- never for a regular user, regardless of how + many documents that user happens to be permitted to view. + """ + + ENDPOINT = "/api/documents/chat/" + + @pytest.fixture + def mocked_stream_chat(self, mocker: MockerFixture) -> mock.MagicMock: + """AI enabled, with stream_chat_with_documents patched so the view + never touches the real vector store; returns the patched callable so + tests can inspect how it was called. + """ + mocker.patch("documents.views.AIConfig").return_value.ai_enabled = True + return mocker.patch( + "documents.views.stream_chat_with_documents", + return_value=iter(()), + ) + + @pytest.fixture + def viewer_client(self, user_client: APIClient, regular_user: User) -> APIClient: + """The conftest regular-user client, additionally granted + view_document -- able to see every document without being a + superuser. + """ + regular_user.user_permissions.add( + *Permission.objects.filter(codename="view_document"), + ) + return user_client + + @pytest.mark.parametrize( + ("client_fixture", "expected_unrestricted"), + [ + pytest.param("admin_client", True, id="superuser_is_unrestricted"), + pytest.param("viewer_client", False, id="regular_user_is_restricted"), + ], + ) + def test_unrestricted_only_for_superuser( + self, + request: pytest.FixtureRequest, + mocked_stream_chat: mock.MagicMock, + client_fixture: str, + *, + expected_unrestricted: bool, + ) -> None: + """ + GIVEN: + - A superuser, or a regular user holding view_document + WHEN: + - They post a chat question with no document_id + THEN: + - stream_chat_with_documents is called with unrestricted=True for + the superuser and unrestricted=False for the regular user, even + though that user can view every document + """ + client: APIClient = request.getfixturevalue(client_fixture) + + client.post( + self.ENDPOINT, + data={"q": "What's in these documents?"}, + format="json", + ) + + assert ( + mocked_stream_chat.call_args.kwargs["unrestricted"] is expected_unrestricted + ) diff --git a/src/documents/views.py b/src/documents/views.py index 4e33d9f9d..258ff9976 100644 --- a/src/documents/views.py +++ b/src/documents/views.py @@ -180,6 +180,7 @@ from documents.permissions import has_system_status_permission from documents.permissions import permitted_document_ids from documents.permissions import permitted_object_ids from documents.permissions import set_permissions_for_object +from documents.permissions import user_is_unrestricted from documents.plugins.date_parsing import get_date_parser from documents.schema import generate_object_with_permissions_schema from documents.search import SearchHit @@ -2318,6 +2319,7 @@ class ChatStreamingView(GenericAPIView[Any]): question = serializer.validated_data["q"] doc_id = serializer.validated_data.get("document_id") + unrestricted = False if doc_id: try: @@ -2333,6 +2335,7 @@ class ChatStreamingView(GenericAPIView[Any]): documents = Document.objects.filter( id__in=permitted_document_ids(request.user), ) + unrestricted = user_is_unrestricted(request.user) output_language = get_llm_output_language( ai_config=ai_config, @@ -2343,6 +2346,7 @@ class ChatStreamingView(GenericAPIView[Any]): stream_chat_with_documents( query_str=question, documents=documents, + unrestricted=unrestricted, output_language=output_language, ), content_type="text/event-stream", diff --git a/src/paperless_ai/chat.py b/src/paperless_ai/chat.py index 03e64af36..80abe3021 100644 --- a/src/paperless_ai/chat.py +++ b/src/paperless_ai/chat.py @@ -95,12 +95,15 @@ def _format_chat_metadata_trailer(references: list[dict[str, int | str]]) -> str def stream_chat_with_documents( query_str: str, documents: QuerySet[Document], + *, + unrestricted: bool = False, output_language: str | None = None, ): try: yield from _stream_chat_with_documents( query_str, documents, + unrestricted=unrestricted, output_language=output_language, ) except Exception as e: @@ -112,6 +115,8 @@ def _stream_chat_with_documents( query_str: str, documents: QuerySet[Document], output_language: str | None = None, + *, + unrestricted: bool = False, ): if not documents.exists(): yield CHAT_NO_CONTENT_MESSAGE @@ -123,8 +128,17 @@ def _stream_chat_with_documents( from llama_index.core.retrievers import VectorIndexRetriever config = AIConfig() - filters = _document_id_filters( - str(pk) for pk in documents.values_list("pk", flat=True) + # An unrestricted caller can see the entire corpus, so an id filter here + # would never narrow the search -- only cost an IN() list that can blow + # past the vector store's bound-parameter safety limit on large + # installs (see _MAX_IN_VALUES in vector_store.py). Skip it and let the + # retriever search the whole index. + filters = ( + None + if unrestricted + else _document_id_filters( + str(pk) for pk in documents.values_list("pk", flat=True) + ) ) # Hold the shared read lock for the whole operation: the query engine diff --git a/src/paperless_ai/tests/test_chat.py b/src/paperless_ai/tests/test_chat.py index 15c5a545d..b59414bfe 100644 --- a/src/paperless_ai/tests/test_chat.py +++ b/src/paperless_ai/tests/test_chat.py @@ -1,4 +1,8 @@ +from __future__ import annotations + import json +from typing import TYPE_CHECKING +from typing import Any from unittest.mock import MagicMock from unittest.mock import patch @@ -18,6 +22,11 @@ from paperless_ai.chat import _build_chat_prompt from paperless_ai.chat import _build_refine_prompt from paperless_ai.chat import stream_chat_with_documents +if TYPE_CHECKING: + from pathlib import Path + + import pytest_mock + @pytest.fixture(autouse=True) def patch_embed_model(): @@ -312,6 +321,30 @@ def test_stream_chat_unexpected_failure_returns_generic_error(caplog) -> None: @pytest.mark.django_db class TestStreamChatRetrieval: + @pytest.fixture + def captured_filters(self, mocker: pytest_mock.MockerFixture) -> list[Any]: + """Stub out the AI client and the retriever, capturing the ``filters`` + kwarg of every VectorIndexRetriever construction. + + VectorIndexRetriever is imported inside _stream_chat_with_documents, + so it is patched at the llama_index source for the lazy import to + pick it up. + """ + captured: list[Any] = [] + retriever = mocker.MagicMock() + retriever.retrieve.return_value = [] + + def capture_retriever(*args, **kwargs) -> pytest_mock.MockType: + captured.append(kwargs.get("filters")) + return retriever + + mocker.patch("paperless_ai.chat.AIClient") + mocker.patch( + "llama_index.core.retrievers.VectorIndexRetriever", + side_effect=capture_retriever, + ) + return captured + def test_no_nodes_yields_no_content_message( self, temp_llm_index_dir, @@ -329,9 +362,9 @@ class TestStreamChatRetrieval: def test_chat_filter_contains_only_requested_document_ids( self, - temp_llm_index_dir, - mock_embed_model, - mocker, + temp_llm_index_dir: Path, + mock_embed_model: pytest_mock.MockType, + captured_filters: list[Any], ) -> None: """The MetadataFilter passed to the retriever must be scoped to the requested documents only — content from other indexed documents must @@ -342,22 +375,6 @@ class TestStreamChatRetrieval: indexing.llm_index_add_or_update_document(included) indexing.llm_index_add_or_update_document(excluded) - # VectorIndexRetriever is imported inside _stream_chat_with_documents; - # patch it at the llama_index source so the lazy import picks it up. - captured_filters = [] - mock_retriever = mocker.MagicMock() - mock_retriever.retrieve.return_value = [] - - def capture_retriever(*args, **kwargs): - captured_filters.append(kwargs.get("filters")) - return mock_retriever - - mocker.patch("paperless_ai.chat.AIClient") - mocker.patch( - "llama_index.core.retrievers.VectorIndexRetriever", - side_effect=capture_retriever, - ) - list( chat.stream_chat_with_documents( "question?", @@ -372,6 +389,37 @@ class TestStreamChatRetrieval: assert str(included.pk) in filter_values assert str(excluded.pk) not in filter_values + def test_unrestricted_chat_skips_document_id_filter( + self, + temp_llm_index_dir: Path, + mock_embed_model: pytest_mock.MockType, + captured_filters: list[Any], + ) -> None: + """ + GIVEN: + - A document indexed in the vector store + WHEN: + - stream_chat_with_documents is called with unrestricted=True + THEN: + - The retriever receives no document id filter (filters=None), + since an unrestricted caller can see every document and an id + filter here would only risk the vector store's IN-filter + safety limit on large installs, never narrow the search + """ + included = DocumentFactory.create(content="included document content") + indexing.llm_index_add_or_update_document(included) + + list( + chat.stream_chat_with_documents( + "question?", + Document.objects.filter(pk=included.pk), + unrestricted=True, + ), + ) + + assert captured_filters, "VectorIndexRetriever was never constructed" + assert captured_filters[0] is None + @pytest.mark.django_db def test_get_document_references_only_queries_referenced_documents( self,