Coverage for agentos/state/schema.py: 0%

329 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-08 01:44 +0800

1""" 

2AgentOS v1.14.0 — 结构化 Agent 状态管理系统。 

3 

4基因来源: LangGraph Pydantic State Schema + AgentOS Checkpoint。 

5 

6核心设计: 

7- AgentState: 强类型的全局 Agent 运行时状态,Pydantic v2 驱动 

8- 支持 JSON Schema 自动生成、验证、序列化 

9- 与 Checkpoint 系统无缝对接 

10- 支持状态合并策略(reducer):append/extend/replace/merge 

11- 支持子状态派生(SubState),实现层级化状态管理 

12""" 

13 

14from __future__ import annotations 

15 

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) 

25 

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 ) 

37 

38# ── Reducers ──────────────────────────────── 

39 

40 

41class ReducerStrategy(StrEnum): 

42 """状态合并策略。""" 

43 

44 REPLACE = "replace" # 直接替换 

45 APPEND = "append" # 追加(list -> extend) 

46 EXTEND = "extend" # 字典合并 

47 MERGE = "merge" # 深度递归合并 

48 KEEP_EXISTING = "keep" # 保留旧值 

49 CUSTOM = "custom" # 自定义 reducer 函数 

50 

51 

52# Type variable for generic state 

53S = TypeVar("S", bound=BaseModel) 

54 

55# Custom reducer type 

56ReducerFn = Callable[[Any, Any], Any] 

57 

58# Default reducers 

59_DEFAULT_REDUCERS: dict[str, ReducerStrategy] = {} 

60 

61 

62def default_reducer(field_name: str, strategy: ReducerStrategy) -> None: 

63 """注册字段的默认合并策略。 

64 

65 Usage: 

66 default_reducer("messages", ReducerStrategy.APPEND) 

67 """ 

68 _DEFAULT_REDUCERS[field_name] = strategy 

69 

70 

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 

90 

91 

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 

103 

104 

105# ── Field Metadata ────────────────────────── 

106 

107 

108class StateFieldInfo(BaseModel): 

109 """状态字段元信息。""" 

110 

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 

117 

118 model_config = ConfigDict(arbitrary_types_allowed=True) 

119 

120 

121# ── AgentState Core ───────────────────────── 

122 

123 

124class BaseAgentState(BaseModel): 

125 """Agent 状态的基类。 

126 

127 所有 Agent 状态必须继承此类。自动提供: 

128 - thread_id / session_id 追踪 

129 - step 计数器 

130 - 状态快照(snapshot)与恢复(restore) 

131 - JSON Schema 生成 

132 - Checkpoint 序列化 

133 

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 """ 

140 

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(用于分支/回溯)") 

159 

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) 

164 

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 ) 

173 

174 def __init__(self, **data): 

175 super().__init__(**data) 

176 self._field_reducers = {} 

177 self._field_custom_reducers = {} 

178 self._sensitive_fields = set() 

179 

180 # ── Reducer Registration ────────────────── 

181 

182 def set_reducer(self, field: str, strategy: ReducerStrategy) -> BaseAgentState: 

183 """为字段设置合并策略。""" 

184 self._field_reducers[field] = strategy 

185 return self 

186 

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 

192 

193 def mark_sensitive(self, *fields: str) -> BaseAgentState: 

194 """标记敏感字段。""" 

195 self._sensitive_fields.update(fields) 

196 return self 

197 

198 # ── State Mutation ──────────────────────── 

199 

200 def update_field( 

201 self, 

202 field: str, 

203 value: Any, 

204 reducer: ReducerStrategy | None = None, 

205 ) -> None: 

206 """更新单个字段,自动应用 Reducer。 

207 

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) 

229 

230 self.updated_at = datetime.now(UTC).isoformat() 

231 

232 def merge(self, other: BaseAgentState | dict) -> BaseAgentState: 

233 """合并另一个状态到当前状态。 

234 

235 Args: 

236 other: 另一个 AgentState 实例或字典 

237 

238 Returns: 

239 self (in-place merge) 

240 """ 

241 if isinstance(other, dict): 

242 other = self.__class__(**other) 

243 

244 for field_name in other.model_fields: 

245 if field_name in ("thread_id", "created_at"): 

246 continue # Immutable fields 

247 

248 other_val = getattr(other, field_name, None) 

249 if other_val is None: 

250 continue 

251 

252 self.update_field(field_name, other_val) 

253 

254 # Merge metadata 

255 if other.metadata: 

256 self.metadata = _deep_merge(self.metadata, other.metadata) 

257 self.updated_at = datetime.now(UTC).isoformat() 

258 

259 # Merge tags 

260 if other.tags: 

261 self.tags = list(set(self.tags + other.tags)) 

262 

263 return self 

264 

265 def increment_step(self) -> int: 

266 """递增步骤计数器,返回新 step。""" 

267 self.step += 1 

268 self.updated_at = datetime.now(UTC).isoformat() 

269 return self.step 

270 

271 # ── Snapshot & Restore ──────────────────── 

272 

273 def snapshot(self) -> dict[str, Any]: 

274 """生成当前状态的完整快照。 

275 

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 

285 

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 

293 

294 @classmethod 

295 def restore(cls, data: dict[str, Any]) -> BaseAgentState: 

