200 lines
6.5 KiB
Python
200 lines
6.5 KiB
Python
"""
|
|
增强的状态管理模块
|
|
|
|
使用 Pydantic 提供类型安全和验证
|
|
"""
|
|
|
|
from typing import Any, Dict, List, Optional, Literal
|
|
from pydantic import BaseModel, Field, field_validator
|
|
from langchain_core.messages import BaseMessage, HumanMessage, AIMessage, SystemMessage
|
|
|
|
|
|
class StateContext(BaseModel):
|
|
"""状态上下文 - 存储工作流执行过程中的数据"""
|
|
|
|
original_input: Optional[str] = Field(None, description="用户原始输入")
|
|
normalized_input: Optional[str] = Field(None, description="规范化后的输入")
|
|
intent: Optional[str] = Field(None, description="识别的意图")
|
|
table_match: Optional[Dict[str, Any]] = Field(None, description="表名匹配结果")
|
|
final_sql: Optional[str] = Field(None, description="生成的 SQL")
|
|
sr_api_result: Optional[Any] = Field(None, description="API 执行结果")
|
|
|
|
class Config:
|
|
extra = "allow"
|
|
|
|
def get(self, key: str, default: Any = None) -> Any:
|
|
"""获取上下文值"""
|
|
return getattr(self, key, default)
|
|
|
|
def set(self, key: str, value: Any) -> None:
|
|
"""设置上下文值"""
|
|
setattr(self, key, value)
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
"""转换为字典"""
|
|
return self.model_dump(exclude_none=True)
|
|
|
|
|
|
class AgentState(BaseModel):
|
|
"""
|
|
Agent 工作流状态定义
|
|
|
|
使用 Pydantic 提供类型安全和验证
|
|
"""
|
|
|
|
messages: List[BaseMessage] = Field(default_factory=list, description="消息历史")
|
|
current_step: str = Field(default="start", description="当前步骤")
|
|
context: StateContext = Field(default_factory=StateContext, description="上下文数据")
|
|
|
|
model_config = {
|
|
"arbitrary_types_allowed": True,
|
|
"extra": "forbid",
|
|
}
|
|
|
|
@field_validator("messages", mode="before")
|
|
@classmethod
|
|
def validate_messages(cls, v):
|
|
"""验证并转换消息列表"""
|
|
if not isinstance(v, list):
|
|
return []
|
|
|
|
result = []
|
|
for msg in v:
|
|
if isinstance(msg, BaseMessage):
|
|
result.append(msg)
|
|
elif isinstance(msg, dict):
|
|
msg_type = msg.get("type", "human")
|
|
content = msg.get("content", "")
|
|
if msg_type == "human":
|
|
result.append(HumanMessage(content=content))
|
|
elif msg_type == "ai":
|
|
result.append(AIMessage(content=content))
|
|
elif msg_type == "system":
|
|
result.append(SystemMessage(content=content))
|
|
return result
|
|
|
|
def add_message(self, message: BaseMessage) -> "AgentState":
|
|
"""添加消息并返回新状态"""
|
|
return AgentState(
|
|
messages=[*self.messages, message],
|
|
current_step=self.current_step,
|
|
context=self.context,
|
|
)
|
|
|
|
def add_human_message(self, content: str) -> "AgentState":
|
|
"""添加用户消息"""
|
|
return self.add_message(HumanMessage(content=content))
|
|
|
|
def add_ai_message(self, content: str) -> "AgentState":
|
|
"""添加 AI 消息"""
|
|
return self.add_message(AIMessage(content=content))
|
|
|
|
def update_step(self, step: str) -> "AgentState":
|
|
"""更新当前步骤"""
|
|
return AgentState(
|
|
messages=self.messages,
|
|
current_step=step,
|
|
context=self.context,
|
|
)
|
|
|
|
def update_context(self, **kwargs) -> "AgentState":
|
|
"""更新上下文"""
|
|
new_context = self.context.model_copy()
|
|
for key, value in kwargs.items():
|
|
new_context.set(key, value)
|
|
return AgentState(
|
|
messages=self.messages,
|
|
current_step=self.current_step,
|
|
context=new_context,
|
|
)
|
|
|
|
def get_last_message(self) -> Optional[BaseMessage]:
|
|
"""获取最后一条消息"""
|
|
return self.messages[-1] if self.messages else None
|
|
|
|
def get_context(self, key: str, default: Any = None) -> Any:
|
|
"""获取上下文值"""
|
|
return self.context.get(key, default)
|
|
|
|
def to_legacy_format(self) -> Dict[str, Any]:
|
|
"""
|
|
转换为旧格式(兼容现有代码)
|
|
|
|
现有代码期望 state 是一个可修改的对象,
|
|
这个方法返回一个兼容的字典格式
|
|
"""
|
|
return {
|
|
"messages": self.messages,
|
|
"current_step": self.current_step,
|
|
"context": self.context.to_dict(),
|
|
}
|
|
|
|
@classmethod
|
|
def from_legacy_format(cls, data: Dict[str, Any]) -> "AgentState":
|
|
"""从旧格式创建"""
|
|
context_data = data.get("context", {})
|
|
if isinstance(context_data, StateContext):
|
|
context = context_data
|
|
else:
|
|
context = StateContext(**context_data) if context_data else StateContext()
|
|
|
|
return cls(
|
|
messages=data.get("messages", []),
|
|
current_step=data.get("current_step", "start"),
|
|
context=context,
|
|
)
|
|
|
|
|
|
class MutableAgentState:
|
|
"""
|
|
可变的 Agent 状态包装器
|
|
|
|
用于兼容现有代码中直接修改 state 的模式
|
|
"""
|
|
|
|
def __init__(self, state: Optional[AgentState] = None):
|
|
self._state = state or AgentState()
|
|
self._context_overrides: Dict[str, Any] = {}
|
|
|
|
@property
|
|
def messages(self) -> List[BaseMessage]:
|
|
return self._state.messages
|
|
|
|
@messages.setter
|
|
def messages(self, value: List[BaseMessage]):
|
|
self._state = AgentState(
|
|
messages=value,
|
|
current_step=self._state.current_step,
|
|
context=self._state.context,
|
|
)
|
|
|
|
@property
|
|
def current_step(self) -> str:
|
|
return self._state.current_step
|
|
|
|
@current_step.setter
|
|
def current_step(self, value: str):
|
|
self._state = AgentState(
|
|
messages=self._state.messages,
|
|
current_step=value,
|
|
context=self._state.context,
|
|
)
|
|
|
|
@property
|
|
def context(self) -> Dict[str, Any]:
|
|
"""返回可修改的上下文字典"""
|
|
result = self._state.context.to_dict()
|
|
result.update(self._context_overrides)
|
|
return result
|
|
|
|
def to_immutable(self) -> AgentState:
|
|
"""转换为不可变状态"""
|
|
context = self._state.context.model_copy()
|
|
for key, value in self._context_overrides.items():
|
|
context.set(key, value)
|
|
return AgentState(
|
|
messages=self._state.messages,
|
|
current_step=self._state.current_step,
|
|
context=context,
|
|
)
|