Compare commits

...
Author SHA1 Message Date
Trenton Holmes 8308ab0552 Fix: pass document chat queries as a QuerySet instead of a materialized list
In tracemalloc based profiling, not materializing the whole Document list
reduced memory to approximately 20% of the baseline, with a peak memory
that scaled with the library size.  Now, the lazt queryset is used and only
the needed pk value is actually contributing to memory
2026-08-10 19:39:57 -07:00
3 changed files with 123 additions and 53 deletions
+1 -1
View File
@@ -2267,7 +2267,7 @@ class ChatStreamingView(GenericAPIView[Any]):
if not has_perms_owner_aware(request.user, "view_document", document): if not has_perms_owner_aware(request.user, "view_document", document):
return HttpResponseForbidden("Insufficient permissions") return HttpResponseForbidden("Insufficient permissions")
documents = [document] documents = Document.objects.filter(pk=document.pk)
else: else:
documents = Document.objects.filter( documents = Document.objects.filter(
id__in=permitted_document_ids(request.user), id__in=permitted_document_ids(request.user),
+21 -6
View File
@@ -2,6 +2,8 @@ import json
import logging import logging
import sys import sys
from django.db.models import QuerySet
from documents.models import Document from documents.models import Document
from paperless.config import AIConfig from paperless.config import AIConfig
from paperless_ai.client import AIClient from paperless_ai.client import AIClient
@@ -82,10 +84,21 @@ def _build_document_reference(
def _get_document_references( def _get_document_references(
documents: list[Document], documents: QuerySet[Document],
top_nodes: list, top_nodes: list,
) -> list[dict[str, int | str]]: ) -> list[dict[str, int | str]]:
allowed_documents = {doc.pk: doc for doc in documents} candidate_ids: set[int] = set()
for node in top_nodes:
try:
candidate_ids.add(int(node.metadata["document_id"]))
except (KeyError, TypeError, ValueError): # pragma: no cover
continue
if not candidate_ids:
return []
allowed_documents = {doc.pk: doc for doc in documents.filter(pk__in=candidate_ids)}
references: list[dict[str, int | str]] = [] references: list[dict[str, int | str]] = []
seen_document_ids: set[int] = set() seen_document_ids: set[int] = set()
@@ -119,7 +132,7 @@ def _format_chat_metadata_trailer(references: list[dict[str, int | str]]) -> str
def stream_chat_with_documents( def stream_chat_with_documents(
query_str: str, query_str: str,
documents: list[Document], documents: QuerySet[Document],
output_language: str | None = None, output_language: str | None = None,
): ):
try: try:
@@ -135,10 +148,10 @@ def stream_chat_with_documents(
def _stream_chat_with_documents( def _stream_chat_with_documents(
query_str: str, query_str: str,
documents: list[Document], documents: QuerySet[Document],
output_language: str | None = None, output_language: str | None = None,
): ):
if not documents: if not documents.exists():
yield CHAT_NO_CONTENT_MESSAGE yield CHAT_NO_CONTENT_MESSAGE
return return
@@ -148,7 +161,9 @@ def _stream_chat_with_documents(
from llama_index.core.retrievers import VectorIndexRetriever from llama_index.core.retrievers import VectorIndexRetriever
config = AIConfig() config = AIConfig()
filters = _document_id_filters(str(doc.pk) for doc in documents) filters = _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 # Hold the shared read lock for the whole operation: the query engine
# retrieves from the vector store again during synthesis, so the connection # retrieves from the vector store again during synthesis, so the connection
+101 -46
View File
@@ -3,10 +3,12 @@ from unittest.mock import MagicMock
from unittest.mock import patch from unittest.mock import patch
import pytest import pytest
from django.db.models.signals import post_init
from llama_index.core import settings as llama_settings from llama_index.core import settings as llama_settings
from llama_index.core.embeddings.mock_embed_model import MockEmbedding from llama_index.core.embeddings.mock_embed_model import MockEmbedding
from llama_index.core.schema import TextNode from llama_index.core.schema import TextNode
from documents.models import Document
from documents.tests.factories import DocumentFactory from documents.tests.factories import DocumentFactory
from paperless_ai import chat from paperless_ai import chat
from paperless_ai import indexing from paperless_ai import indexing
@@ -36,16 +38,6 @@ def patch_embed_nodes():
yield mock_embed_nodes yield mock_embed_nodes
@pytest.fixture
def mock_document():
doc = MagicMock()
doc.pk = 1
doc.title = "Test Document"
doc.filename = "test_file.pdf"
doc.content = "This is the document content."
return doc
def assert_chat_output( def assert_chat_output(
output: list[str], output: list[str],
*, *,
@@ -61,6 +53,13 @@ def assert_chat_output(
} }
def _fake_documents_queryset(pks: list[int]) -> MagicMock:
qs = MagicMock()
qs.exists.return_value = bool(pks)
qs.values_list.return_value = pks
return qs
@pytest.mark.parametrize( @pytest.mark.parametrize(
("output_language", "expected_language_line"), ("output_language", "expected_language_line"),
[ [
@@ -107,9 +106,10 @@ def test_build_refine_prompt(
@pytest.mark.django_db @pytest.mark.django_db
def test_stream_chat_with_one_document_retrieval( def test_stream_chat_with_one_document_retrieval(
mock_document,
patch_embed_nodes, patch_embed_nodes,
) -> None: ) -> None:
document = DocumentFactory.create(title="Test Document", content="ignored")
documents = Document.objects.filter(pk=document.pk)
with ( with (
patch("paperless_ai.chat.AIClient") as mock_client_cls, patch("paperless_ai.chat.AIClient") as mock_client_cls,
patch("paperless_ai.chat.load_or_build_index") as mock_load_index, patch("paperless_ai.chat.load_or_build_index") as mock_load_index,
@@ -124,22 +124,19 @@ def test_stream_chat_with_one_document_retrieval(
mock_client_cls.return_value = mock_client mock_client_cls.return_value = mock_client
mock_client.llm = MagicMock() mock_client.llm = MagicMock()
mock_node = TextNode(
text="This is node content.",
metadata={"document_id": str(mock_document.pk), "title": "Test Document"},
)
mock_index = MagicMock() mock_index = MagicMock()
# Simulate get_nodes returning nodes (content exists) mock_index.vector_store.get_nodes.return_value = [
mock_index.vector_store.get_nodes.return_value = [mock_node] TextNode(
text="This is node content.",
metadata={"document_id": str(document.pk), "title": "Test Document"},
),
]
mock_load_index.return_value = mock_index mock_load_index.return_value = mock_index
mock_retriever_instance = MagicMock() mock_retriever_instance = MagicMock()
mock_retriever_instance.retrieve.return_value = [ mock_retriever_instance.retrieve.return_value = [
MagicMock( MagicMock(
metadata={ metadata={"document_id": str(document.pk), "title": "Test Document"},
"document_id": str(mock_document.pk),
"title": "Test Document",
},
), ),
] ]
@@ -153,7 +150,7 @@ def test_stream_chat_with_one_document_retrieval(
"llama_index.core.retrievers.VectorIndexRetriever", "llama_index.core.retrievers.VectorIndexRetriever",
return_value=mock_retriever_instance, return_value=mock_retriever_instance,
): ):
output = list(stream_chat_with_documents("What is this?", [mock_document])) output = list(stream_chat_with_documents("What is this?", documents))
mock_query_engine.query.assert_called_once_with("What is this?") mock_query_engine.query.assert_called_once_with("What is this?")
synthesizer_kwargs = mock_get_response_synthesizer.call_args.kwargs synthesizer_kwargs = mock_get_response_synthesizer.call_args.kwargs
@@ -166,13 +163,16 @@ def test_stream_chat_with_one_document_retrieval(
output, output,
expected_chunks=["chunk1", "chunk2"], expected_chunks=["chunk1", "chunk2"],
expected_references=[ expected_references=[
{"id": mock_document.pk, "title": "Test Document"}, {"id": document.pk, "title": "Test Document"},
], ],
) )
@pytest.mark.django_db @pytest.mark.django_db
def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> None: def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> None:
doc1 = DocumentFactory.create(title="Document 1", content="ignored")
doc2 = DocumentFactory.create(title="Document 2", content="ignored")
documents = Document.objects.filter(pk__in=[doc1.pk, doc2.pk])
with ( with (
patch("paperless_ai.chat.AIClient") as mock_client_cls, patch("paperless_ai.chat.AIClient") as mock_client_cls,
patch("paperless_ai.chat.load_or_build_index") as mock_load_index, patch("paperless_ai.chat.load_or_build_index") as mock_load_index,
@@ -184,23 +184,23 @@ def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> Non
mock_client_cls.return_value = mock_client mock_client_cls.return_value = mock_client
mock_client.llm = MagicMock() mock_client.llm = MagicMock()
mock_node1 = TextNode(
text="Content for doc 1.",
metadata={"document_id": "1", "title": "Document 1"},
)
mock_node2 = TextNode(
text="Content for doc 2.",
metadata={"document_id": "2", "title": "Document 2"},
)
mock_index = MagicMock() mock_index = MagicMock()
# Simulate get_nodes returning nodes (content exists) mock_index.vector_store.get_nodes.return_value = [
mock_index.vector_store.get_nodes.return_value = [mock_node1, mock_node2] TextNode(
text="Content for doc 1.",
metadata={"document_id": str(doc1.pk), "title": "Document 1"},
),
TextNode(
text="Content for doc 2.",
metadata={"document_id": str(doc2.pk), "title": "Document 2"},
),
]
mock_load_index.return_value = mock_index mock_load_index.return_value = mock_index
mock_retriever_instance = MagicMock() mock_retriever_instance = MagicMock()
mock_retriever_instance.retrieve.return_value = [ mock_retriever_instance.retrieve.return_value = [
MagicMock(metadata={"document_id": "1", "title": "Document 1"}), MagicMock(metadata={"document_id": str(doc1.pk), "title": "Document 1"}),
MagicMock(metadata={"document_id": "2", "title": "Document 2"}), MagicMock(metadata={"document_id": str(doc2.pk), "title": "Document 2"}),
] ]
mock_response_stream = MagicMock() mock_response_stream = MagicMock()
@@ -210,14 +210,11 @@ def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> Non
mock_query_engine_cls.return_value = mock_query_engine mock_query_engine_cls.return_value = mock_query_engine
mock_query_engine.query.return_value = mock_response_stream mock_query_engine.query.return_value = mock_response_stream
doc1 = MagicMock(pk=1, title="Document 1", filename="doc1.pdf")
doc2 = MagicMock(pk=2, title="Document 2", filename="doc2.pdf")
with patch( with patch(
"llama_index.core.retrievers.VectorIndexRetriever", "llama_index.core.retrievers.VectorIndexRetriever",
return_value=mock_retriever_instance, return_value=mock_retriever_instance,
): ):
output = list(stream_chat_with_documents("What's up?", [doc1, doc2])) output = list(stream_chat_with_documents("What's up?", documents))
mock_query_engine.query.assert_called_once_with("What's up?") mock_query_engine.query.assert_called_once_with("What's up?")
patch_embed_nodes.assert_not_called() patch_embed_nodes.assert_not_called()
@@ -225,15 +222,15 @@ def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> Non
output, output,
expected_chunks=["chunk1", "chunk2"], expected_chunks=["chunk1", "chunk2"],
expected_references=[ expected_references=[
{"id": 1, "title": "Document 1"}, {"id": doc1.pk, "title": "Document 1"},
{"id": 2, "title": "Document 2"}, {"id": doc2.pk, "title": "Document 2"},
], ],
) )
def test_stream_chat_empty_document_list() -> None: def test_stream_chat_empty_document_list() -> None:
with patch("paperless_ai.chat.load_or_build_index") as mock_load_index: with patch("paperless_ai.chat.load_or_build_index") as mock_load_index:
output = list(stream_chat_with_documents("Any info?", [])) output = list(stream_chat_with_documents("Any info?", Document.objects.none()))
mock_load_index.assert_not_called() mock_load_index.assert_not_called()
assert output == ["Sorry, I couldn't find any content to answer your question."] assert output == ["Sorry, I couldn't find any content to answer your question."]
@@ -253,7 +250,9 @@ def test_stream_chat_no_matching_nodes() -> None:
mock_index.vector_store.get_nodes.return_value = [] mock_index.vector_store.get_nodes.return_value = []
mock_load_index.return_value = mock_index mock_load_index.return_value = mock_index
output = list(stream_chat_with_documents("Any info?", [MagicMock(pk=1)])) output = list(
stream_chat_with_documents("Any info?", _fake_documents_queryset([1])),
)
assert output == ["Sorry, I couldn't find any content to answer your question."] assert output == ["Sorry, I couldn't find any content to answer your question."]
@@ -282,7 +281,9 @@ def test_stream_chat_unexpected_failure_returns_generic_error(caplog) -> None:
) )
mock_retriever_cls.return_value = mock_retriever mock_retriever_cls.return_value = mock_retriever
output = list(stream_chat_with_documents("Any info?", [MagicMock(pk=1)])) output = list(
stream_chat_with_documents("Any info?", _fake_documents_queryset([1])),
)
assert output == [CHAT_ERROR_MESSAGE] assert output == [CHAT_ERROR_MESSAGE]
assert "Failed to stream document chat response" in caplog.text assert "Failed to stream document chat response" in caplog.text
@@ -298,7 +299,12 @@ class TestStreamChatRetrieval:
) -> None: ) -> None:
doc = DocumentFactory.create(content="hello world") doc = DocumentFactory.create(content="hello world")
# Nothing indexed for this document yet. # Nothing indexed for this document yet.
out = list(chat.stream_chat_with_documents("question?", [doc])) out = list(
chat.stream_chat_with_documents(
"question?",
Document.objects.filter(pk=doc.pk),
),
)
assert chat.CHAT_NO_CONTENT_MESSAGE in out assert chat.CHAT_NO_CONTENT_MESSAGE in out
def test_chat_filter_contains_only_requested_document_ids( def test_chat_filter_contains_only_requested_document_ids(
@@ -332,7 +338,12 @@ class TestStreamChatRetrieval:
side_effect=capture_retriever, side_effect=capture_retriever,
) )
list(chat.stream_chat_with_documents("question?", [included])) list(
chat.stream_chat_with_documents(
"question?",
Document.objects.filter(pk=included.pk),
),
)
assert captured_filters, "VectorIndexRetriever was never constructed" assert captured_filters, "VectorIndexRetriever was never constructed"
filt = captured_filters[0] filt = captured_filters[0]
@@ -340,3 +351,47 @@ class TestStreamChatRetrieval:
filter_values = filt.filters[0].value filter_values = filt.filters[0].value
assert str(included.pk) in filter_values assert str(included.pk) in filter_values
assert str(excluded.pk) not in filter_values assert str(excluded.pk) not in filter_values
@pytest.mark.django_db
def test_get_document_references_only_queries_referenced_documents(
self,
django_assert_num_queries,
) -> None:
"""Building references must not hydrate every document the caller is
permitted to see -- only the (<= CHAT_RETRIEVER_TOP_K) documents that
the retriever actually returned nodes for.
"""
referenced = DocumentFactory.create(title="Referenced Document")
# Many more documents are "accessible" but never referenced by a node.
DocumentFactory.create_batch(200)
documents = Document.objects.all()
top_nodes = [
MagicMock(
metadata={
"document_id": str(referenced.pk),
"title": "Referenced Document",
},
),
]
hydrated_count = 0
def _count_hydration(sender, instance, **kwargs):
nonlocal hydrated_count
hydrated_count += 1
post_init.connect(_count_hydration, sender=Document)
try:
# One query: `documents.filter(pk__in=candidate_ids)` for the single
# referenced id. No query should scale with the 200 unreferenced documents.
with django_assert_num_queries(1):
references = chat._get_document_references(documents, top_nodes)
finally:
post_init.disconnect(_count_hydration, sender=Document)
# The bug this guards against: the old code hydrated all 201 accessible
# documents via `{doc.pk: doc for doc in documents}` before filtering by
# top_nodes. Only the referenced document should ever be constructed.
assert hydrated_count == 1
assert references == [{"id": referenced.pk, "title": "Referenced Document"}]