""" 增强的状态管理模块 使用 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, )