ai-chatbot / app /agents /nodes.py
appQQQ's picture
chore: upload app/agents/nodes.py
cdc883e verified
Raw
History Blame
16.9 kB
"""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],
}