修复多次RAG检索溯源覆盖

This commit is contained in:
ChenRui 2026-07-22 17:35:19 +08:00
parent 6c470d55a5
commit 263cb9f723
2 changed files with 74 additions and 11 deletions

14
app.py
View File

@ -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", [])

71
utils/source_citations.py Normal file
View File

@ -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