appQQQ commited on
Commit
2b21030
·
verified ·
1 Parent(s): cdc883e

chore: upload app/agents/state.py

Browse files
Files changed (1) hide show
  1. app/agents/state.py +115 -0
app/agents/state.py ADDED
@@ -0,0 +1,115 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """AgentState - LangGraph 跨节点状态.
2
+
3
+ 使用 TypedDict 而不是 Pydantic, 因为 LangGraph 内置 add_messages reducer
4
+ 需要 Annotated[Sequence[BaseMessage], add_messages] 这样的类型签名.
5
+ """
6
+ from __future__ import annotations
7
+
8
+ from typing import Annotated, Any, Literal, Sequence, TypedDict
9
+
10
+ from langchain_core.documents import Document
11
+ from langchain_core.messages import BaseMessage
12
+ from langgraph.graph.message import add_messages
13
+
14
+ from app.services.vector_store import RetrievalHit
15
+
16
+
17
+ RouteDecision = Literal["direct", "retrieve", "multi_step"]
18
+
19
+
20
+ class AgentState(TypedDict, total=False):
21
+ """LangGraph 跨节点状态. total=False 让所有字段可选, 各 node 自行填充."""
22
+
23
+ # ===== 基础 =====
24
+ messages: Annotated[Sequence[BaseMessage], add_messages]
25
+ session_id: str
26
+ user_id: str
27
+ locale: str # "zh" | "en"
28
+
29
+ # ===== 路由 =====
30
+ route_decision: RouteDecision
31
+ query_rewritten: str # 改写后的查询
32
+ plan: list[str] # 多步拆解的子任务
33
+
34
+ # ===== 检索 =====
35
+ retrieved: list[RetrievalHit]
36
+ reranked: list[RetrievalHit]
37
+ retrieved_doc_ids: list[str] # 命中 doc 列表 (用于 SSE retrieval 事件)
38
+
39
+ # ===== 引用 =====
40
+ citations: list[dict[str, Any]] # [{doc_id, page, snippet, score, source}]
41
+
42
+ # ===== 工具 =====
43
+ tool_calls: list[dict[str, Any]] # 工具调用历史
44
+ tool_results: list[dict[str, Any]] # 工具结果摘要
45
+
46
+ # ===== CRAG 自校正 =====
47
+ iteration: int
48
+ max_iterations: int
49
+ relevance_score: float # top-1 rerank score, 0-1
50
+ relevance_verdict: Literal["relevant", "ambiguous", "irrelevant"]
51
+ needs_more_retrieval: bool
52
+ crag_finished: bool
53
+
54
+ # ===== 元数据 =====
55
+ elapsed_ms: int
56
+ final_answer: str
57
+ error: str | None
58
+
59
+
60
+ # ========== State helpers ==========
61
+ def empty_state_for(session_id: str, user_id: str = "default", locale: str = "zh") -> AgentState:
62
+ from app.config import settings
63
+
64
+ return AgentState(
65
+ messages=[],
66
+ session_id=session_id,
67
+ user_id=user_id,
68
+ locale=locale,
69
+ route_decision="retrieve",
70
+ query_rewritten="",
71
+ plan=[],
72
+ retrieved=[],
73
+ reranked=[],
74
+ retrieved_doc_ids=[],
75
+ citations=[],
76
+ tool_calls=[],
77
+ tool_results=[],
78
+ iteration=0,
79
+ max_iterations=settings.crag_max_iterations,
80
+ relevance_score=0.0,
81
+ relevance_verdict="ambiguous",
82
+ needs_more_retrieval=False,
83
+ crag_finished=False,
84
+ elapsed_ms=0,
85
+ final_answer="",
86
+ error=None,
87
+ )
88
+
89
+
90
+ def messages_to_lc(messages: list[dict[str, Any]] | list[BaseMessage]) -> list[BaseMessage]:
91
+ """业务侧 dict 消息列表 (来自 SQLite) 转为 LangChain BaseMessage 列表."""
92
+ from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage
93
+
94
+ out: list[BaseMessage] = []
95
+ for m in messages:
96
+ if hasattr(m, "type"):
97
+ out.append(m) # 已经是 BaseMessage
98
+ continue
99
+ role = m.get("role", "user")
100
+ content = m.get("content", "")
101
+ if role == "system":
102
+ out.append(SystemMessage(content=content))
103
+ elif role == "user":
104
+ out.append(HumanMessage(content=content))
105
+ elif role == "assistant":
106
+ extra = {}
107
+ if m.get("tool_calls"):
108
+ extra["tool_calls"] = m["tool_calls"]
109
+ out.append(AIMessage(content=content, **extra))
110
+ elif role == "tool":
111
+ out.append(ToolMessage(
112
+ content=content,
113
+ tool_call_id=m.get("tool_call_id", ""),
114
+ ))
115
+ return out