Coverage for paperless_ai/chat.py: 25%
97 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 09:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 09:07 +0000
1import json
2import logging
3import sys
5from django.db.models import QuerySet
7from documents.models import Document
8from paperless.config import AIConfig
9from paperless_ai.client import AIClient
10from paperless_ai.db import db_connection_released
11from paperless_ai.indexing import document_id_filters
12from paperless_ai.indexing import exclude_document_ids_filter
13from paperless_ai.indexing import get_rag_prompt_helper
14from paperless_ai.indexing import load_or_build_index
15from paperless_ai.indexing import read_store
16from paperless_ai.prompts.context import ChatQaPromptContext
17from paperless_ai.prompts.context import ChatRefinePromptContext
18from paperless_ai.prompts.render import render_prompt
20logger = logging.getLogger("paperless_ai.chat")
22CHAT_METADATA_DELIMITER = "\n\n__PAPERLESS_CHAT_METADATA__"
23CHAT_ERROR_MESSAGE = "Sorry, something went wrong while generating a response."
24CHAT_NO_CONTENT_MESSAGE = "Sorry, I couldn't find any content to answer your question."
25MAX_CHAT_REFERENCES = 3
26CHAT_RETRIEVER_TOP_K = 5
29def _build_chat_prompt(output_language: str | None) -> str:
30 return render_prompt(ChatQaPromptContext(output_language=output_language))
33def _build_refine_prompt(output_language: str | None) -> str:
34 return render_prompt(
35 ChatRefinePromptContext(output_language=output_language),
36 )
39def _build_document_reference(
40 document: Document,
41 title: str | None = None,
42) -> dict[str, int | str]:
43 return {
44 "id": document.pk,
45 "title": title or document.title or document.filename,
46 }
49def _get_document_references(
50 documents: QuerySet[Document],
51 top_nodes: list,
52) -> list[dict[str, int | str]]:
53 candidate_ids: set[int] = set()
54 for node in top_nodes:
55 try:
56 candidate_ids.add(int(node.metadata["document_id"]))
57 except (KeyError, TypeError, ValueError): # pragma: no cover
58 continue
60 if not candidate_ids:
61 return []
63 allowed_documents = {doc.pk: doc for doc in documents.filter(pk__in=candidate_ids)}
65 references: list[dict[str, int | str]] = []
66 seen_document_ids: set[int] = set()
68 for node in top_nodes:
69 try:
70 document_id = int(node.metadata["document_id"])
71 except (KeyError, TypeError, ValueError): # pragma: no cover
72 continue
74 if document_id in seen_document_ids or document_id not in allowed_documents:
75 continue
77 seen_document_ids.add(document_id)
78 document = allowed_documents[document_id]
79 references.append(
80 _build_document_reference(document, node.metadata.get("title")),
81 )
83 if len(references) >= MAX_CHAT_REFERENCES: # pragma: no cover
84 break
86 return references
89def _format_chat_metadata_trailer(references: list[dict[str, int | str]]) -> str:
90 return (
91 f"{CHAT_METADATA_DELIMITER}"
92 f"{json.dumps({'references': references}, separators=(',', ':'))}"
93 )
96def stream_chat_with_documents(
97 query_str: str,
98 documents: QuerySet[Document],
99 *,
100 unrestricted: bool = False,
101 output_language: str | None = None,
102):
103 try:
104 yield from _stream_chat_with_documents(
105 query_str,
106 documents,
107 unrestricted=unrestricted,
108 output_language=output_language,
109 )
110 except Exception as e:
111 logger.exception("Failed to stream document chat response: %s", e)
112 yield CHAT_ERROR_MESSAGE
115def _stream_chat_with_documents(
116 query_str: str,
117 documents: QuerySet[Document],
118 *,
119 unrestricted: bool = False,
120 output_language: str | None = None,
121):
122 if not documents.exists():
123 yield CHAT_NO_CONTENT_MESSAGE
124 return
126 from llama_index.core.prompts import PromptTemplate
127 from llama_index.core.query_engine import RetrieverQueryEngine
128 from llama_index.core.response_synthesizers import get_response_synthesizer
129 from llama_index.core.retrievers import VectorIndexRetriever
131 config = AIConfig()
132 if unrestricted:
133 # Exclude trashed ids (usually few) instead of an IN filter over the
134 # full permitted set, which risks the vector store's bound parameter
135 # limit (_MAX_IN_VALUES) on large installs. Trashed documents stay
136 # indexed until permanent deletion (delete_document_from_llm_index
137 # hangs off post_delete, not trash), so must be excluded explicitly.
138 trashed_ids = Document.deleted_objects.values_list("pk", flat=True)
139 filters = exclude_document_ids_filter(str(pk) for pk in trashed_ids)
140 else:
141 filters = document_id_filters(
142 str(pk) for pk in documents.values_list("pk", flat=True)
143 )
145 # Hold the shared read lock for the whole operation: the query engine
146 # retrieves from the vector store again during synthesis, so the connection
147 # must stay open (and the swap must not run) until the stream finishes.
148 with read_store() as store:
149 index = load_or_build_index(config, store)
150 retriever = VectorIndexRetriever(
151 index=index,
152 similarity_top_k=CHAT_RETRIEVER_TOP_K,
153 filters=filters,
154 )
156 # Slow query-embedding + vector search; no Django ORM access happens
157 # during it, so release the pooled DB connection for its duration. See
158 # #12976.
159 with db_connection_released():
160 top_nodes = retriever.retrieve(query_str)
161 if not top_nodes:
162 logger.warning("No nodes found for the given documents.")
163 yield CHAT_NO_CONTENT_MESSAGE
164 return
166 client = AIClient()
168 references = _get_document_references(documents, top_nodes)
170 prompt_template = PromptTemplate(template=_build_chat_prompt(output_language))
171 refine_template = PromptTemplate(template=_build_refine_prompt(output_language))
172 response_synthesizer = get_response_synthesizer(
173 llm=client.llm,
174 prompt_helper=get_rag_prompt_helper(
175 chunk_size=config.llm_embedding_chunk_size,
176 context_size=config.llm_context_size,
177 ),
178 text_qa_template=prompt_template,
179 refine_template=refine_template,
180 streaming=True,
181 )
182 query_engine = RetrieverQueryEngine.from_args(
183 retriever=retriever,
184 llm=client.llm,
185 response_synthesizer=response_synthesizer,
186 streaming=True,
187 )
189 logger.debug("Document chat query: %s", query_str)
190 # Release the pooled DB connection for the slow streaming LLM response
191 # so it is not pinned for the whole stream; see paperless_ai.db and
192 # #12976.
193 with db_connection_released():
194 response_stream = query_engine.query(query_str)
195 for chunk in response_stream.response_gen:
196 yield chunk
197 sys.stdout.flush()
199 if references:
200 yield _format_chat_metadata_trailer(references)