296 """从快照字典恢复状态。 

297 

298 Args: 

299 data: snapshot() 返回的字典 

300 

301 Returns: 

302 新的 AgentState 实例 

303 """ 

304 

305 def _clean_private(d: dict) -> dict: 

306 return {k: v for k, v in d.items() if not k.startswith("_")} 

307 

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 

321 

322 def diff(self, other: BaseAgentState) -> dict[str, tuple]: 

323 """计算两个状态之间的差异。 

324 

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 

335 

336 # ── JSON Schema ─────────────────────────── 

337 

338 @classmethod 

339 def generate_schema(cls) -> dict[str, Any]: 

340 """生成状态的 JSON Schema(符合 OpenAI Function Calling 格式)。 

341 

342 Returns: 

343 JSON Schema dict,可直接用作 tool/function 的 parameters 定义 

344 """ 

345 return cls.model_json_schema() 

346 

347 @classmethod 

348 def validate_json_input(cls, data: dict) -> BaseAgentState: 

349 """从 JSON 字典验证并创建实例。""" 

350 return cls.model_validate(data) 

351 

352 

353# ── Specialized States ────────────────────── 

354 

355 

356class AgentState(BaseAgentState): 

357 """通用 Agent 运行状态。 

358 

359 预配置了常用字段和默认 reducer: 

360 - messages: APPEND(对话消息累积) 

361 - tools_result: MERGE(工具结果合并) 

362 - errors: APPEND(错误收集) 

363 - intermediate: REPLACE(中间结果替换) 

364 """ 

365 

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 ) 

400 

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 

409 

410 @property 

411 def last_message(self) -> dict[str, Any] | None: 

412 """获取最后一条消息。""" 

413 return self.messages[-1] if self.messages else None 

414 

415 @property 

416 def error_count(self) -> int: 

417 """累计错误数。""" 

418 return len(self.errors) 

419 

420 @property 

421 def should_abort(self) -> bool: 

422 """是否需要中止执行?""" 

423 return self.abort_reason is not None 

424 

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 

430 

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 

442 

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 

459 

460 def clear_human_interrupts(self) -> AgentState: 

461 """清除所有待处理的人工干预请求。""" 

462 self.human_interrupts = [] 

463 return self 

464 

465 def abort(self, reason: str) -> AgentState: 

466 """标记任务需要中止。""" 

467 self.abort_reason = reason 

468 return self 

469 

470 def reset_abort(self) -> AgentState: 

471 """清除中止标记。""" 

472 self.abort_reason = None 

473 return self 

474 

475 

476class MultiAgentState(BaseAgentState): 

477 """多 Agent 协作状态。 

478 

479 管理多个子 Agent 的状态、消息路由、角色分配。 

480 """ 

481 

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 ) 

502 

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 

508 

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 

520 

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 

527 

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 

538 

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 

556 

557 def _touch(self) -> None: 

558 self.updated_at = datetime.now(UTC).isoformat() 

559 

560 

561class ToolCallState(BaseAgentState): 

562 """工具调用追踪状态。 

563 

564 用于细粒度监控每次工具调用的入参/出参/耗时。 

565 """ 

566 

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 ) 

579 

580 def __init__(self, **data): 

581 super().__init__(**data) 

582 self._field_reducers["calls"] = ReducerStrategy.APPEND 

583 

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 

594 

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

611 

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 

620 

621 return self 

622 

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 } 

632 

633 @property 

634 def total_tool_calls(self) -> int: 

635 """总工具调用次数。""" 

636 return len(self.calls) 

637 

638 @property 

639 def failed_calls(self) -> int: 

640 """失败的工具调用次数。""" 

641 return sum(1 for c in self.calls if not c.get("success", True)) 

642 

643 

644# ── Schema Registry ────────────────────────── 

645 

646 

647class StateSchemaRegistry: 

648 """状态 Schema 注册中心。 

649 

650 支持按名称查找、注册、验证状态类型。 

651 """ 

652 

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 

659 

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 

667 

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] 

675 

676 def list_schemas(self) -> list[str]: 

677 """列出所有已注册的状态类型。""" 

678 return list(self._schemas.keys()) 

679 

680 def create_state(self, name: str, **kwargs) -> BaseAgentState: 

681 """创建已注册类型的实例。""" 

682 cls = self.get(name) 

683 return cls(**kwargs) 

684 

685 def validate(self, name: str, data: dict) -> BaseAgentState: 

686 """验证并创建状态实例。""" 

687 cls = self.get(name) 

688 return cls.model_validate(data) 

689 

690 @property 

691 def default_state_class(self) -> type[BaseAgentState]: 

692 """默认状态类型。""" 

693 return self._schemas[self._default_name] 

694 

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 

700 

701 

702# ── 全局单例 ───────────────────────────────── 

703 

704state_registry = StateSchemaRegistry() 

705 

706 

707# ── State Reducers (test compatibility) ── 

708class StateReducer: 

709 """Base state reducer with merge strategy.""" 

710 

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) 

717 

718 

719class LastWriteWinsReducer: 

720 """Reducer: newer state wins based on version.""" 

721 

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 

731 

732 

733class AppendOnlyReducer: 

734 """Reducer: append-only for list fields.""" 

735 

736 MERGE_FIELDS = ["messages", "logs"] 

737 

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)