appQQQ commited on
Commit
cdc883e
·
verified ·
1 Parent(s): 2193e90

chore: upload app/agents/nodes.py

Browse files
Files changed (1) hide show
  1. app/agents/nodes.py +480 -0
app/agents/nodes.py ADDED
@@ -0,0 +1,480 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """LangGraph 节点实现.
2
+
3
+ 每个 node 接收 AgentState, 返回部分更新的 dict.
4
+ 节点间通过 state 自动传递, 不直接耦合.
5
+
6
+ 设计要点:
7
+ - 节点只做一件事, 容易测试
8
+ - LLM 调用统一走工厂, 走 LLM cache
9
+ - 异常不抛, 写入 state['error'], 让图走 fallback 边
10
+ """
11
+ from __future__ import annotations
12
+
13
+ import json
14
+ import logging
15
+ import re
16
+ import time
17
+ from typing import Any
18
+
19
+ from langchain_core.messages import AIMessage, HumanMessage, SystemMessage
20
+
21
+ from app.agents.prompts import (
22
+ ANSWER_PROMPT,
23
+ CRAG_EVAL_PROMPT,
24
+ MULTI_STEP_PROMPT,
25
+ QUERY_REWRITE_PROMPT,
26
+ ROUTE_PROMPT,
27
+ )
28
+ from app.agents.state import AgentState
29
+ from app.agents.tools import TOOL_SCHEMAS, execute_tool
30
+ from app.config import settings
31
+ from app.llm.base import LLMMessage
32
+ from app.llm.factory import get_llm
33
+ from app.services.embedding import get_embedder
34
+ from app.services.llm_cache import CachedAnswer, lookup as cache_lookup, store as cache_store
35
+ from app.services.reranker import get_reranker_service
36
+ from app.services.vector_store import RetrievalHit, hybrid_query
37
+
38
+ logger = logging.getLogger(__name__)
39
+
40
+
41
+ # ========== 工具: 取最近用户消息文本 ==========
42
+ def _last_user_query(state: AgentState) -> str:
43
+ for m in reversed(list(state.get("messages") or [])):
44
+ if hasattr(m, "type") and m.type == "human":
45
+ return m.content if isinstance(m.content, str) else str(m.content)
46
+ if isinstance(m, dict) and m.get("role") == "user":
47
+ return m.get("content", "")
48
+ return ""
49
+
50
+
51
+ def _safe_json(text: str) -> dict | None:
52
+ """尽量从 LLM 输出中抠 JSON. 失败返回 None."""
53
+ if not text:
54
+ return None
55
+ # 尝试直接 parse
56
+ try:
57
+ return json.loads(text)
58
+ except json.JSONDecodeError:
59
+ pass
60
+ # 抠 ```json ... ```
61
+ m = re.search(r"```(?:json)?\s*(\{.*?\}|\[.*?\])\s*```", text, re.DOTALL)
62
+ if m:
63
+ try:
64
+ return json.loads(m.group(1))
65
+ except json.JSONDecodeError:
66
+ pass
67
+ # 抠第一个 { ... }
68
+ m = re.search(r"\{.*\}", text, re.DOTALL)
69
+ if m:
70
+ try:
71
+ return json.loads(m.group(0))
72
+ except json.JSONDecodeError:
73
+ pass
74
+ return None
75
+
76
+
77
+ # ========== Node 1: route ==========
78
+ async def route_node(state: AgentState) -> dict[str, Any]:
79
+ """判断 query 走向: direct / retrieve / multi_step."""
80
+ started = time.time()
81
+ query = _last_user_query(state)
82
+ if not query:
83
+ return {"route_decision": "retrieve"}
84
+
85
+ # 启发式快路 (避免每次都 LLM 调用)
86
+ ql = query.strip().lower()
87
+ if len(ql) <= 12 and any(g in ql for g in (
88
+ "你好", "您好", "hi", "hello", "hey", "你是谁", "what's up", "how are you",
89
+ "thanks", "thank you", "谢谢", "再见", "bye",
90
+ )):
91
+ return {"route_decision": "direct", "query_rewritten": query}
92
+
93
+ try:
94
+ llm = get_llm()
95
+ resp = await llm.chat(
96
+ messages=[
97
+ LLMMessage(role="system", content=ROUTE_PROMPT),
98
+ LLMMessage(role="user", content=query),
99
+ ],
100
+ temperature=0.0,
101
+ max_tokens=80,
102
+ )
103
+ data = _safe_json(resp.content or "")
104
+ decision = data.get("route", "retrieve") if data else "retrieve"
105
+ if decision not in ("direct", "retrieve", "multi_step"):
106
+ decision = "retrieve"
107
+ logger.debug("route: %s (%sms) reason=%s", decision, int((time.time() - started) * 1000),
108
+ (data or {}).get("reason", ""))
109
+ return {"route_decision": decision, "query_rewritten": query}
110
+ except Exception as e: # noqa: BLE001
111
+ logger.warning("route_node failed: %s, default to retrieve", e)
112
+ return {"route_decision": "retrieve", "query_rewritten": query}
113
+
114
+
115
+ # ========== Node 2: query_rewrite ==========
116
+ async def query_rewrite_node(state: AgentState) -> dict[str, Any]:
117
+ """对 query 改写, 提升召回. multi_step 时拆子问题."""
118
+ started = time.time()
119
+ query = state.get("query_rewritten") or _last_user_query(state)
120
+ if state.get("route_decision") == "direct":
121
+ return {"query_rewritten": query, "plan": []}
122
+
123
+ try:
124
+ llm = get_llm()
125
+ if state.get("route_decision") == "multi_step":
126
+ resp = await llm.chat(
127
+ messages=[
128
+ LLMMessage(role="system", content=MULTI_STEP_PROMPT),
129
+ LLMMessage(role="user", content=query),
130
+ ],
131
+ temperature=0.2,
132
+ max_tokens=200,
133
+ )
134
+ data = _safe_json(resp.content or "")
135
+ steps = (data or {}).get("steps", [query])
136
+ if not isinstance(steps, list) or not steps:
137
+ steps = [query]
138
+ return {"query_rewritten": steps[0], "plan": steps}
139
+
140
+ resp = await llm.chat(
141
+ messages=[
142
+ LLMMessage(role="system", content=QUERY_REWRITE_PROMPT.format(query=query)),
143
+ ],
144
+ temperature=0.3,
145
+ max_tokens=200,
146
+ )
147
+ data = _safe_json(resp.content or "")
148
+ rewrites = (data or {}).get("rewrites", [])
149
+ if not isinstance(rewrites, list) or not rewrites:
150
+ rewrites = [query]
151
+ # 拼接为最终检索串
152
+ merged = " | ".join([query] + list(rewrites[:2]))
153
+ logger.debug("query_rewrite: %d variants, %dms", len(rewrites), int((time.time() - started) * 1000))
154
+ return {"query_rewritten": merged, "plan": []}
155
+ except Exception as e: # noqa: BLE001
156
+ logger.warning("query_rewrite failed: %s", e)
157
+ return {"query_rewritten": query, "plan": []}
158
+
159
+
160
+ # ========== Node 3: retrieve ==========
161
+ async def retrieve_node(state: AgentState) -> dict[str, Any]:
162
+ """混合检索 top-K."""
163
+ started = time.time()
164
+ query = state.get("query_rewritten") or _last_user_query(state)
165
+ if not query:
166
+ return {"retrieved": [], "retrieved_doc_ids": []}
167
+
168
+ embedder = get_embedder()
169
+ out = await embedder.encode_query(query)
170
+ dense = out["dense"]
171
+ # dense shape: (1, 1024) or (1024,) depending on encode return
172
+ if dense.ndim == 2:
173
+ dense_vec = dense[0]
174
+ else:
175
+ dense_vec = dense
176
+ sparse = out.get("sparse", [{}])[0] if out.get("sparse") else None
177
+ colbert = out.get("colbert", [None])[0] if out.get("colbert") else None
178
+
179
+ # multi_step: 每步独立检索再合并
180
+ plan = state.get("plan") or []
181
+ all_hits: list[RetrievalHit] = []
182
+ seen: set[str] = set()
183
+ queries_to_run = plan if plan else [query]
184
+ for q in queries_to_run:
185
+ if q == query and all_hits:
186
+ continue # 主 query 已跑过
187
+ if q != query:
188
+ sub_out = await embedder.encode_query(q)
189
+ sub_dense = sub_out["dense"][0] if sub_out["dense"].ndim == 2 else sub_out["dense"]
190
+ sub_sparse = sub_out.get("sparse", [{}])[0] if sub_out.get("sparse") else None
191
+ sub_colbert = sub_out.get("colbert", [None])[0] if sub_out.get("colbert") else None
192
+ hits = hybrid_query(
193
+ query_emb=sub_dense,
194
+ query_sparse=sub_sparse,
195
+ query_colbert_emb=sub_colbert,
196
+ k=settings.rerank_top_n * 4,
197
+ )
198
+ else:
199
+ hits = hybrid_query(
200
+ query_emb=dense_vec,
201
+ query_sparse=sparse,
202
+ query_colbert_emb=colbert,
203
+ k=settings.rerank_top_n * 4,
204
+ )
205
+ for h in hits:
206
+ if h.chunk_id not in seen:
207
+ seen.add(h.chunk_id)
208
+ all_hits.append(h)
209
+
210
+ # 按 score 截前 N
211
+ all_hits.sort(key=lambda h: h.score, reverse=True)
212
+ all_hits = all_hits[: settings.rerank_top_n * 4]
213
+
214
+ doc_ids = list({h.doc_id for h in all_hits if h.doc_id})
215
+ logger.debug("retrieve: %d hits, %d docs, %dms",
216
+ len(all_hits), len(doc_ids), int((time.time() - started) * 1000))
217
+ return {"retrieved": all_hits, "retrieved_doc_ids": doc_ids}
218
+
219
+
220
+ # ========== Node 4: rerank ==========
221
+ async def rerank_node(state: AgentState) -> dict[str, Any]:
222
+ """BGE-reranker 精排 top-N + 产出引用."""
223
+ started = time.time()
224
+ hits = state.get("retrieved") or []
225
+ query = state.get("query_rewritten") or _last_user_query(state)
226
+ if not hits:
227
+ return {"reranked": [], "citations": [], "relevance_score": 0.0, "relevance_verdict": "irrelevant"}
228
+
229
+ reranker = get_reranker_service()
230
+ reranked = await reranker.rerank(query, hits, top_n=settings.rerank_top_n)
231
+
232
+ # 构造引用 (前 5 个, 按 rerank 分数)
233
+ citations: list[dict[str, Any]] = []
234
+ for i, h in enumerate(reranked):
235
+ doc = _doc_meta_brief(h.doc_id)
236
+ citations.append({
237
+ "doc_id": h.doc_id,
238
+ "filename": doc.get("filename", "未知"),
239
+ "page": h.page_no,
240
+ "heading": h.heading,
241
+ "snippet": (h.text or "")[:240],
242
+ "score": round(h.rerank_score, 4),
243
+ "rank": i + 1,
244
+ })
245
+
246
+ top_score = reranked[0].rerank_score if reranked else 0.0
247
+ if top_score >= settings.crag_relevance_threshold:
248
+ verdict = "relevant"
249
+ elif top_score < 0.3:
250
+ verdict = "irrelevant"
251
+ else:
252
+ verdict = "ambiguous"
253
+
254
+ logger.debug("rerank: top=%.3f verdict=%s %dms", top_score, verdict, int((time.time() - started) * 1000))
255
+ return {
256
+ "reranked": reranked,
257
+ "citations": citations,
258
+ "relevance_score": top_score,
259
+ "relevance_verdict": verdict,
260
+ }
261
+
262
+
263
+ def _doc_meta_brief(doc_id: str) -> dict[str, Any]:
264
+ try:
265
+ from app.models import db
266
+ d = db.doc_get(doc_id)
267
+ return d or {}
268
+ except Exception: # noqa: BLE001
269
+ return {}
270
+
271
+
272
+ # ========== Node 5: answer (流式 LLM 调用) ==========
273
+ async def answer_node_stream(
274
+ state: AgentState,
275
+ on_token: Any = None, # async callable(content: str) -> None
276
+ on_citation: Any = None,
277
+ on_thinking: Any = None,
278
+ ) -> dict[str, Any]:
279
+ """生成最终答案. 通过 on_token 回调逐 token 推送.
280
+
281
+ 流程:
282
+ 1. 拼装 context (从 reranked hits)
283
+ 2. 查 LLM 缓存
284
+ 3. 命中: 回放 tokens
285
+ 4. 未命中: 调 LLM 流式 + 缓存结果
286
+ """
287
+ query = state.get("query_rewritten") or _last_user_query(state)
288
+ reranked = state.get("reranked") or []
289
+ locale = state.get("locale", "zh")
290
+
291
+ # 拼 context
292
+ if reranked:
293
+ ctx_lines: list[str] = []
294
+ for i, h in enumerate(reranked, 1):
295
+ tag = f"[{i}]"
296
+ prefix_bits = []
297
+ if h.heading:
298
+ prefix_bits.append(f"章节: {h.heading}")
299
+ if h.page_no:
300
+ prefix_bits.append(f"页码: {h.page_no}")
301
+ if h.context_prefix:
302
+ prefix_bits.append(f"上下文: {h.context_prefix}")
303
+ meta = " | ".join(prefix_bits)
304
+ ctx_lines.append(f"{tag} {('('+meta+')') if meta else ''}\n{h.text}")
305
+ context = "\n\n".join(ctx_lines)
306
+ else:
307
+ if on_thinking:
308
+ await on_thinking("未在知识库中找到相关文档, 直接基于通用知识回答。")
309
+ context = "(无相关文档)"
310
+
311
+ prompt = ANSWER_PROMPT.format(context=context, query=query, LOCALE=locale)
312
+ system_msg = "你是私人智能客服, 回答需基于 context 引用, 用对应 locale 回答。"
313
+
314
+ # 缓存 key
315
+ top_doc_ids = [c["doc_id"] for c in state.get("citations", [])]
316
+ cached = cache_lookup(query, top_doc_ids, 0.7)
317
+
318
+ started = time.time()
319
+ if cached is not None:
320
+ # 回放
321
+ if on_thinking:
322
+ await on_thinking("(cache hit, 跳过 LLM)")
323
+ for tok in cached.tokens:
324
+ if on_token:
325
+ await on_token(tok)
326
+ return {
327
+ "final_answer": cached.content,
328
+ "messages": [AIMessage(content=cached.content)],
329
+ "elapsed_ms": int((time.time() - started) * 1000),
330
+ }
331
+
332
+ # 实际 LLM 流式
333
+ llm = get_llm()
334
+ collected: list[str] = []
335
+ full_text = ""
336
+ try:
337
+ async for chunk in llm.stream_chat(
338
+ messages=[
339
+ LLMMessage(role="system", content=system_msg),
340
+ LLMMessage(role="user", content=prompt),
341
+ ],
342
+ temperature=0.7,
343
+ max_tokens=1200,
344
+ ):
345
+ if chunk.content:
346
+ collected.append(chunk.content)
347
+ full_text += chunk.content
348
+ if on_token:
349
+ await on_token(chunk.content)
350
+ except Exception as e: # noqa: BLE001
351
+ logger.exception("answer_node_stream failed")
352
+ err_msg = f"抱歉, 生成答案时出错: {e}"
353
+ if on_token:
354
+ await on_token(err_msg)
355
+ return {
356
+ "final_answer": err_msg,
357
+ "messages": [AIMessage(content=err_msg)],
358
+ "error": str(e),
359
+ }
360
+
361
+ # 缓存结果
362
+ cache_store(query, top_doc_ids, 0.7, CachedAnswer(
363
+ content=full_text,
364
+ citations=state.get("citations", []),
365
+ tool_calls=[],
366
+ tokens=collected,
367
+ ))
368
+
369
+ # 推引用 (在 answer 末尾)
370
+ if on_citation and state.get("citations"):
371
+ for c in state["citations"]:
372
+ await on_citation(c)
373
+
374
+ return {
375
+ "final_answer": full_text,
376
+ "messages": [AIMessage(content=full_text)],
377
+ "elapsed_ms": int((time.time() - started) * 1000),
378
+ }
379
+
380
+
381
+ # ========== Node 6: evaluate (CRAG) ==========
382
+ async def evaluate_node(state: AgentState) -> dict[str, Any]:
383
+ """CRAG 自校正判定.
384
+
385
+ 阶段 1 (必走, 极快): 用 rerank top-1 score 做硬阈值判断
386
+ 阶段 2 (仅模糊区间): LLM judge 二次判定, 决定是否回 retrieve
387
+ """
388
+ iteration = state.get("iteration", 0) + 1
389
+ score = state.get("relevance_score", 0.0)
390
+ verdict = state.get("relevance_verdict", "ambiguous")
391
+ max_iter = state.get("max_iterations", settings.crag_max_iterations)
392
+
393
+ # 阶段 1: rerank 分数硬阈值
394
+ if verdict == "relevant":
395
+ return {
396
+ "iteration": iteration,
397
+ "needs_more_retrieval": False,
398
+ "crag_finished": True,
399
+ }
400
+ if verdict == "irrelevant":
401
+ # 直接告知用户, 不再循环
402
+ return {
403
+ "iteration": iteration,
404
+ "needs_more_retrieval": False,
405
+ "crag_finished": True,
406
+ }
407
+
408
+ # 阶段 2: 模糊区间, 调 LLM judge (用便宜的 judge model)
409
+ if iteration >= max_iter:
410
+ # 超过上限, 收口
411
+ return {
412
+ "iteration": iteration,
413
+ "needs_more_retrieval": False,
414
+ "crag_finished": True,
415
+ }
416
+
417
+ try:
418
+ llm = get_llm() # 用同 model (个人项目成本可接受)
419
+ query = state.get("query_rewritten") or _last_user_query(state)
420
+ reranked = state.get("reranked") or []
421
+ docs_summary = "\n".join(
422
+ f"[{i+1}] {h.heading or '无标题'}: {(h.text or '')[:120]}"
423
+ for i, h in enumerate(reranked[:5])
424
+ )
425
+ resp = await llm.chat(
426
+ messages=[
427
+ LLMMessage(role="system", content=CRAG_EVAL_PROMPT.format(
428
+ query=query, n=len(reranked[:5]), docs_summary=docs_summary,
429
+ )),
430
+ ],
431
+ temperature=0.0,
432
+ max_tokens=120,
433
+ )
434
+ data = _safe_json(resp.content or "")
435
+ v = (data or {}).get("verdict", "sufficient")
436
+ needs = v == "insufficient" and iteration < max_iter
437
+ return {
438
+ "iteration": iteration,
439
+ "needs_more_retrieval": needs,
440
+ "crag_finished": not needs,
441
+ }
442
+ except Exception as e: # noqa: BLE001
443
+ logger.warning("evaluate_node LLM judge failed: %s", e)
444
+ return {
445
+ "iteration": iteration,
446
+ "needs_more_retrieval": False,
447
+ "crag_finished": True,
448
+ }
449
+
450
+
451
+ # ========== Node 7: tool_executor ==========
452
+ async def tool_executor_node(state: AgentState) -> dict[str, Any]:
453
+ """执行 LLM 在 answer 阶段请求的工具调用."""
454
+ msgs = list(state.get("messages") or [])
455
+ last_ai = next((m for m in reversed(msgs)
456
+ if hasattr(m, "type") and m.type == "ai"), None)
457
+ tool_calls = getattr(last_ai, "tool_calls", None) or []
458
+
459
+ if not tool_calls:
460
+ return {"tool_results": []}
461
+
462
+ results: list[dict[str, Any]] = []
463
+ for tc in tool_calls:
464
+ name = tc.get("name", "")
465
+ args = tc.get("args", {})
466
+ if isinstance(args, str):
467
+ try:
468
+ args = json.loads(args)
469
+ except json.JSONDecodeError:
470
+ args = {}
471
+ try:
472
+ out = await execute_tool(name, args)
473
+ except Exception as e: # noqa: BLE001
474
+ out = f"工具执行异常: {e}"
475
+ results.append({"name": name, "args": args, "output": out})
476
+
477
+ return {
478
+ "tool_results": results,
479
+ "tool_calls": [{"name": r["name"], "args": r["args"]} for r in results],
480
+ }