Coverage for agentos/state/schema.py: 0%
329 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 21:26 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 21:26 +0800
1"""
2AgentOS v1.14.0 — 结构化 Agent 状态管理系统。
4基因来源: LangGraph Pydantic State Schema + AgentOS Checkpoint。
6核心设计:
7- AgentState: 强类型的全局 Agent 运行时状态,Pydantic v2 驱动
8- 支持 JSON Schema 自动生成、验证、序列化
9- 与 Checkpoint 系统无缝对接
10- 支持状态合并策略(reducer):append/extend/replace/merge
11- 支持子状态派生(SubState),实现层级化状态管理
12"""
14from __future__ import annotations
16import time
17import uuid
18from collections.abc import Callable
19from datetime import UTC, datetime
20from enum import StrEnum
21from typing import (
22 Any,
23 TypeVar,
24)
26try:
27 from pydantic import (
28 BaseModel,
29 ConfigDict,
30 Field,
31 PrivateAttr,
32 )
33except ImportError:
34 raise ImportError(
35 "pydantic>=2.0 is required for agentos.state. " "Install with: pip install pydantic>=2.0"
36 )
38# ── Reducers ────────────────────────────────
41class ReducerStrategy(StrEnum):
42 """状态合并策略。"""
44 REPLACE = "replace" # 直接替换
45 APPEND = "append" # 追加(list -> extend)
46 EXTEND = "extend" # 字典合并
47 MERGE = "merge" # 深度递归合并
48 KEEP_EXISTING = "keep" # 保留旧值
49 CUSTOM = "custom" # 自定义 reducer 函数
52# Type variable for generic state
53S = TypeVar("S", bound=BaseModel)
55# Custom reducer type
56ReducerFn = Callable[[Any, Any], Any]
58# Default reducers
59_DEFAULT_REDUCERS: dict[str, ReducerStrategy] = {}
62def default_reducer(field_name: str, strategy: ReducerStrategy) -> None:
63 """注册字段的默认合并策略。
65 Usage:
66 default_reducer("messages", ReducerStrategy.APPEND)
67 """
68 _DEFAULT_REDUCERS[field_name] = strategy
71def _apply_reducer(old_val: Any, new_val: Any, strategy: ReducerStrategy) -> Any:
72 """应用合并策略。"""
73 if strategy == ReducerStrategy.REPLACE or old_val is None:
74 return new_val
75 if new_val is None:
76 return old_val
77 if strategy == ReducerStrategy.KEEP_EXISTING:
78 return old_val
79 if strategy == ReducerStrategy.APPEND:
80 if isinstance(old_val, list) and isinstance(new_val, list):
81 return old_val + new_val
82 return [old_val, new_val]
83 if strategy == ReducerStrategy.EXTEND:
84 if isinstance(old_val, dict) and isinstance(new_val, dict):
85 return {**old_val, **new_val}
86 return new_val
87 if strategy == ReducerStrategy.MERGE:
88 return _deep_merge(old_val, new_val)
89 return new_val
92def _deep_merge(old: Any, new: Any) -> Any:
93 """递归深度合并两个字典。"""
94 if not isinstance(old, dict) or not isinstance(new, dict):
95 return new
96 result = dict(old)
97 for k, v in new.items():
98 if k in result and isinstance(result[k], dict) and isinstance(v, dict):
99 result[k] = _deep_merge(result[k], v)
100 else:
101 result[k] = v
102 return result
105# ── Field Metadata ──────────────────────────
108class StateFieldInfo(BaseModel):
109 """状态字段元信息。"""
111 reducer: ReducerStrategy = ReducerStrategy.REPLACE
112 custom_reducer: ReducerFn | None = None
113 description: str = ""
114 required: bool = False
115 sensitive: bool = False # 敏感字段,序列化时脱敏
116 checkpointed: bool = True # 是否持久化到 Checkpoint
118 model_config = ConfigDict(arbitrary_types_allowed=True)
121# ── AgentState Core ─────────────────────────
124class BaseAgentState(BaseModel):
125 """Agent 状态的基类。
127 所有 Agent 状态必须继承此类。自动提供:
128 - thread_id / session_id 追踪
129 - step 计数器
130 - 状态快照(snapshot)与恢复(restore)
131 - JSON Schema 生成
132 - Checkpoint 序列化
134 Usage:
135 class MyState(BaseAgentState):
136 messages: list[dict] = Field(default_factory=list)
137 tools_result: dict = Field(default_factory=dict)
138 task_progress: float = 0.0
139 """
141 thread_id: str = Field(
142 default_factory=lambda: f"thread-{uuid.uuid4().hex[:8]}",
143 description="对话线程唯一标识",
144 )
145 messages: list[Any] = Field(default_factory=list, description="对话消息列表")
146 metrics: dict[str, Any] = Field(default_factory=dict, description="运行时指标")
147 step: int = Field(default=0, ge=0, description="当前执行步骤")
148 created_at: str = Field(
149 default_factory=lambda: datetime.now(UTC).isoformat(),
150 description="创建时间",
151 )
152 updated_at: str = Field(
153 default_factory=lambda: datetime.now(UTC).isoformat(),
154 description="最后更新时间",
155 )
156 tags: list[str] = Field(default_factory=list, description="标签")
157 metadata: dict[str, Any] = Field(default_factory=dict, description="自定义元数据")
158 parent_state_id: str | None = Field(default=None, description="父状态 ID(用于分支/回溯)")
160 # 字段级 Reducer 配置
161 _field_reducers: dict[str, ReducerStrategy] = PrivateAttr(default_factory=dict)
162 _field_custom_reducers: dict[str, ReducerFn] = PrivateAttr(default_factory=dict)
163 _sensitive_fields: set[str] = PrivateAttr(default_factory=set)
165 model_config = ConfigDict(
166 extra="allow",
167 validate_assignment=True,
168 json_schema_extra={
169 "title": "AgentState",
170 "description": "AgentOS Structured Agent State",
171 },
172 )
174 def __init__(self, **data):
175 super().__init__(**data)
176 self._field_reducers = {}
177 self._field_custom_reducers = {}
178 self._sensitive_fields = set()
180 # ── Reducer Registration ──────────────────
182 def set_reducer(self, field: str, strategy: ReducerStrategy) -> BaseAgentState:
183 """为字段设置合并策略。"""
184 self._field_reducers[field] = strategy
185 return self
187 def set_custom_reducer(self, field: str, fn: ReducerFn) -> BaseAgentState:
188 """为字段设置自定义合并函数。"""
189 self._field_custom_reducers[field] = fn
190 self._field_reducers[field] = ReducerStrategy.CUSTOM
191 return self
193 def mark_sensitive(self, *fields: str) -> BaseAgentState:
194 """标记敏感字段。"""
195 self._sensitive_fields.update(fields)
196 return self
198 # ── State Mutation ────────────────────────
200 def update_field(
201 self,
202 field: str,
203 value: Any,
204 reducer: ReducerStrategy | None = None,
205 ) -> None:
206 """更新单个字段,自动应用 Reducer。
208 Args:
209 field: 字段名
210 value: 新值
211 reducer: 合并策略,不传则使用注册的 reducer
212 """
213 if field not in self.model_fields and field not in self.model_computed_fields:
214 # Dynamic field — store in metadata
215 old = self.metadata.get(field)
216 strategy = reducer or self._field_reducers.get(
217 field,
218 ReducerStrategy.REPLACE,
219 )
220 self.metadata[field] = _apply_reducer(old, value, strategy)
221 else:
222 old = getattr(self, field, None)
223 strategy = reducer or self._field_reducers.get(
224 field,
225 ReducerStrategy.REPLACE,
226 )
227 new_val = _apply_reducer(old, value, strategy)
228 setattr(self, field, new_val)
230 self.updated_at = datetime.now(UTC).isoformat()
232 def merge(self, other: BaseAgentState | dict) -> BaseAgentState:
233 """合并另一个状态到当前状态。
235 Args:
236 other: 另一个 AgentState 实例或字典
238 Returns:
239 self (in-place merge)
240 """
241 if isinstance(other, dict):
242 other = self.__class__(**other)
244 for field_name in other.model_fields:
245 if field_name in ("thread_id", "created_at"):
246 continue # Immutable fields
248 other_val = getattr(other, field_name, None)
249 if other_val is None:
250 continue
252 self.update_field(field_name, other_val)
254 # Merge metadata
255 if other.metadata:
256 self.metadata = _deep_merge(self.metadata, other.metadata)
257 self.updated_at = datetime.now(UTC).isoformat()
259 # Merge tags
260 if other.tags:
261 self.tags = list(set(self.tags + other.tags))
263 return self
265 def increment_step(self) -> int:
266 """递增步骤计数器,返回新 step。"""
267 self.step += 1
268 self.updated_at = datetime.now(UTC).isoformat()
269 return self.step
271 # ── Snapshot & Restore ────────────────────
273 def snapshot(self) -> dict[str, Any]:
274 """生成当前状态的完整快照。
276 Returns:
277 可序列化的状态字典(可直接存入 Checkpoint)
278 """
279 data = self.model_dump(mode="python", exclude_none=False)
280 # Remove private attrs
281 data.pop("_field_reducers", None)
282 data.pop("_field_custom_reducers", None)
283 data.pop("_sensitive_fields", None)
284 return data
286 def sanitized_snapshot(self) -> dict[str, Any]:
287 """生成脱敏快照(敏感字段替换为 '***')"""
288 data = self.snapshot()
289 for field in self._sensitive_fields:
290 if field in data:
291 data[field] = "***"
292 return data
294 @classmethod
295 def restore(cls, data: dict[str, Any]) -> BaseAgentState:
296 """从快照字典恢复状态。
298 Args:
299 data: snapshot() 返回的字典
301 Returns:
302 新的 AgentState 实例
303 """
305 def _clean_private(d: dict) -> dict:
306 return {k: v for k, v in d.items() if not k.startswith("_")}
308 cleaned = _clean_private(data)
309 instance = cls(**cleaned)
310 # Restore reducer configs from saved metadata if present
311 if "metadata" in cleaned and isinstance(cleaned["metadata"], dict):
312 reducer_config = cleaned["metadata"].get("_reducer_config")
313 if isinstance(reducer_config, dict):
314 for field, strategy_name in reducer_config.items():
315 try:
316 strategy = ReducerStrategy(strategy_name)
317 instance._field_reducers[field] = strategy
318 except ValueError:
319 pass
320 return instance
322 def diff(self, other: BaseAgentState) -> dict[str, tuple]:
323 """计算两个状态之间的差异。
325 Returns:
326 {field: (old_val, new_val)} 的字典
327 """
328 diffs = {}
329 for field in self.model_fields:
330 old = getattr(self, field, None)
331 new = getattr(other, field, None)
332 if old != new:
333 diffs[field] = (old, new)
334 return diffs
336 # ── JSON Schema ───────────────────────────
338 @classmethod
339 def generate_schema(cls) -> dict[str, Any]:
340 """生成状态的 JSON Schema(符合 OpenAI Function Calling 格式)。
342 Returns:
343 JSON Schema dict,可直接用作 tool/function 的 parameters 定义
344 """
345 return cls.model_json_schema()
347 @classmethod
348 def validate_json_input(cls, data: dict) -> BaseAgentState:
349 """从 JSON 字典验证并创建实例。"""
350 return cls.model_validate(data)
353# ── Specialized States ──────────────────────
356class AgentState(BaseAgentState):
357 """通用 Agent 运行状态。
359 预配置了常用字段和默认 reducer:
360 - messages: APPEND(对话消息累积)
361 - tools_result: MERGE(工具结果合并)
362 - errors: APPEND(错误收集)
363 - intermediate: REPLACE(中间结果替换)
364 """
366 messages: list[dict[str, Any]] = Field(
367 default_factory=list,
368 description="对话消息历史",
369 )
370 tools_result: dict[str, Any] = Field(
371 default_factory=dict,
372 description="最近一次工具调用结果",
373 )
374 intermediate: dict[str, Any] = Field(
375 default_factory=dict,
376 description="中间计算结果(每次替换)",
377 )
378 errors: list[dict[str, Any]] = Field(
379 default_factory=list,
380 description="错误堆栈",
381 )
382 human_interrupts: list[dict[str, Any]] = Field(
383 default_factory=list,
384 description="人工干预请求队列",
385 )
386 context_summary: str = Field(
387 default="",
388 description="上下文摘要(自动分页用)",
389 )
390 task_progress: float | None = Field(
391 default=None,
392 ge=0.0,
393 le=1.0,
394 description="任务进度 0.0~1.0",
395 )
396 abort_reason: str | None = Field(
397 default=None,
398 description="中止原因(非空表示需中止)",
399 )
401 def __init__(self, **data):
402 super().__init__(**data)
403 # Default reducers for AgentState
404 self._field_reducers["messages"] = ReducerStrategy.APPEND
405 self._field_reducers["tools_result"] = ReducerStrategy.MERGE
406 self._field_reducers["errors"] = ReducerStrategy.APPEND
407 self._field_reducers["human_interrupts"] = ReducerStrategy.APPEND
408 self._field_reducers["intermediate"] = ReducerStrategy.REPLACE
410 @property
411 def last_message(self) -> dict[str, Any] | None:
412 """获取最后一条消息。"""
413 return self.messages[-1] if self.messages else None
415 @property
416 def error_count(self) -> int:
417 """累计错误数。"""
418 return len(self.errors)
420 @property
421 def should_abort(self) -> bool:
422 """是否需要中止执行?"""
423 return self.abort_reason is not None
425 def add_message(self, role: str, content: str, **extra) -> AgentState:
426 """添加一条消息。"""
427 msg = {"role": role, "content": content, **extra}
428 self.update_field("messages", [msg])
429 return self
431 def add_error(self, error_type: str, message: str, **extra) -> AgentState:
432 """记录一个错误。"""
433 err = {
434 "type": error_type,
435 "message": message,
436 "step": self.step,
437 "timestamp": datetime.now(UTC).isoformat(),
438 **extra,
439 }
440 self.update_field("errors", [err])
441 return self
443 def request_human_input(
444 self,
445 prompt: str,
446 options: list[str] | None = None,
447 **extra,
448 ) -> AgentState:
449 """发起人工干预请求。"""
450 interrupt = {
451 "prompt": prompt,
452 "options": options,
453 "step": self.step,
454 "timestamp": datetime.now(UTC).isoformat(),
455 **extra,
456 }
457 self.update_field("human_interrupts", [interrupt])
458 return self
460 def clear_human_interrupts(self) -> AgentState:
461 """清除所有待处理的人工干预请求。"""
462 self.human_interrupts = []
463 return self
465 def abort(self, reason: str) -> AgentState:
466 """标记任务需要中止。"""
467 self.abort_reason = reason
468 return self
470 def reset_abort(self) -> AgentState:
471 """清除中止标记。"""
472 self.abort_reason = None
473 return self
476class MultiAgentState(BaseAgentState):
477 """多 Agent 协作状态。
479 管理多个子 Agent 的状态、消息路由、角色分配。
480 """
482 agents: dict[str, dict[str, Any]] = Field(
483 default_factory=dict,
484 description="所有子 Agent 的状态 {agent_id: {state dict}}",
485 )
486 message_queue: list[dict[str, Any]] = Field(
487 default_factory=list,
488 description="Agent 间消息队列",
489 )
490 roles: dict[str, str] = Field(
491 default_factory=dict,
492 description="Agent 角色分配 {agent_id: role_name}",
493 )
494 coordinator_state: dict[str, Any] = Field(
495 default_factory=dict,
496 description="协调器内部状态",
497 )
498 handoff_log: list[dict[str, Any]] = Field(
499 default_factory=list,
500 description="任务移交记录",
501 )
503 def __init__(self, **data):
504 super().__init__(**data)
505 self._field_reducers["agents"] = ReducerStrategy.MERGE
506 self._field_reducers["message_queue"] = ReducerStrategy.APPEND
507 self._field_reducers["handoff_log"] = ReducerStrategy.APPEND
509 def register_agent(self, agent_id: str, role: str = "worker", **meta) -> MultiAgentState:
510 """注册一个子 Agent。"""
511 self.agents[agent_id] = {
512 "role": role,
513 "state": "idle",
514 "step": 0,
515 "errors": 0,
516 **meta,
517 }
518 self.roles[agent_id] = role
519 return self
521 def update_agent_state(self, agent_id: str, updates: dict[str, Any]) -> MultiAgentState:
522 """更新子 Agent 的状态。"""
523 if agent_id in self.agents:
524 self.agents[agent_id].update(updates)
525 self._touch()
526 return self
528 def send_message(self, from_agent: str, to_agent: str, content: Any) -> MultiAgentState:
529 """Agent 间发送消息。"""
530 msg = {
531 "from": from_agent,
532 "to": to_agent,
533 "content": content,
534 "timestamp": datetime.now(UTC).isoformat(),
535 }
536 self.update_field("message_queue", [msg])
537 return self
539 def log_handoff(
540 self, from_agent: str, to_agent: str, task_id: str, reason: str = ""
541 ) -> MultiAgentState:
542 """记录任务移交。"""
543 self.update_field(
544 "handoff_log",
545 [
546 {
547 "from": from_agent,
548 "to": to_agent,
549 "task_id": task_id,
550 "reason": reason,
551 "timestamp": datetime.now(UTC).isoformat(),
552 }
553 ],
554 )
555 return self
557 def _touch(self) -> None:
558 self.updated_at = datetime.now(UTC).isoformat()
561class ToolCallState(BaseAgentState):
562 """工具调用追踪状态。
564 用于细粒度监控每次工具调用的入参/出参/耗时。
565 """
567 calls: list[dict[str, Any]] = Field(
568 default_factory=list,
569 description="工具调用历史",
570 )
571 pending_calls: dict[str, dict[str, Any]] = Field(
572 default_factory=dict,
573 description="进行中的工具调用 {call_id: {tool, args, start_time}}",
574 )
575 tool_stats: dict[str, dict[str, Any]] = Field(
576 default_factory=dict,
577 description="工具统计 {tool_name: {count, total_ms, errors, avg_ms}}",
578 )
580 def __init__(self, **data):
581 super().__init__(**data)
582 self._field_reducers["calls"] = ReducerStrategy.APPEND
584 def start_call(self, call_id: str, tool: str, args: dict) -> ToolCallState:
585 """记录工具调用开始。"""
586 self.pending_calls[call_id] = {
587 "tool": tool,
588 "args": args,
589 "start_time": time.time(),
590 "step": self.step,
591 }
592 self._init_stats(tool)
593 return self
595 def complete_call(self, call_id: str, result: Any, error: str | None = None) -> ToolCallState:
596 """记录工具调用完成。"""
597 if call_id not in self.pending_calls:
598 return self # 幂等
599 call_info = self.pending_calls.pop(call_id)
600 elapsed_ms = (time.time() - call_info["start_time"]) * 1000
601 record = {
602 **call_info,
603 "call_id": call_id,
604 "elapsed_ms": round(elapsed_ms, 2),
605 "result": result,
606 "error": error,
607 "success": error is None,
608 "timestamp": datetime.now(UTC).isoformat(),
609 }
610 self.update_field("calls", [record])
612 # Update stats
613 tool = call_info["tool"]
614 stats = self.tool_stats.get(tool, {})
615 stats["count"] = stats.get("count", 0) + 1
616 stats["total_ms"] = stats.get("total_ms", 0) + elapsed_ms
617 stats["errors"] = stats.get("errors", 0) + (1 if error else 0)
618 stats["avg_ms"] = stats["total_ms"] / stats["count"]
619 self.tool_stats[tool] = stats
621 return self
623 def _init_stats(self, tool: str) -> None:
624 """初始化工具统计。"""
625 if tool not in self.tool_stats:
626 self.tool_stats[tool] = {
627 "count": 0,
628 "total_ms": 0,
629 "errors": 0,
630 "avg_ms": 0,
631 }
633 @property
634 def total_tool_calls(self) -> int:
635 """总工具调用次数。"""
636 return len(self.calls)
638 @property
639 def failed_calls(self) -> int:
640 """失败的工具调用次数。"""
641 return sum(1 for c in self.calls if not c.get("success", True))
644# ── Schema Registry ──────────────────────────
647class StateSchemaRegistry:
648 """状态 Schema 注册中心。
650 支持按名称查找、注册、验证状态类型。
651 """
653 def __init__(self):
654 self._schemas: dict[str, type[BaseAgentState]] = {}
655 self._default_name: str = "AgentState"
656 self._schemas[self._default_name] = AgentState
657 self._schemas["MultiAgentState"] = MultiAgentState
658 self._schemas["ToolCallState"] = ToolCallState
660 def register(self, name: str, schema_cls: type[BaseAgentState]) -> None:
661 """注册自定义状态类型。"""
662 if not issubclass(schema_cls, BaseAgentState):
663 raise TypeError(
664 f"State class must inherit from BaseAgentState, " f"got {schema_cls.__name__}"
665 )
666 self._schemas[name] = schema_cls
668 def get(self, name: str) -> type[BaseAgentState]:
669 """获取已注册的状态类型。"""
670 if name not in self._schemas:
671 raise KeyError(
672 f"Unknown state schema: '{name}'. " f"Available: {list(self._schemas.keys())}"
673 )
674 return self._schemas[name]
676 def list_schemas(self) -> list[str]:
677 """列出所有已注册的状态类型。"""
678 return list(self._schemas.keys())
680 def create_state(self, name: str, **kwargs) -> BaseAgentState:
681 """创建已注册类型的实例。"""
682 cls = self.get(name)
683 return cls(**kwargs)
685 def validate(self, name: str, data: dict) -> BaseAgentState:
686 """验证并创建状态实例。"""
687 cls = self.get(name)
688 return cls.model_validate(data)
690 @property
691 def default_state_class(self) -> type[BaseAgentState]:
692 """默认状态类型。"""
693 return self._schemas[self._default_name]
695 @default_state_class.setter
696 def default_state_class(self, name: str) -> None:
697 if name not in self._schemas:
698 raise KeyError(f"Unknown state schema: '{name}'")
699 self._default_name = name
702# ── 全局单例 ─────────────────────────────────
704state_registry = StateSchemaRegistry()
707# ── State Reducers (test compatibility) ──
708class StateReducer:
709 """Base state reducer with merge strategy."""
711 @staticmethod
712 def merge(base_state, new_state):
713 data = base_state.model_dump()
714 new_data = new_state.model_dump(exclude_unset=True)
715 data.update(new_data)
716 return type(base_state)(**data)
719class LastWriteWinsReducer:
720 """Reducer: newer state wins based on version."""
722 @staticmethod
723 def merge(base_state, new_state):
724 bv = getattr(base_state, "version", 0)
725 nv = getattr(new_state, "version", 0)
726 if nv >= bv:
727 data = base_state.model_dump()
728 data.update(new_state.model_dump(exclude_unset=True))
729 return type(base_state)(**data)
730 return base_state
733class AppendOnlyReducer:
734 """Reducer: append-only for list fields."""
736 MERGE_FIELDS = ["messages", "logs"]
738 @staticmethod
739 def merge(base_state, new_state):
740 data = base_state.model_dump()
741 new_data = new_state.model_dump(exclude_unset=True)
742 for field in AppendOnlyReducer.MERGE_FIELDS:
743 if field in new_data and field in data:
744 data[field] = list(data[field]) + list(new_data[field])
745 data.update({k: v for k, v in new_data.items() if k not in AppendOnlyReducer.MERGE_FIELDS})
746 return type(base_state)(**data)