diff --git a/src/paperless_ai/ai_classifier.py b/src/paperless_ai/ai_classifier.py index 0559446d7..391edca30 100644 --- a/src/paperless_ai/ai_classifier.py +++ b/src/paperless_ai/ai_classifier.py @@ -6,6 +6,7 @@ from django.contrib.auth.models import User from documents.models import Document from documents.permissions import permitted_document_ids from documents.permissions import user_is_unrestricted +from documents.versioning import annotate_effective_content from paperless.config import AIConfig from paperless_ai.base_model import ClassificationSuggestions from paperless_ai.base_model import TaxonomyChoiceDict @@ -111,7 +112,7 @@ def build_prompt_without_rag( ) -> str: filename = document.filename or "" content = truncate_content( - document.content[:4000] or "", + (document.get_effective_content() or "")[:4000], chunk_size=config.llm_embedding_chunk_size, context_size=config.llm_context_size, ) @@ -225,7 +226,9 @@ def get_taxonomy_context( # similar_documents is already ordered by descending weight; don't lose it. similar_document_ids = [s["document_id"] for s in similar_documents] - similar_documents_by_id = Document.objects.in_bulk(similar_document_ids) + similar_documents_by_id = annotate_effective_content( + Document.objects.all(), + ).in_bulk(similar_document_ids) similar_docs = [ similar_documents_by_id[document_id] for document_id in similar_document_ids @@ -233,7 +236,7 @@ def get_taxonomy_context( ][:max_docs] context_blocks = [] for similar in similar_docs: - text = similar.content[:1000] or "" + text = (similar.get_effective_content() or "")[:1000] title = similar.title or similar.filename or "Untitled" context_blocks.append(f"TITLE: {title}\n{text}") except Exception: diff --git a/src/paperless_ai/embedding.py b/src/paperless_ai/embedding.py index 0fd71423f..f4466686a 100644 --- a/src/paperless_ai/embedding.py +++ b/src/paperless_ai/embedding.py @@ -135,6 +135,6 @@ def build_llm_index_text(doc: Document) -> str: lines.append(f"Custom Field - {instance.field.name}: {instance}") lines.append("\nContent:\n") - lines.append(doc.content or "") + lines.append(doc.get_effective_content() or "") return _normalize_llm_index_text("\n".join(lines)) diff --git a/src/paperless_ai/indexing.py b/src/paperless_ai/indexing.py index 681ef7ad5..91a2e7e30 100644 --- a/src/paperless_ai/indexing.py +++ b/src/paperless_ai/indexing.py @@ -17,6 +17,7 @@ from documents.models import PaperlessTask from documents.utils import IterWrapper from documents.utils import QuerySetStream from documents.utils import identity +from documents.versioning import annotate_effective_content from documents.versioning import root_document_ids from paperless.config import AIConfig from paperless_ai.db import db_connection_released @@ -444,10 +445,10 @@ def update_llm_index( "Skipping LLM index update: migration check deferred; " "will retry next run." ) - documents = ( + documents = annotate_effective_content( Document.objects.filter(root_document__isnull=True) .select_related("correspondent", "document_type", "storage_path") - .prefetch_related("tags", "notes", "custom_fields__field") + .prefetch_related("tags", "notes", "custom_fields__field"), ) no_documents = not documents.exists() @@ -694,7 +695,7 @@ def retrieve_similar_nodes( ) query_text = truncate_embedding_query( - (document.title or "") + "\n" + (document.content or ""), + (document.title or "") + "\n" + (document.get_effective_content() or ""), chunk_size=config.llm_embedding_chunk_size, ) # Hold the shared read lock for the whole retrieval so the connection is diff --git a/src/paperless_ai/tests/test_ai_classifier.py b/src/paperless_ai/tests/test_ai_classifier.py index 1e018a62f..ccc2cd08a 100644 --- a/src/paperless_ai/tests/test_ai_classifier.py +++ b/src/paperless_ai/tests/test_ai_classifier.py @@ -55,6 +55,7 @@ def mock_document(): doc.storage_path = None doc.archive_serial_number = "12345" doc.content = "This is the document content." + doc.get_effective_content.return_value = "This is the document content." cf1 = MagicMock(__str__=lambda x: "Value1") cf1.field = MagicMock() @@ -434,6 +435,55 @@ def test_get_taxonomy_context_preserves_similarity_order_and_distinct_documents( ) +@pytest.mark.django_db +class TestClassifierEffectiveContent: + """A root document's text for the LLM is its newest version's content.""" + + @staticmethod + def _root_with_version() -> Document: + root = DocumentFactory(title="Statement", content="stale text") + DocumentFactory(root_document=root, version_index=1, content="latest text") + return root + + def test_prompt_uses_the_newest_versions_content(self) -> None: + """ + GIVEN: + - A root document with a version + WHEN: + - The classification prompt is built for the root + THEN: + - It contains the newest version's content + """ + prompt = build_prompt_without_rag(self._root_with_version(), AIConfig()) + + assert "latest text" in prompt + assert "stale text" not in prompt + + @override_settings(LLM_EMBEDDING_BACKEND="huggingface") + def test_similar_document_context_uses_the_newest_versions_content(self) -> None: + """ + GIVEN: + - A similar root document with a version + WHEN: + - The similar-document context is built + THEN: + - It contains the newest version's content + """ + similar = self._root_with_version() + document = DocumentFactory(content="Some content") + fake_nodes = [ + SimpleNamespace(metadata={"document_id": str(similar.pk)}, score=0.9), + ] + + with patch( + "paperless_ai.ai_classifier.retrieve_similar_nodes", + return_value=fake_nodes, + ): + _candidates, context = get_taxonomy_context(document, user=None) + + assert context == "TITLE: Statement\nlatest text" + + @pytest.mark.django_db @override_settings(LLM_EMBEDDING_BACKEND="huggingface") def test_get_taxonomy_context_no_similar_docs(): diff --git a/src/paperless_ai/tests/test_ai_indexing.py b/src/paperless_ai/tests/test_ai_indexing.py index 8c91dcf3a..70285b590 100644 --- a/src/paperless_ai/tests/test_ai_indexing.py +++ b/src/paperless_ai/tests/test_ai_indexing.py @@ -438,6 +438,31 @@ class TestLlmIndexVersions: assert set(indexed) == {str(root.pk)} + def test_rebuild_indexes_the_newest_versions_content( + self, + temp_llm_index_dir: Path, + mock_embed_model: FakeEmbedding, + mocker: pytest_mock.MockerFixture, + ) -> None: + """ + GIVEN: + - A root document with a version + WHEN: + - The LLM index is rebuilt + THEN: + - The root's text for the index is the version's content, answered + from the query rather than a query per document + """ + root = DocumentFactory(content="stale text") + DocumentFactory(root_document=root, version_index=1, content="latest text") + spy = mocker.spy(indexing, "build_document_node") + + indexing.update_llm_index(rebuild=True) + + indexed = spy.call_args.args[0] + assert indexed.pk == root.pk + assert indexed.effective_content == "latest text" + def test_incremental_update_by_version_id_refreshes_the_root( self, temp_llm_index_dir: Path, diff --git a/src/paperless_ai/tests/test_embedding.py b/src/paperless_ai/tests/test_embedding.py index 126a1a9f3..3d2c04cb7 100644 --- a/src/paperless_ai/tests/test_embedding.py +++ b/src/paperless_ai/tests/test_embedding.py @@ -12,6 +12,7 @@ from paperless_ai.embedding import _normalize_llm_index_text from paperless_ai.embedding import build_llm_index_text from paperless_ai.embedding import get_configured_model_name from paperless_ai.embedding import get_embedding_model +from paperless_testing.factories import DocumentFactory @pytest.fixture @@ -46,6 +47,7 @@ def mock_document(): doc.correspondent.name = "Test Correspondent" doc.archive_serial_number = "12345" doc.content = "This is the document content." + doc.get_effective_content.return_value = "This is the document content." cf1 = MagicMock(__str__=lambda x: "Value1") cf1.field = MagicMock() @@ -280,7 +282,7 @@ def test_build_llm_index_text(mock_document): def test_build_llm_index_text_normalizes_ocr_punctuation_runs(mock_document): - mock_document.content = ( + mock_document.get_effective_content.return_value = ( "Introduction ................................................ 7\n" "Hardware Limitation ________________________________________ 9\n" "Keep short punctuation like INV-100 and ellipses..." @@ -294,6 +296,43 @@ def test_build_llm_index_text_normalizes_ocr_punctuation_runs(mock_document): assert "ellipses..." in result +@pytest.mark.django_db +class TestBuildLlmIndexTextVersions: + """A root document is indexed with its effective content, like in the search index.""" + + def test_root_uses_the_newest_versions_content(self) -> None: + """ + GIVEN: + - A root document with two versions + WHEN: + - The LLM index text is built for the root + THEN: + - It contains the newest version's content and not the others' + """ + root = DocumentFactory(content="stale text") + DocumentFactory(root_document=root, version_index=1, content="older text") + DocumentFactory(root_document=root, version_index=2, content="latest text") + + text = build_llm_index_text(root) + + assert "latest text" in text + assert "stale text" not in text + assert "older text" not in text + + def test_root_without_versions_uses_its_own_content(self) -> None: + """ + GIVEN: + - A root document without versions + WHEN: + - The LLM index text is built for it + THEN: + - It contains the document's own content + """ + root = DocumentFactory(content="own text") + + assert "own text" in build_llm_index_text(root) + + def test_normalize_llm_index_text_collapses_ocr_leaders_without_joining_lines(): assert _normalize_llm_index_text("A........B\nC____D----E") == "A B\nC D E"