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
This commit is contained in:
Trenton Holmes
2026-08-10 19:39:57 -07:00
committed by stumpylog
parent fc242bb570
commit 8308ab0552
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):
return HttpResponseForbidden("Insufficient permissions")
documents = [document]
documents = Document.objects.filter(pk=document.pk)
else:
documents = Document.objects.filter(
id__in=permitted_document_ids(request.user),
+21 -6
View File
@@ -2,6 +2,8 @@ import json
import logging
import sys
from django.db.models import QuerySet
from documents.models import Document
from paperless.config import AIConfig
from paperless_ai.client import AIClient
@@ -82,10 +84,21 @@ def _build_document_reference(
def _get_document_references(
documents: list[Document],
documents: QuerySet[Document],
top_nodes: list,
) -> 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]] = []
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(
query_str: str,
documents: list[Document],
documents: QuerySet[Document],
output_language: str | None = None,
):
try:
@@ -135,10 +148,10 @@ def stream_chat_with_documents(
def _stream_chat_with_documents(
query_str: str,
documents: list[Document],
documents: QuerySet[Document],
output_language: str | None = None,
):
if not documents:
if not documents.exists():
yield CHAT_NO_CONTENT_MESSAGE
return
@@ -148,7 +161,9 @@ def _stream_chat_with_documents(
from llama_index.core.retrievers import VectorIndexRetriever
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
# 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
import pytest
from django.db.models.signals import post_init
from llama_index.core import settings as llama_settings
from llama_index.core.embeddings.mock_embed_model import MockEmbedding
from llama_index.core.schema import TextNode
from documents.models import Document
from documents.tests.factories import DocumentFactory
from paperless_ai import chat
from paperless_ai import indexing
@@ -36,16 +38,6 @@ def patch_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(
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(
("output_language", "expected_language_line"),
[
@@ -107,9 +106,10 @@ def test_build_refine_prompt(
@pytest.mark.django_db
def test_stream_chat_with_one_document_retrieval(
mock_document,
patch_embed_nodes,
) -> None:
document = DocumentFactory.create(title="Test Document", content="ignored")
documents = Document.objects.filter(pk=document.pk)
with (
patch("paperless_ai.chat.AIClient") as mock_client_cls,
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.llm = MagicMock()
mock_node = TextNode(
text="This is node content.",
metadata={"document_id": str(mock_document.pk), "title": "Test Document"},
)
mock_index = MagicMock()
# Simulate get_nodes returning nodes (content exists)
mock_index.vector_store.get_nodes.return_value = [mock_node]
mock_index.vector_store.get_nodes.return_value = [
TextNode(
text="This is node content.",
metadata={"document_id": str(document.pk), "title": "Test Document"},
),
]
mock_load_index.return_value = mock_index
mock_retriever_instance = MagicMock()
mock_retriever_instance.retrieve.return_value = [
MagicMock(
metadata={
"document_id": str(mock_document.pk),
"title": "Test Document",
},
metadata={"document_id": str(document.pk), "title": "Test Document"},
),
]
@@ -153,7 +150,7 @@ def test_stream_chat_with_one_document_retrieval(
"llama_index.core.retrievers.VectorIndexRetriever",
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?")
synthesizer_kwargs = mock_get_response_synthesizer.call_args.kwargs
@@ -166,13 +163,16 @@ def test_stream_chat_with_one_document_retrieval(
output,
expected_chunks=["chunk1", "chunk2"],
expected_references=[
{"id": mock_document.pk, "title": "Test Document"},
{"id": document.pk, "title": "Test Document"},
],
)
@pytest.mark.django_db
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 (
patch("paperless_ai.chat.AIClient") as mock_client_cls,
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.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()
# Simulate get_nodes returning nodes (content exists)
mock_index.vector_store.get_nodes.return_value = [mock_node1, mock_node2]
mock_index.vector_store.get_nodes.return_value = [
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_retriever_instance = MagicMock()
mock_retriever_instance.retrieve.return_value = [
MagicMock(metadata={"document_id": "1", "title": "Document 1"}),
MagicMock(metadata={"document_id": "2", "title": "Document 2"}),
MagicMock(metadata={"document_id": str(doc1.pk), "title": "Document 1"}),
MagicMock(metadata={"document_id": str(doc2.pk), "title": "Document 2"}),
]
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.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(
"llama_index.core.retrievers.VectorIndexRetriever",
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?")
patch_embed_nodes.assert_not_called()
@@ -225,15 +222,15 @@ def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> Non
output,
expected_chunks=["chunk1", "chunk2"],
expected_references=[
{"id": 1, "title": "Document 1"},
{"id": 2, "title": "Document 2"},
{"id": doc1.pk, "title": "Document 1"},
{"id": doc2.pk, "title": "Document 2"},
],
)
def test_stream_chat_empty_document_list() -> None:
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()
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_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."]
@@ -282,7 +281,9 @@ def test_stream_chat_unexpected_failure_returns_generic_error(caplog) -> None:
)
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 "Failed to stream document chat response" in caplog.text
@@ -298,7 +299,12 @@ class TestStreamChatRetrieval:
) -> None:
doc = DocumentFactory.create(content="hello world")
# 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
def test_chat_filter_contains_only_requested_document_ids(
@@ -332,7 +338,12 @@ class TestStreamChatRetrieval:
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"
filt = captured_filters[0]
@@ -340,3 +351,47 @@ class TestStreamChatRetrieval:
filter_values = filt.filters[0].value
assert str(included.pk) 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"}]