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

1import json 

2import logging 

3import sys 

4 

5from django.db.models import QuerySet 

6 

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 

19 

20logger = logging.getLogger("paperless_ai.chat") 

21 

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 

27 

28 

29def _build_chat_prompt(output_language: str | None) -> str: 

30 return render_prompt(ChatQaPromptContext(output_language=output_language)) 

31 

32 

33def _build_refine_prompt(output_language: str | None) -> str: 

34 return render_prompt( 

35 ChatRefinePromptContext(output_language=output_language), 

36 ) 

37 

38 

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 } 

47 

48 

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 

59 

60 if not candidate_ids: 

61 return [] 

62 

63 allowed_documents = {doc.pk: doc for doc in documents.filter(pk__in=candidate_ids)} 

64 

65 references: list[dict[str, int | str]] = [] 

66 seen_document_ids: set[int] = set() 

67 

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 

73 

74 if document_id in seen_document_ids or document_id not in allowed_documents: 

75 continue 

76 

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 ) 

82 

83 if len(references) >= MAX_CHAT_REFERENCES: # pragma: no cover 

84 break 

85 

86 return references 

87 

88 

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 ) 

94 

95 

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 

113 

114 

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 

125 

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 

130 

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 ) 

144 

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 ) 

155 

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 

165 

166 client = AIClient() 

167 

168 references = _get_document_references(documents, top_nodes) 

169 

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 ) 

188 

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() 

198 

199 if references: 

200 yield _format_chat_metadata_trailer(references)