Files
more_dots/core/state.py
T
2026-03-11 23:40:39 +08:00

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,
)