mirror of
https://github.com/paperless-ngx/paperless-ngx.git
synced 2026-10-04 15:20:30 +00:00
* Fix: skip vector store document id filter for unrestricted chat users ChatStreamingView built an IN filter from every permitted document id for the "chat over all documents" case, which exceeds the vector store's SQLite bound-parameter safety limit on installs with more than ~32700 documents, silently returning no context. For a user who can see every document (an active superuser), that filter never narrows anything, so skip it and let the retriever search the whole index instead. * Minor improvements from a Claude review * When a user is unrestricted chatting, still exclude trashed documents using a 'NOT IN' SQL statement. Wire that up where we need it * Update src/paperless_ai/chat.py Co-authored-by: shamoon <4887959+shamoon@users.noreply.github.com>
201 lines
6.8 KiB
Python
201 lines
6.8 KiB
Python
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
|
|
from paperless_ai.db import db_connection_released
|
|
from paperless_ai.indexing import document_id_filters
|
|
from paperless_ai.indexing import exclude_document_ids_filter
|
|
from paperless_ai.indexing import get_rag_prompt_helper
|
|
from paperless_ai.indexing import load_or_build_index
|
|
from paperless_ai.indexing import read_store
|
|
from paperless_ai.prompts.context import ChatQaPromptContext
|
|
from paperless_ai.prompts.context import ChatRefinePromptContext
|
|
from paperless_ai.prompts.render import render_prompt
|
|
|
|
logger = logging.getLogger("paperless_ai.chat")
|
|
|
|
CHAT_METADATA_DELIMITER = "\n\n__PAPERLESS_CHAT_METADATA__"
|
|
CHAT_ERROR_MESSAGE = "Sorry, something went wrong while generating a response."
|
|
CHAT_NO_CONTENT_MESSAGE = "Sorry, I couldn't find any content to answer your question."
|
|
MAX_CHAT_REFERENCES = 3
|
|
CHAT_RETRIEVER_TOP_K = 5
|
|
|
|
|
|
def _build_chat_prompt(output_language: str | None) -> str:
|
|
return render_prompt(ChatQaPromptContext(output_language=output_language))
|
|
|
|
|
|
def _build_refine_prompt(output_language: str | None) -> str:
|
|
return render_prompt(
|
|
ChatRefinePromptContext(output_language=output_language),
|
|
)
|
|
|
|
|
|
def _build_document_reference(
|
|
document: Document,
|
|
title: str | None = None,
|
|
) -> dict[str, int | str]:
|
|
return {
|
|
"id": document.pk,
|
|
"title": title or document.title or document.filename,
|
|
}
|
|
|
|
|
|
def _get_document_references(
|
|
documents: QuerySet[Document],
|
|
top_nodes: list,
|
|
) -> list[dict[str, int | str]]:
|
|
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()
|
|
|
|
for node in top_nodes:
|
|
try:
|
|
document_id = int(node.metadata["document_id"])
|
|
except (KeyError, TypeError, ValueError): # pragma: no cover
|
|
continue
|
|
|
|
if document_id in seen_document_ids or document_id not in allowed_documents:
|
|
continue
|
|
|
|
seen_document_ids.add(document_id)
|
|
document = allowed_documents[document_id]
|
|
references.append(
|
|
_build_document_reference(document, node.metadata.get("title")),
|
|
)
|
|
|
|
if len(references) >= MAX_CHAT_REFERENCES: # pragma: no cover
|
|
break
|
|
|
|
return references
|
|
|
|
|
|
def _format_chat_metadata_trailer(references: list[dict[str, int | str]]) -> str:
|
|
return (
|
|
f"{CHAT_METADATA_DELIMITER}"
|
|
f"{json.dumps({'references': references}, separators=(',', ':'))}"
|
|
)
|
|
|
|
|
|
def stream_chat_with_documents(
|
|
query_str: str,
|
|
documents: QuerySet[Document],
|
|
*,
|
|
unrestricted: bool = False,
|
|
output_language: str | None = None,
|
|
):
|
|
try:
|
|
yield from _stream_chat_with_documents(
|
|
query_str,
|
|
documents,
|
|
unrestricted=unrestricted,
|
|
output_language=output_language,
|
|
)
|
|
except Exception as e:
|
|
logger.exception("Failed to stream document chat response: %s", e)
|
|
yield CHAT_ERROR_MESSAGE
|
|
|
|
|
|
def _stream_chat_with_documents(
|
|
query_str: str,
|
|
documents: QuerySet[Document],
|
|
*,
|
|
unrestricted: bool = False,
|
|
output_language: str | None = None,
|
|
):
|
|
if not documents.exists():
|
|
yield CHAT_NO_CONTENT_MESSAGE
|
|
return
|
|
|
|
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
|
|
from llama_index.core.retrievers import VectorIndexRetriever
|
|
|
|
config = AIConfig()
|
|
if unrestricted:
|
|
# Exclude trashed ids (usually few) instead of an IN filter over the
|
|
# full permitted set, which risks the vector store's bound parameter
|
|
# limit (_MAX_IN_VALUES) on large installs. Trashed documents stay
|
|
# indexed until permanent deletion (delete_document_from_llm_index
|
|
# hangs off post_delete, not trash), so must be excluded explicitly.
|
|
trashed_ids = Document.deleted_objects.values_list("pk", flat=True)
|
|
filters = exclude_document_ids_filter(str(pk) for pk in trashed_ids)
|
|
else:
|
|
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
|
|
# must stay open (and the swap must not run) until the stream finishes.
|
|
with read_store() as store:
|
|
index = load_or_build_index(config, store)
|
|
retriever = VectorIndexRetriever(
|
|
index=index,
|
|
similarity_top_k=CHAT_RETRIEVER_TOP_K,
|
|
filters=filters,
|
|
)
|
|
|
|
# Slow query-embedding + vector search; no Django ORM access happens
|
|
# during it, so release the pooled DB connection for its duration. See
|
|
# #12976.
|
|
with db_connection_released():
|
|
top_nodes = retriever.retrieve(query_str)
|
|
if not top_nodes:
|
|
logger.warning("No nodes found for the given documents.")
|
|
yield CHAT_NO_CONTENT_MESSAGE
|
|
return
|
|
|
|
client = AIClient()
|
|
|
|
references = _get_document_references(documents, top_nodes)
|
|
|
|
prompt_template = PromptTemplate(template=_build_chat_prompt(output_language))
|
|
refine_template = PromptTemplate(template=_build_refine_prompt(output_language))
|
|
response_synthesizer = get_response_synthesizer(
|
|
llm=client.llm,
|
|
prompt_helper=get_rag_prompt_helper(
|
|
chunk_size=config.llm_embedding_chunk_size,
|
|
context_size=config.llm_context_size,
|
|
),
|
|
text_qa_template=prompt_template,
|
|
refine_template=refine_template,
|
|
streaming=True,
|
|
)
|
|
query_engine = RetrieverQueryEngine.from_args(
|
|
retriever=retriever,
|
|
llm=client.llm,
|
|
response_synthesizer=response_synthesizer,
|
|
streaming=True,
|
|
)
|
|
|
|
logger.debug("Document chat query: %s", query_str)
|
|
# Release the pooled DB connection for the slow streaming LLM response
|
|
# so it is not pinned for the whole stream; see paperless_ai.db and
|
|
# #12976.
|
|
with db_connection_released():
|
|
response_stream = query_engine.query(query_str)
|
|
for chunk in response_stream.response_gen:
|
|
yield chunk
|
|
sys.stdout.flush()
|
|
|
|
if references:
|
|
yield _format_chat_metadata_trailer(references)
|