""" 百科问答工作流。 核心链路借鉴 WeKnora: 1. 意图判断默认偏向检索 2. 基于历史重写检索 query,并生成多个检索变体 3. RAG / Graph 并行召回 4. 多源 evidence 去重、融合排序、rerank 5. 只基于 evidence 生成可追溯回答 """ import asyncio import json import os import random import re import sys from typing import Any, Dict, List, Optional, TypedDict from langgraph.graph import END, START, StateGraph sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from modelsAPI.model_api import OpenaiAPI from prompts import BAIKE_PROMPTS from tools.function_tool import graph_rag_search from tools.rag_tools import rag_search from utils.function_tracker import get_all_callbacks, track_function_calls from utils.image_utils import extract_image_paths from workflows.history_manager import DEFAULT_HISTORY_MAX_MESSAGES, init_history from workflows.workflow_utils import dedupe_rag_results, stream_generate_and_postprocess BAIKE_RAG_TOP_K = 8 BAIKE_GRAPH_TOP_K = 240 BAIKE_RESULT_LIMIT = 8 BAIKE_EVIDENCE_LIMIT = 10 class QAState(TypedDict): extracted_text: str combined_query: str retrieval_query: str retrieval_queries: List[str] route_flag: str history_message: str history_messages: List[Dict[str, Any]] need_retrieval: Optional[bool] rag_search_result: Optional[List[Dict[str, Any]]] graph_search_result: Optional[str] evidence_items: Optional[List[Dict[str, Any]]] response: str actions: List[Dict[str, Any]] suggestedReplies: List[Dict[str, Any]] error_message: str def get_baike_history_context(state: QAState) -> Dict[str, Any]: return init_history(state.get("history_message", "") or "") async def judge_need_retrieval(state: QAState) -> Dict[str, Any]: original_query = (state.get("combined_query") or state.get("extracted_text") or "").strip() history_messages = state.get("history_messages", []) or [] prompt = BAIKE_PROMPTS["judge_need_retrieval"].format( history_messages=history_messages if history_messages else "无", original_query=original_query, ) substantive_query = _looks_like_substantive_query(original_query) need_retrieval = substantive_query try: result = await OpenaiAPI.open_api_chat_without_thinking(prompt, model=None) if not substantive_query and result and result.strip().lower() == "false": need_retrieval = False except Exception as exc: print(f"判断是否需要检索失败: {exc}") print(f"是否需要检索: {need_retrieval}") return {"need_retrieval": need_retrieval} def _looks_like_substantive_query(query: str) -> bool: text = re.sub(r"\s+", "", (query or "").strip()) if not text: return False if re.fullmatch(r"(你好|您好|hi|hello|在吗|谢谢|多谢|好的|好|嗯|哦|确认|取消|不用了)[。!!?\s]*", text, re.I): return False if len(text) >= 4: return True return bool(re.search(r"[??]|[A-Za-z0-9]{2,}", text)) async def rewrite_query_with_history(state: QAState) -> Dict[str, Any]: original_query = (state.get("combined_query") or state.get("extracted_text") or "").strip() history_messages = state.get("history_messages", []) or [] trimmed_history = history_messages[-DEFAULT_HISTORY_MAX_MESSAGES:] if history_messages else [] history_str = _json_dumps(trimmed_history) last_user_query = "" for msg in reversed(trimmed_history): if msg.get("role") == "user": last_user_query = str(msg.get("content") or "").strip() break if not original_query: return {"retrieval_query": "", "retrieval_queries": []} prompt = BAIKE_PROMPTS["rewrite_query_with_history"].format( last_user_query=last_user_query or "无", history_str=history_str, original_query=original_query, ) rewritten_query = original_query query_variants: List[str] = [] try: raw = await OpenaiAPI.open_api_chat_without_thinking(prompt, model=None, json_output=True) data = _parse_json_object(raw) if data: candidate = str(data.get("rewrite_query") or "").strip() if candidate: rewritten_query = candidate[:180] for q in data.get("search_queries") or []: if isinstance(q, str) and q.strip(): query_variants.append(q.strip()[:180]) except Exception as exc: print(f"query 重写失败: {exc}") query_variants = _build_search_queries(rewritten_query, original_query, query_variants) _emit_event("✨正在分析中,请稍等...", "正在生成检索策略...") print("原始检索query:", original_query) print("重写后检索query:", rewritten_query) print("检索query变体:", query_variants) return {"retrieval_query": rewritten_query, "retrieval_queries": query_variants} @track_function_calls async def search_rag_and_graph(state: QAState) -> Dict[str, Any]: kb_id = "3" if state.get("route_flag") == "rules" else None query = (state.get("retrieval_query") or state.get("combined_query") or state.get("extracted_text") or "").strip() extracted_text = (state.get("extracted_text") or "").strip() query_variants = _build_search_queries(query, extracted_text, state.get("retrieval_queries") or []) try: rag_task = _search_rag_variants(query_variants[:3], kb_id=kb_id) graph_task = graph_rag_search.ainvoke({"query": query, "top_k": BAIKE_GRAPH_TOP_K}) rag_results, graph_results = await asyncio.gather(rag_task, graph_task) evidence_items = _build_evidence_items( rag_results=rag_results, graph_results=graph_results if isinstance(graph_results, str) else "", query=query, ) evidence_items = await _rerank_evidence(query, evidence_items) return { "rag_search_result": rag_results, "graph_search_result": graph_results if isinstance(graph_results, str) else "", "evidence_items": evidence_items, } except Exception as exc: return { "rag_search_result": [], "graph_search_result": "", "evidence_items": [], "error_message": f"检索失败: {exc}", } async def generate_suggested_replies_baike(answer: str, question: str) -> List[Dict[str, Any]]: prompt = BAIKE_PROMPTS["generate_suggested_replies"].format(question=question, answer=answer[:500]) try: result = await OpenaiAPI.open_api_chat_without_thinking(prompt, model=None, json_output=True) json_match = re.search(r"\[.*\]", result or "", re.DOTALL) if json_match: replies = json.loads(json_match.group(0)) formatted = [ {"title": str(x["title"]), "content": str(x["content"])} for x in replies[:2] if isinstance(x, dict) and x.get("title") and x.get("content") ] if len(formatted) == 2: return formatted except Exception: pass return [ {"title": "查看相关依据", "content": "这个回答对应的资料依据有哪些?"}, {"title": "继续展开说明", "content": "请把这个问题再展开说明一下。"}, ] @track_function_calls async def generate_answer(state: QAState) -> Dict[str, Any]: question = state.get("combined_query") or state.get("extracted_text") or "" evidence_items = state.get("evidence_items") or [] need_retrieval = state.get("need_retrieval", True) history_messages = state.get("history_messages", []) or [] try: source_text = "" original_image_paths: List[str] = [] if need_retrieval and evidence_items: source_text = _format_evidence_context(evidence_items) original_image_paths = extract_image_paths(source_text) source_counts = _count_sources(evidence_items) _emit_event( f"✨已融合 {len(evidence_items)} 条证据", f"RAG {source_counts.get('rag', 0)} 条,图谱 {source_counts.get('graph', 0)} 条", ) domain_desc = ( "舰船维修相关的法规、规范、行业标准及技术文件" if state.get("route_flag") == "rules" else "工业设备的结构原理、维护规程、故障诊断与维修操作" ) if need_retrieval and source_text.strip(): system_prompt = BAIKE_PROMPTS["generate_answer_with_retrieval_system"].format(domain_desc=domain_desc) user_prompt = BAIKE_PROMPTS["generate_answer_with_retrieval_user"].format( extracted_text=question, rag_count=len(evidence_items), all_source_text=source_text, ) else: system_prompt = BAIKE_PROMPTS["generate_answer_without_retrieval_system"].format(domain_desc=domain_desc) user_prompt = BAIKE_PROMPTS["generate_answer_without_retrieval_user"].format(extracted_text=question) messages = [] if history_messages: messages.extend(history_messages) messages.append({"role": "user", "content": user_prompt}) _emit_event(random.choice(["🔍 证据整理中", "🧠 正在生成回答"]), "正在基于检索资料生成回答...") answer = await stream_generate_and_postprocess( system_prompt=system_prompt, messages=messages, original_image_paths=original_image_paths, temperature=0.2, fallback="抱歉,无法生成回答。", ) suggested_replies = await generate_suggested_replies_baike(answer, question) return {"response": answer, "actions": [], "result_tag": "qa", "suggestedReplies": suggested_replies} except Exception as exc: return {"response": f"生成回答失败:{exc}", "actions": [], "suggestedReplies": []} async def _search_rag_variants(queries: List[str], kb_id: Optional[str] = None) -> List[Dict[str, Any]]: tasks = [rag_search(query=q, kb_id=kb_id, top_k=BAIKE_RAG_TOP_K) for q in queries if q.strip()] results = await asyncio.gather(*tasks, return_exceptions=True) if tasks else [] items: List[Dict[str, Any]] = [] for result in results: if isinstance(result, Exception) or not result.get("success"): continue for resource, rows in (result.get("sourceCitation") or {}).items(): for row in rows: item = dict(row) item.setdefault("resource", resource) items.append(item) items = sorted(items, key=lambda x: float(x.get("score") or 0), reverse=True) return dedupe_rag_results(items, limit=BAIKE_RESULT_LIMIT) def _build_evidence_items( rag_results: List[Dict[str, Any]], graph_results: str, query: str, ) -> List[Dict[str, Any]]: candidates: List[Dict[str, Any]] = [] for item in rag_results or []: candidates.append({ "source_type": "rag", "title": item.get("filename") or item.get("resource") or "RAG 检索结果", "filename": item.get("filename") or item.get("resource") or "", "text": item.get("text") or "", "score": float(item.get("score") or 0), "file_id": str(item.get("file_id") or ""), "slice_id": str(item.get("id") or item.get("slice_id") or ""), "metadata": item, }) graph_text = str(graph_results or "").strip() if graph_text and "图谱检索完成,但未返回有效数据。" not in graph_text and graph_text != "查询无结果": candidates.append({ "source_type": "graph", "title": "图谱检索结果", "filename": "", "text": graph_text, "score": 0.45, "file_id": "", "slice_id": "", "metadata": {}, }) deduped: Dict[str, Dict[str, Any]] = {} for item in candidates: text = _clean_text(item.get("text") or "") if not text: continue item["text"] = text[:5000] key = item.get("slice_id") or _content_key(text) local_score = _local_evidence_score(query, item) item["local_score"] = local_score existing = deduped.get(key) if not existing or local_score > float(existing.get("local_score") or 0): deduped[key] = item return sorted(deduped.values(), key=lambda x: float(x.get("local_score") or 0), reverse=True) async def _rerank_evidence(query: str, evidence_items: List[Dict[str, Any]]) -> List[Dict[str, Any]]: if not query.strip() or not evidence_items: return evidence_items[:BAIKE_EVIDENCE_LIMIT] selected = evidence_items[:24] documents = [ "\n".join([ str(item.get("title") or ""), str(item.get("filename") or ""), str(item.get("text") or ""), ])[:4000] for item in selected ] rows = await OpenaiAPI.rerank(query=query, documents=documents, top_k=BAIKE_EVIDENCE_LIMIT) if not rows: return evidence_items[:BAIKE_EVIDENCE_LIMIT] ranked: List[Dict[str, Any]] = [] seen = set() for row in rows: idx = int(row.get("index") or 0) if idx < 0 or idx >= len(selected) or idx in seen: continue seen.add(idx) item = dict(selected[idx]) item["rerank_score"] = float(row.get("relevance_score") or 0.0) ranked.append(item) for idx, item in enumerate(selected): if idx not in seen and len(ranked) < BAIKE_EVIDENCE_LIMIT: ranked.append(item) return ranked[:BAIKE_EVIDENCE_LIMIT] def _format_evidence_context(items: List[Dict[str, Any]]) -> str: parts = [] label_map = {"rag": "RAG", "graph": "Graph"} for idx, item in enumerate(items[:BAIKE_EVIDENCE_LIMIT], start=1): source = label_map.get(item.get("source_type"), item.get("source_type") or "unknown") score = item.get("rerank_score", item.get("local_score", item.get("score", ""))) meta_lines = [ f"[Evidence {idx}] source={source} score={score}", f"Title: {item.get('title') or ''}", ] if item.get("filename"): meta_lines.append(f"File: {item.get('filename')}") if item.get("file_id"): meta_lines.append(f"File ID: {item.get('file_id')}") if item.get("slice_id"): meta_lines.append(f"Slice ID: {item.get('slice_id')}") meta_lines.append("Content:") meta_lines.append(str(item.get("text") or "")) parts.append("\n".join(meta_lines)) return "\n\n---\n\n".join(parts) def _local_evidence_score(query: str, item: Dict[str, Any]) -> float: base = {"rag": 1.0, "graph": 0.65}.get(item.get("source_type"), 0.8) raw = float(item.get("score") or 0) score = base + min(max(raw, 0.0), 1.0) text = " ".join([ str(item.get("title") or ""), str(item.get("filename") or ""), str(item.get("text") or "")[:1500], ]).lower() title = str(item.get("title") or "").lower() for term in _extract_search_terms(query): t = term.lower() if not t: continue if t in title: score += 0.7 elif t in text: score += 0.25 return score def _build_search_queries(main_query: str, original_query: str = "", extra_queries: Optional[List[str]] = None) -> List[str]: candidates: List[str] = [] for text in [main_query, original_query, *(extra_queries or [])]: text = _strip_query_noise(text) if not text: continue candidates.append(text) regex_terms = _extract_search_terms(text) if regex_terms: candidates.append(" ".join(regex_terms[:4])) candidates.extend(regex_terms[:4]) return _dedupe_strings(candidates)[:10] def _extract_search_terms(query: str) -> List[str]: stop = {"什么", "如何", "怎么", "哪些", "主要", "内容", "相关", "信息", "一下", "这个", "那个", "是否", "有没有"} terms = [] for token in re.findall(r"[\u4e00-\u9fff]{2,12}|[A-Za-z0-9][A-Za-z0-9_\-]{1,30}", query or ""): token = token.strip("::,,。.;;、()()[]【】") if len(token) < 2 or token in stop or re.fullmatch(r"\d+", token): continue terms.append(token) return _dedupe_strings(terms) def _strip_query_noise(text: str) -> str: text = str(text or "").strip().strip('"“”') if "用户问题:" in text: text = text.split("用户问题:", 1)[1].strip() return re.sub(r"\s+", " ", text) def _parse_json_object(raw: str) -> Dict[str, Any]: if not raw: return {} try: return json.loads(raw) except Exception: match = re.search(r"\{.*\}", raw, re.DOTALL) if match: try: return json.loads(match.group(0)) except Exception: return {} return {} def _json_dumps(value: Any) -> str: try: return json.dumps(value, ensure_ascii=False) except Exception: return str(value) def _content_key(text: str) -> str: return re.sub(r"\s+", "", text or "")[:220] def _clean_text(text: str) -> str: return re.sub(r"\n{3,}", "\n\n", str(text or "")).strip() def _dedupe_strings(values: List[str]) -> List[str]: out = [] seen = set() for value in values: value = str(value or "").strip() if not value or value in seen: continue seen.add(value) out.append(value) return out def _count_sources(items: List[Dict[str, Any]]) -> Dict[str, int]: counts: Dict[str, int] = {} for item in items: key = str(item.get("source_type") or "unknown") counts[key] = counts.get(key, 0) + 1 return counts def _emit_event(title: str, details: str) -> None: event = {"type": "function_execution", "title": title, "details": details} for callback in get_all_callbacks(): try: callback(event) except Exception: pass def route_after_judge(state: QAState) -> str: return "rewrite_query_with_history" if state.get("need_retrieval", True) else "generate_answer" def create_baike_workflow(): workflow = StateGraph(QAState) workflow.add_node("get_baike_history_context", get_baike_history_context) workflow.add_node("judge_need_retrieval", judge_need_retrieval) workflow.add_node("rewrite_query_with_history", rewrite_query_with_history) workflow.add_node("search_rag_and_graph", search_rag_and_graph) workflow.add_node("generate_answer", generate_answer) workflow.add_edge(START, "get_baike_history_context") workflow.add_edge("get_baike_history_context", "judge_need_retrieval") workflow.add_conditional_edges( "judge_need_retrieval", route_after_judge, {"rewrite_query_with_history": "rewrite_query_with_history", "generate_answer": "generate_answer"}, ) workflow.add_edge("rewrite_query_with_history", "search_rag_and_graph") workflow.add_edge("search_rag_and_graph", "generate_answer") workflow.add_edge("generate_answer", END) return workflow.compile() async def run_baike_workflow( extracted_text: str, combined_query: str, history_message: str = "", route_flag: str = "", ) -> Dict[str, Any]: app = create_baike_workflow() initial_state = { "extracted_text": extracted_text, "combined_query": combined_query, "retrieval_query": "", "retrieval_queries": [], "route_flag": route_flag, "history_message": history_message, "history_messages": [], "need_retrieval": None, "rag_search_result": None, "graph_search_result": None, "evidence_items": None, "response": "", "actions": [], "suggestedReplies": [], "error_message": "", } try: final_state = await app.ainvoke(initial_state) if final_state.get("error_message"): return { "response": f"普通问答工作流执行失败: {final_state['error_message']}", "actions": [], "result_tag": "qa", "suggestedReplies": [], } return { "response": final_state.get("response", ""), "actions": final_state.get("actions", []), "result_tag": "qa", "suggestedReplies": final_state.get("suggestedReplies", []), } except Exception as exc: return { "response": f"普通问答工作流执行失败: {exc}", "actions": [], "result_tag": "qa", "suggestedReplies": [], }