Coverage for src / agent_contracts / state_accessors.py: 98%
56 statements
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-09 00:42 +0900
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-09 00:42 +0900
1"""State Accessors - Type-safe state access pattern.
3Provides immutable, type-safe access to state fields following
4the Redux selector pattern. All state modifications return new
5state dictionaries rather than mutating in place.
7Example:
8 >>> state = {"_internal": {"turn_count": 5}}
9 >>> count = Internal.turn_count.get(state) # 5
10 >>> new_state = Internal.turn_count.set(state, 10)
11 >>> Internal.turn_count.get(new_state) # 10
12 >>> Internal.turn_count.get(state) # 5 (original unchanged)
13"""
14from __future__ import annotations
16from typing import TypeVar, Generic, overload
18T = TypeVar("T")
21class StateAccessor(Generic[T]):
22 """Type-safe accessor for state fields.
24 Provides get/set methods that work immutably on state dictionaries.
25 The accessor knows its slice name, field name, and default value.
27 Attributes:
28 slice_name: Name of the state slice (e.g., "_internal", "request")
29 field_name: Name of the field within the slice
30 default: Default value if field is not present
32 Example:
33 >>> turn_count = StateAccessor("_internal", "turn_count", 0)
34 >>> state = {}
35 >>> turn_count.get(state) # Returns 0 (default)
36 >>> new_state = turn_count.set(state, 5)
37 >>> turn_count.get(new_state) # Returns 5
38 """
40 __slots__ = ("slice_name", "field_name", "default")
42 def __init__(self, slice_name: str, field_name: str, default: T) -> None:
43 """Initialize the accessor.
45 Args:
46 slice_name: Name of the state slice
47 field_name: Name of the field within the slice
48 default: Default value to return if field is absent
49 """
50 self.slice_name = slice_name
51 self.field_name = field_name
52 self.default = default
54 def get(self, state: dict) -> T:
55 """Get the field value from state.
57 Args:
58 state: The state dictionary
60 Returns:
61 The field value, or default if not present
62 """
63 slice_data = state.get(self.slice_name)
64 if not isinstance(slice_data, dict):
65 return self.default
66 return slice_data.get(self.field_name, self.default)
68 def set(self, state: dict, value: T) -> dict:
69 """Set the field value and return a new state (immutable).
71 Args:
72 state: The current state dictionary
73 value: The new value to set
75 Returns:
76 A new state dictionary with the updated value
77 """
78 current_slice = state.get(self.slice_name, {})
79 if not isinstance(current_slice, dict):
80 current_slice = {}
81 new_slice = {**current_slice, self.field_name: value}
82 return {**state, self.slice_name: new_slice}
84 def update(self, state: dict, func) -> dict:
85 """Update the field using a function (immutable).
87 Args:
88 state: The current state dictionary
89 func: A function that takes the current value and returns new value
91 Returns:
92 A new state dictionary with the updated value
93 """
94 current_value = self.get(state)
95 new_value = func(current_value)
96 return self.set(state, new_value)
98 def __repr__(self) -> str:
99 return f"StateAccessor({self.slice_name!r}, {self.field_name!r}, default={self.default!r})"
102# =============================================================================
103# Standard Accessors: _internal slice
104# =============================================================================
106class Internal:
107 """Accessors for _internal slice fields.
109 These fields are used by the supervisor/router and should not be
110 accessed directly by nodes.
111 """
113 # Core control fields
114 turn_count: StateAccessor[int] = StateAccessor("_internal", "turn_count", 0)
115 is_first_turn: StateAccessor[bool] = StateAccessor("_internal", "is_first_turn", True)
116 active_mode: StateAccessor[str | None] = StateAccessor("_internal", "active_mode", None)
117 next_node: StateAccessor[str | None] = StateAccessor("_internal", "next_node", None)
118 decision: StateAccessor[str | None] = StateAccessor("_internal", "decision", None)
119 error: StateAccessor[str | None] = StateAccessor("_internal", "error", None)
122# =============================================================================
123# Standard Accessors: request slice
124# =============================================================================
126class Request:
127 """Accessors for request slice fields.
129 Request slice contains information from the API request.
130 Typically read-only for nodes.
131 """
133 session_id: StateAccessor[str] = StateAccessor("request", "session_id", "")
134 action: StateAccessor[str] = StateAccessor("request", "action", "")
135 params: StateAccessor[dict | None] = StateAccessor("request", "params", None)
136 message: StateAccessor[str | None] = StateAccessor("request", "message", None)
137 image: StateAccessor[str | None] = StateAccessor("request", "image", None)
140# =============================================================================
141# Standard Accessors: response slice
142# =============================================================================
144class Response:
145 """Accessors for response slice fields.
147 Response slice contains information to return as API response.
148 """
150 response_type: StateAccessor[str | None] = StateAccessor("response", "response_type", None)
151 response_data: StateAccessor[dict | None] = StateAccessor("response", "response_data", None)
152 response_message: StateAccessor[str | None] = StateAccessor("response", "response_message", None)
155# =============================================================================
156# Convenience Functions
157# =============================================================================
159def reset_response(state: dict) -> dict:
160 """Reset the response slice to empty values (immutable).
162 Args:
163 state: Current state
165 Returns:
166 New state with response slice cleared
167 """
168 state = Response.response_type.set(state, None)
169 state = Response.response_data.set(state, None)
170 state = Response.response_message.set(state, None)
171 return state
174def increment_turn(state: dict) -> dict:
175 """Increment turn count and set is_first_turn to False (immutable).
177 Args:
178 state: Current state
180 Returns:
181 New state with incremented turn count
182 """
183 state = Internal.turn_count.update(state, lambda x: x + 1)
184 state = Internal.is_first_turn.set(state, False)
185 return state
188def set_error(state: dict, error: str) -> dict:
189 """Set error in state (immutable).
191 Args:
192 state: Current state
193 error: Error message
195 Returns:
196 New state with error set
197 """
198 return Internal.error.set(state, error)
201def clear_error(state: dict) -> dict:
202 """Clear error in state (immutable).
204 Args:
205 state: Current state
207 Returns:
208 New state with error cleared
209 """
210 return Internal.error.set(state, None)