From 968578e26b42b0d5eea6a2cb9fa53bd7b9c3bb30 Mon Sep 17 00:00:00 2001 From: shamoon <4887959+shamoon@users.noreply.github.com> Date: Mon, 25 May 2026 07:07:50 -0700 Subject: [PATCH] Enhancement: dont re-build vector index with each chat --- src/paperless_ai/chat.py | 86 +++++++++++++++++++++++++++-- src/paperless_ai/tests/test_chat.py | 44 ++++++++++----- 2 files changed, 113 insertions(+), 17 deletions(-) diff --git a/src/paperless_ai/chat.py b/src/paperless_ai/chat.py index 40c901db7..8608a9cde 100644 --- a/src/paperless_ai/chat.py +++ b/src/paperless_ai/chat.py @@ -68,6 +68,84 @@ def _format_chat_metadata_trailer(references: list[dict[str, int | str]]) -> str ) +def _get_document_filtered_retriever(index, doc_ids: set[str], similarity_top_k: int): + from llama_index.core.base.base_retriever import BaseRetriever + from llama_index.core.schema import NodeWithScore + from llama_index.core.vector_stores import VectorStoreQuery + + class DocumentFilteredFaissRetriever(BaseRetriever): + def __init__(self): + super().__init__() + self._cached_query_str = None + self._cached_nodes = [] + + def _retrieve(self, query_bundle): + if query_bundle.query_str == self._cached_query_str: + return self._cached_nodes + + if query_bundle.embedding is None: + query_bundle.embedding = ( + index._embed_model.get_agg_embedding_from_queries( + query_bundle.embedding_strs, + ) + ) + + faiss_index = index.vector_store._faiss_index + max_top_k = faiss_index.ntotal + if max_top_k == 0: + self._cached_query_str = query_bundle.query_str + self._cached_nodes = [] + return [] + + query_top_k = min(max(similarity_top_k, 1), max_top_k) + allowed_nodes: list[NodeWithScore] = [] + seen_node_ids: set[str] = set() + + while query_top_k <= max_top_k: + query_result = index.vector_store.query( + VectorStoreQuery( + query_embedding=query_bundle.embedding, + similarity_top_k=query_top_k, + ), + ) + + allowed_nodes = [] + seen_node_ids = set() + for vector_id, score in zip( + query_result.ids or [], + query_result.similarities or [], + strict=False, + ): + node_id = index.index_struct.nodes_dict.get(vector_id) + if node_id is None or node_id in seen_node_ids: + continue + + node = index.docstore.docs.get(node_id) + if node is None or node.metadata.get("document_id") not in doc_ids: + continue + + seen_node_ids.add(node_id) + allowed_nodes.append(NodeWithScore(node=node, score=score)) + + if len(allowed_nodes) >= similarity_top_k: + self._cached_query_str = query_bundle.query_str + self._cached_nodes = allowed_nodes + return allowed_nodes + + if query_top_k == max_top_k: + self._cached_query_str = query_bundle.query_str + self._cached_nodes = allowed_nodes + return allowed_nodes + + query_top_k = min(query_top_k * 2, max_top_k) + + self._cached_query_str = query_bundle.query_str + self._cached_nodes = allowed_nodes + return allowed_nodes + + return DocumentFilteredFaissRetriever() + + def stream_chat_with_documents(query_str: str, documents: list[Document]): client = AIClient() index = load_or_build_index() @@ -86,14 +164,14 @@ def stream_chat_with_documents(query_str: str, documents: list[Document]): yield "Sorry, I couldn't find any content to answer your question." return - from llama_index.core import VectorStoreIndex from llama_index.core.prompts import PromptTemplate from llama_index.core.query_engine import RetrieverQueryEngine from llama_index.core.response_synthesizers import get_response_synthesizer - local_index = VectorStoreIndex(nodes=nodes) - retriever = local_index.as_retriever( - similarity_top_k=CHAT_RETRIEVER_TOP_K, + retriever = _get_document_filtered_retriever( + index, + set(doc_ids), + CHAT_RETRIEVER_TOP_K, ) top_nodes = retriever.retrieve(query_str) diff --git a/src/paperless_ai/tests/test_chat.py b/src/paperless_ai/tests/test_chat.py index c7beb50d0..2511a8bbd 100644 --- a/src/paperless_ai/tests/test_chat.py +++ b/src/paperless_ai/tests/test_chat.py @@ -3,7 +3,6 @@ from unittest.mock import MagicMock from unittest.mock import patch import pytest -from llama_index.core import VectorStoreIndex from llama_index.core.schema import TextNode from paperless_ai.chat import CHAT_METADATA_DELIMITER @@ -29,7 +28,7 @@ def patch_embed_nodes(): mock_embed_nodes.side_effect = lambda nodes, *_args, **_kwargs: { node.node_id: [0.1] * 1536 for node in nodes } - yield + yield mock_embed_nodes @pytest.fixture @@ -57,7 +56,25 @@ def assert_chat_output( } -def test_stream_chat_with_one_document_retrieval(mock_document) -> None: +def add_vector_query_results(mock_index, nodes: list[TextNode]) -> None: + mock_index.index_struct.nodes_dict = { + str(vector_id): node.node_id for vector_id, node in enumerate(nodes) + } + mock_index.docstore.docs.get.side_effect = { + node.node_id: node for node in nodes + }.get + mock_index.vector_store._faiss_index.ntotal = len(nodes) + mock_index.vector_store.query.return_value = MagicMock( + ids=list(mock_index.index_struct.nodes_dict), + similarities=[0.1] * len(nodes), + ) + mock_index._embed_model.get_agg_embedding_from_queries.return_value = [0.1] * 1536 + + +def test_stream_chat_with_one_document_retrieval( + mock_document, + patch_embed_nodes, +) -> None: with ( patch("paperless_ai.chat.AIClient") as mock_client_cls, patch("paperless_ai.chat.load_or_build_index") as mock_load_index, @@ -75,6 +92,7 @@ def test_stream_chat_with_one_document_retrieval(mock_document) -> None: ) mock_index = MagicMock() mock_index.docstore.docs.values.return_value = [mock_node] + add_vector_query_results(mock_index, [mock_node]) mock_load_index.return_value = mock_index mock_response_stream = MagicMock() @@ -86,6 +104,7 @@ def test_stream_chat_with_one_document_retrieval(mock_document) -> None: output = list(stream_chat_with_documents("What is this?", [mock_document])) mock_query_engine.query.assert_called_once_with("What is this?") + patch_embed_nodes.assert_not_called() assert_chat_output( output, expected_chunks=["chunk1", "chunk2"], @@ -102,7 +121,6 @@ def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> Non patch( "llama_index.core.query_engine.RetrieverQueryEngine.from_args", ) as mock_query_engine_cls, - patch.object(VectorStoreIndex, "as_retriever") as mock_as_retriever, ): # Mock AIClient and LLM mock_client = MagicMock() @@ -118,12 +136,6 @@ def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> Non text="Content for doc 2.", metadata={"document_id": "2", "title": "Document 2"}, ) - mock_index = MagicMock() - mock_index.docstore.docs.values.return_value = [mock_node1, mock_node2] - mock_load_index.return_value = mock_index - - # Patch as_retriever to return a retriever whose retrieve() returns mock_node1 and mock_node2 - mock_retriever = MagicMock() mock_duplicate_node = TextNode( text="More content for doc 1.", metadata={"document_id": "1", "title": "Document 1 Duplicate"}, @@ -132,13 +144,18 @@ def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> Non text="Content for doc 3.", metadata={"document_id": "3", "title": "Document 3"}, ) - mock_retriever.retrieve.return_value = [ + mock_index = MagicMock() + mock_index.docstore.docs.values.return_value = [ mock_node1, - mock_duplicate_node, mock_node2, + mock_duplicate_node, mock_foreign_node, ] - mock_as_retriever.return_value = mock_retriever + add_vector_query_results( + mock_index, + [mock_node1, mock_duplicate_node, mock_node2, mock_foreign_node], + ) + mock_load_index.return_value = mock_index # Mock response stream mock_response_stream = MagicMock() @@ -156,6 +173,7 @@ def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> Non output = list(stream_chat_with_documents("What's up?", [doc1, doc2])) mock_query_engine.query.assert_called_once_with("What's up?") + patch_embed_nodes.assert_not_called() assert_chat_output( output, expected_chunks=["chunk1", "chunk2"],