diff --git a/app.py b/app.py index f28840a..7e87171 100644 --- a/app.py +++ b/app.py @@ -31,6 +31,7 @@ from workflow_registry import WORKFLOW_CONFIG, VALID_ROUTE_FLAGS from main_agent import extract_conversation_title from workflows.history_manager import filter_image_urls, preprocess_from_request from config import set_request_user_config +from utils.source_citations import merge_rag_source_citations from checkpointer_config import ( CheckpointerManager, checkpointer_manager @@ -473,7 +474,7 @@ async def stream_main_agent_execution( agent_task = asyncio.create_task(run_agent()) execution_complete_received = False - rag_result = [] + rag_result = {} graph_result = [] while True: try: @@ -493,16 +494,7 @@ async def stream_main_agent_execution( if title == "知识库搜索工具" and result and isinstance(result, dict): print("===============================", result) rag_result_raw = result.get("sourceCitation", {}) - rag_result = {} - for doc_name, chunks in rag_result_raw.items(): - if isinstance(chunks, list): - rag_result[doc_name] = [ - {k: v for k, v in item.items() if k != "text"} - if isinstance(item, dict) else item - for item in chunks - ] - else: - rag_result[doc_name] = chunks + merge_rag_source_citations(rag_result, rag_result_raw) elif title == "图谱检索工具" and result and isinstance(result, dict): print("===============================", result) graph_result = result.get("xxxx.graph", []) diff --git a/utils/source_citations.py b/utils/source_citations.py new file mode 100644 index 0000000..c072972 --- /dev/null +++ b/utils/source_citations.py @@ -0,0 +1,71 @@ +"""Utilities for accumulating RAG source citations across multiple searches.""" + +import json +from typing import Any, Dict + + +def _citation_item_key(item: Any) -> str: + """Build a stable identity for one RAG citation chunk.""" + if not isinstance(item, dict): + return json.dumps(item, ensure_ascii=False, sort_keys=True, default=str) + + chunk_id = item.get("id") + if chunk_id not in (None, ""): + return f"id:{chunk_id}" + + location = { + "file_id": item.get("file_id"), + "page_idx": item.get("page_idx"), + "positions": item.get("positions"), + } + if any(value not in (None, "", []) for value in location.values()): + return "location:" + json.dumps( + location, ensure_ascii=False, sort_keys=True, default=str + ) + + # index/score may differ between rewritten queries and do not identify a source. + fallback = { + key: value + for key, value in item.items() + if key not in {"index", "score", "text"} + } + return "fallback:" + json.dumps( + fallback, ensure_ascii=False, sort_keys=True, default=str + ) + + +def merge_rag_source_citations( + accumulated: Dict[str, Any], incoming: Any +) -> Dict[str, Any]: + """ + Append a RAG call's citations without losing earlier calls or duplicating chunks. + + Citation text is intentionally removed here because the stream response has never + exposed it; the full text remains available inside the workflow for answer generation. + """ + if not isinstance(incoming, dict) or not incoming: + return accumulated + + for doc_name, chunks in incoming.items(): + if isinstance(chunks, list): + target = accumulated.setdefault(doc_name, []) + if not isinstance(target, list): + # Be defensive about malformed/legacy data while preserving both values. + target = [target] + accumulated[doc_name] = target + + seen = {_citation_item_key(item) for item in target} + for item in chunks: + sanitized_item = ( + {key: value for key, value in item.items() if key != "text"} + if isinstance(item, dict) + else item + ) + item_key = _citation_item_key(sanitized_item) + if item_key not in seen: + target.append(sanitized_item) + seen.add(item_key) + elif doc_name not in accumulated: + accumulated[doc_name] = chunks + + return accumulated