x
This commit is contained in:
+199
@@ -0,0 +1,199 @@
|
||||
"""
|
||||
增强的状态管理模块
|
||||
|
||||
使用 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,
|
||||
)
|
||||
Reference in New Issue
Block a user