Spaces:
Sleeping
Sleeping
| """LangGraph 节点实现. | |
| 每个 node 接收 AgentState, 返回部分更新的 dict. | |
| 节点间通过 state 自动传递, 不直接耦合. | |
| 设计要点: | |
| - 节点只做一件事, 容易测试 | |
| - LLM 调用统一走工厂, 走 LLM cache | |
| - 异常不抛, 写入 state['error'], 让图走 fallback 边 | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import logging | |
| import re | |
| import time | |
| from typing import Any | |
| from langchain_core.messages import AIMessage, HumanMessage, SystemMessage | |
| from app.agents.prompts import ( | |
| ANSWER_PROMPT, | |
| CRAG_EVAL_PROMPT, | |
| MULTI_STEP_PROMPT, | |
| QUERY_REWRITE_PROMPT, | |
| ROUTE_PROMPT, | |
| ) | |
| from app.agents.state import AgentState | |
| from app.agents.tools import TOOL_SCHEMAS, execute_tool | |
| from app.config import settings | |
| from app.llm.base import LLMMessage | |
| from app.llm.factory import get_llm | |
| from app.services.embedding import get_embedder | |
| from app.services.llm_cache import CachedAnswer, lookup as cache_lookup, store as cache_store | |
| from app.services.reranker import get_reranker_service | |
| from app.services.vector_store import RetrievalHit, hybrid_query | |
| logger = logging.getLogger(__name__) | |
| # ========== 工具: 取最近用户消息文本 ========== | |
| def _last_user_query(state: AgentState) -> str: | |
| for m in reversed(list(state.get("messages") or [])): | |
| if hasattr(m, "type") and m.type == "human": | |
| return m.content if isinstance(m.content, str) else str(m.content) | |
| if isinstance(m, dict) and m.get("role") == "user": | |
| return m.get("content", "") | |
| return "" | |
| def _safe_json(text: str) -> dict | None: | |
| """尽量从 LLM 输出中抠 JSON. 失败返回 None.""" | |
| if not text: | |
| return None | |
| # 尝试直接 parse | |
| try: | |
| return json.loads(text) | |
| except json.JSONDecodeError: | |
| pass | |
| # 抠 ```json ... ``` | |
| m = re.search(r"```(?:json)?\s*(\{.*?\}|\[.*?\])\s*```", text, re.DOTALL) | |
| if m: | |
| try: | |
| return json.loads(m.group(1)) | |
| except json.JSONDecodeError: | |
| pass | |
| # 抠第一个 { ... } | |
| m = re.search(r"\{.*\}", text, re.DOTALL) | |
| if m: | |
| try: | |
| return json.loads(m.group(0)) | |
| except json.JSONDecodeError: | |
| pass | |
| return None | |
| # ========== Node 1: route ========== | |
| async def route_node(state: AgentState) -> dict[str, Any]: | |
| """判断 query 走向: direct / retrieve / multi_step.""" | |
| started = time.time() | |
| query = _last_user_query(state) | |
| if not query: | |
| return {"route_decision": "retrieve"} | |
| # 启发式快路 (避免每次都 LLM 调用) | |
| ql = query.strip().lower() | |
| if len(ql) <= 12 and any(g in ql for g in ( | |
| "你好", "您好", "hi", "hello", "hey", "你是谁", "what's up", "how are you", | |
| "thanks", "thank you", "谢谢", "再见", "bye", | |
| )): | |
| return {"route_decision": "direct", "query_rewritten": query} | |
| try: | |
| llm = get_llm() | |
| resp = await llm.chat( | |
| messages=[ | |
| LLMMessage(role="system", content=ROUTE_PROMPT), | |
| LLMMessage(role="user", content=query), | |
| ], | |
| temperature=0.0, | |
| max_tokens=80, | |
| ) | |
| data = _safe_json(resp.content or "") | |
| decision = data.get("route", "retrieve") if data else "retrieve" | |
| if decision not in ("direct", "retrieve", "multi_step"): | |
| decision = "retrieve" | |
| logger.debug("route: %s (%sms) reason=%s", decision, int((time.time() - started) * 1000), | |
| (data or {}).get("reason", "")) | |
| return {"route_decision": decision, "query_rewritten": query} | |
| except Exception as e: # noqa: BLE001 | |
| logger.warning("route_node failed: %s, default to retrieve", e) | |
| return {"route_decision": "retrieve", "query_rewritten": query} | |
| # ========== Node 2: query_rewrite ========== | |
| async def query_rewrite_node(state: AgentState) -> dict[str, Any]: | |
| """对 query 改写, 提升召回. multi_step 时拆子问题.""" | |
| started = time.time() | |
| query = state.get("query_rewritten") or _last_user_query(state) | |
| if state.get("route_decision") == "direct": | |
| return {"query_rewritten": query, "plan": []} | |
| try: | |
| llm = get_llm() | |
| if state.get("route_decision") == "multi_step": | |
| resp = await llm.chat( | |
| messages=[ | |
| LLMMessage(role="system", content=MULTI_STEP_PROMPT), | |
| LLMMessage(role="user", content=query), | |
| ], | |
| temperature=0.2, | |
| max_tokens=200, | |
| ) | |
| data = _safe_json(resp.content or "") | |
| steps = (data or {}).get("steps", [query]) | |
| if not isinstance(steps, list) or not steps: | |
| steps = [query] | |
| return {"query_rewritten": steps[0], "plan": steps} | |
| resp = await llm.chat( | |
| messages=[ | |
| LLMMessage(role="system", content=QUERY_REWRITE_PROMPT.format(query=query)), | |
| ], | |
| temperature=0.3, | |
| max_tokens=200, | |
| ) | |
| data = _safe_json(resp.content or "") | |
| rewrites = (data or {}).get("rewrites", []) | |
| if not isinstance(rewrites, list) or not rewrites: | |
| rewrites = [query] | |
| # 拼接为最终检索串 | |
| merged = " | ".join([query] + list(rewrites[:2])) | |
| logger.debug("query_rewrite: %d variants, %dms", len(rewrites), int((time.time() - started) * 1000)) | |
| return {"query_rewritten": merged, "plan": []} | |
| except Exception as e: # noqa: BLE001 | |
| logger.warning("query_rewrite failed: %s", e) | |
| return {"query_rewritten": query, "plan": []} | |
| # ========== Node 3: retrieve ========== | |
| async def retrieve_node(state: AgentState) -> dict[str, Any]: | |
| """混合检索 top-K.""" | |
| started = time.time() | |
| query = state.get("query_rewritten") or _last_user_query(state) | |
| if not query: | |
| return {"retrieved": [], "retrieved_doc_ids": []} | |
| embedder = get_embedder() | |
| out = await embedder.encode_query(query) | |
| dense = out["dense"] | |
| # dense shape: (1, 1024) or (1024,) depending on encode return | |
| if dense.ndim == 2: | |
| dense_vec = dense[0] | |
| else: | |
| dense_vec = dense | |
| sparse = out.get("sparse", [{}])[0] if out.get("sparse") else None | |
| colbert = out.get("colbert", [None])[0] if out.get("colbert") else None | |
| # multi_step: 每步独立检索再合并 | |
| plan = state.get("plan") or [] | |
| all_hits: list[RetrievalHit] = [] | |
| seen: set[str] = set() | |
| queries_to_run = plan if plan else [query] | |
| for q in queries_to_run: | |
| if q == query and all_hits: | |
| continue # 主 query 已跑过 | |
| if q != query: | |
| sub_out = await embedder.encode_query(q) | |
| sub_dense = sub_out["dense"][0] if sub_out["dense"].ndim == 2 else sub_out["dense"] | |
| sub_sparse = sub_out.get("sparse", [{}])[0] if sub_out.get("sparse") else None | |
| sub_colbert = sub_out.get("colbert", [None])[0] if sub_out.get("colbert") else None | |
| hits = hybrid_query( | |
| query_emb=sub_dense, | |
| query_sparse=sub_sparse, | |
| query_colbert_emb=sub_colbert, | |
| k=settings.rerank_top_n * 4, | |
| ) | |
| else: | |
| hits = hybrid_query( | |
| query_emb=dense_vec, | |
| query_sparse=sparse, | |
| query_colbert_emb=colbert, | |
| k=settings.rerank_top_n * 4, | |
| ) | |
| for h in hits: | |
| if h.chunk_id not in seen: | |
| seen.add(h.chunk_id) | |
| all_hits.append(h) | |
| # 按 score 截前 N | |
| all_hits.sort(key=lambda h: h.score, reverse=True) | |
| all_hits = all_hits[: settings.rerank_top_n * 4] | |
| doc_ids = list({h.doc_id for h in all_hits if h.doc_id}) | |
| logger.debug("retrieve: %d hits, %d docs, %dms", | |
| len(all_hits), len(doc_ids), int((time.time() - started) * 1000)) | |
| return {"retrieved": all_hits, "retrieved_doc_ids": doc_ids} | |
| # ========== Node 4: rerank ========== | |
| async def rerank_node(state: AgentState) -> dict[str, Any]: | |
| """BGE-reranker 精排 top-N + 产出引用.""" | |
| started = time.time() | |
| hits = state.get("retrieved") or [] | |
| query = state.get("query_rewritten") or _last_user_query(state) | |
| if not hits: | |
| return {"reranked": [], "citations": [], "relevance_score": 0.0, "relevance_verdict": "irrelevant"} | |
| reranker = get_reranker_service() | |
| reranked = await reranker.rerank(query, hits, top_n=settings.rerank_top_n) | |
| # 构造引用 (前 5 个, 按 rerank 分数) | |
| citations: list[dict[str, Any]] = [] | |
| for i, h in enumerate(reranked): | |
| doc = _doc_meta_brief(h.doc_id) | |
| citations.append({ | |
| "doc_id": h.doc_id, | |
| "filename": doc.get("filename", "未知"), | |
| "page": h.page_no, | |
| "heading": h.heading, | |
| "snippet": (h.text or "")[:240], | |
| "score": round(h.rerank_score, 4), | |
| "rank": i + 1, | |
| }) | |
| top_score = reranked[0].rerank_score if reranked else 0.0 | |
| if top_score >= settings.crag_relevance_threshold: | |
| verdict = "relevant" | |
| elif top_score < 0.3: | |
| verdict = "irrelevant" | |
| else: | |
| verdict = "ambiguous" | |
| logger.debug("rerank: top=%.3f verdict=%s %dms", top_score, verdict, int((time.time() - started) * 1000)) | |
| return { | |
| "reranked": reranked, | |
| "citations": citations, | |
| "relevance_score": top_score, | |
| "relevance_verdict": verdict, | |
| } | |
| def _doc_meta_brief(doc_id: str) -> dict[str, Any]: | |
| try: | |
| from app.models import db | |
| d = db.doc_get(doc_id) | |
| return d or {} | |
| except Exception: # noqa: BLE001 | |
| return {} | |
| # ========== Node 5: answer (流式 LLM 调用) ========== | |
| async def answer_node_stream( | |
| state: AgentState, | |
| on_token: Any = None, # async callable(content: str) -> None | |
| on_citation: Any = None, | |
| on_thinking: Any = None, | |
| ) -> dict[str, Any]: | |
| """生成最终答案. 通过 on_token 回调逐 token 推送. | |
| 流程: | |
| 1. 拼装 context (从 reranked hits) | |
| 2. 查 LLM 缓存 | |
| 3. 命中: 回放 tokens | |
| 4. 未命中: 调 LLM 流式 + 缓存结果 | |
| """ | |
| query = state.get("query_rewritten") or _last_user_query(state) | |
| reranked = state.get("reranked") or [] | |
| locale = state.get("locale", "zh") | |
| # 拼 context | |
| if reranked: | |
| ctx_lines: list[str] = [] | |
| for i, h in enumerate(reranked, 1): | |
| tag = f"[{i}]" | |
| prefix_bits = [] | |
| if h.heading: | |
| prefix_bits.append(f"章节: {h.heading}") | |
| if h.page_no: | |
| prefix_bits.append(f"页码: {h.page_no}") | |
| if h.context_prefix: | |
| prefix_bits.append(f"上下文: {h.context_prefix}") | |
| meta = " | ".join(prefix_bits) | |
| ctx_lines.append(f"{tag} {('('+meta+')') if meta else ''}\n{h.text}") | |
| context = "\n\n".join(ctx_lines) | |
| else: | |
| if on_thinking: | |
| await on_thinking("未在知识库中找到相关文档, 直接基于通用知识回答。") | |
| context = "(无相关文档)" | |
| prompt = ANSWER_PROMPT.format(context=context, query=query, LOCALE=locale) | |
| system_msg = "你是私人智能客服, 回答需基于 context 引用, 用对应 locale 回答。" | |
| # 缓存 key | |
| top_doc_ids = [c["doc_id"] for c in state.get("citations", [])] | |
| cached = cache_lookup(query, top_doc_ids, 0.7) | |
| started = time.time() | |
| if cached is not None: | |
| # 回放 | |
| if on_thinking: | |
| await on_thinking("(cache hit, 跳过 LLM)") | |
| for tok in cached.tokens: | |
| if on_token: | |
| await on_token(tok) | |
| return { | |
| "final_answer": cached.content, | |
| "messages": [AIMessage(content=cached.content)], | |
| "elapsed_ms": int((time.time() - started) * 1000), | |
| } | |
| # 实际 LLM 流式 | |
| llm = get_llm() | |
| collected: list[str] = [] | |
| full_text = "" | |
| try: | |
| async for chunk in llm.stream_chat( | |
| messages=[ | |
| LLMMessage(role="system", content=system_msg), | |
| LLMMessage(role="user", content=prompt), | |
| ], | |
| temperature=0.7, | |
| max_tokens=1200, | |
| ): | |
| if chunk.content: | |
| collected.append(chunk.content) | |
| full_text += chunk.content | |
| if on_token: | |
| await on_token(chunk.content) | |
| except Exception as e: # noqa: BLE001 | |
| logger.exception("answer_node_stream failed") | |
| err_msg = f"抱歉, 生成答案时出错: {e}" | |
| if on_token: | |
| await on_token(err_msg) | |
| return { | |
| "final_answer": err_msg, | |
| "messages": [AIMessage(content=err_msg)], | |
| "error": str(e), | |
| } | |
| # 缓存结果 | |
| cache_store(query, top_doc_ids, 0.7, CachedAnswer( | |
| content=full_text, | |
| citations=state.get("citations", []), | |
| tool_calls=[], | |
| tokens=collected, | |
| )) | |
| # 推引用 (在 answer 末尾) | |
| if on_citation and state.get("citations"): | |
| for c in state["citations"]: | |
| await on_citation(c) | |
| return { | |
| "final_answer": full_text, | |
| "messages": [AIMessage(content=full_text)], | |
| "elapsed_ms": int((time.time() - started) * 1000), | |
| } | |
| # ========== Node 6: evaluate (CRAG) ========== | |
| async def evaluate_node(state: AgentState) -> dict[str, Any]: | |
| """CRAG 自校正判定. | |
| 阶段 1 (必走, 极快): 用 rerank top-1 score 做硬阈值判断 | |
| 阶段 2 (仅模糊区间): LLM judge 二次判定, 决定是否回 retrieve | |
| """ | |
| iteration = state.get("iteration", 0) + 1 | |
| score = state.get("relevance_score", 0.0) | |
| verdict = state.get("relevance_verdict", "ambiguous") | |
| max_iter = state.get("max_iterations", settings.crag_max_iterations) | |
| # 阶段 1: rerank 分数硬阈值 | |
| if verdict == "relevant": | |
| return { | |
| "iteration": iteration, | |
| "needs_more_retrieval": False, | |
| "crag_finished": True, | |
| } | |
| if verdict == "irrelevant": | |
| # 直接告知用户, 不再循环 | |
| return { | |
| "iteration": iteration, | |
| "needs_more_retrieval": False, | |
| "crag_finished": True, | |
| } | |
| # 阶段 2: 模糊区间, 调 LLM judge (用便宜的 judge model) | |
| if iteration >= max_iter: | |
| # 超过上限, 收口 | |
| return { | |
| "iteration": iteration, | |
| "needs_more_retrieval": False, | |
| "crag_finished": True, | |
| } | |
| try: | |
| llm = get_llm() # 用同 model (个人项目成本可接受) | |
| query = state.get("query_rewritten") or _last_user_query(state) | |
| reranked = state.get("reranked") or [] | |
| docs_summary = "\n".join( | |
| f"[{i+1}] {h.heading or '无标题'}: {(h.text or '')[:120]}" | |
| for i, h in enumerate(reranked[:5]) | |
| ) | |
| resp = await llm.chat( | |
| messages=[ | |
| LLMMessage(role="system", content=CRAG_EVAL_PROMPT.format( | |
| query=query, n=len(reranked[:5]), docs_summary=docs_summary, | |
| )), | |
| ], | |
| temperature=0.0, | |
| max_tokens=120, | |
| ) | |
| data = _safe_json(resp.content or "") | |
| v = (data or {}).get("verdict", "sufficient") | |
| needs = v == "insufficient" and iteration < max_iter | |
| return { | |
| "iteration": iteration, | |
| "needs_more_retrieval": needs, | |
| "crag_finished": not needs, | |
| } | |
| except Exception as e: # noqa: BLE001 | |
| logger.warning("evaluate_node LLM judge failed: %s", e) | |
| return { | |
| "iteration": iteration, | |
| "needs_more_retrieval": False, | |
| "crag_finished": True, | |
| } | |
| # ========== Node 7: tool_executor ========== | |
| async def tool_executor_node(state: AgentState) -> dict[str, Any]: | |
| """执行 LLM 在 answer 阶段请求的工具调用.""" | |
| msgs = list(state.get("messages") or []) | |
| last_ai = next((m for m in reversed(msgs) | |
| if hasattr(m, "type") and m.type == "ai"), None) | |
| tool_calls = getattr(last_ai, "tool_calls", None) or [] | |
| if not tool_calls: | |
| return {"tool_results": []} | |
| results: list[dict[str, Any]] = [] | |
| for tc in tool_calls: | |
| name = tc.get("name", "") | |
| args = tc.get("args", {}) | |
| if isinstance(args, str): | |
| try: | |
| args = json.loads(args) | |
| except json.JSONDecodeError: | |
| args = {} | |
| try: | |
| out = await execute_tool(name, args) | |
| except Exception as e: # noqa: BLE001 | |
| out = f"工具执行异常: {e}" | |
| results.append({"name": name, "args": args, "output": out}) | |
| return { | |
| "tool_results": results, | |
| "tool_calls": [{"name": r["name"], "args": r["args"]} for r in results], | |
| } | |