Coverage for src / agent_contracts / runtime / state_ops.py: 97%
65 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 Operations - Higher-level state manipulation helpers.
3Provides immutable helper functions for common state operations like
4merging sessions, resetting flags, and creating initial state.
6All functions follow immutable patterns - they return new state dictionaries
7rather than mutating the input.
8"""
9from __future__ import annotations
11from typing import Any
13from agent_contracts.state_accessors import (
14 StateAccessor,
15 Internal,
16 Request,
17 Response,
18 reset_response,
19)
22def ensure_slices(state: dict, slice_names: list[str]) -> dict:
23 """Ensure specified slices exist in state (immutable).
25 Creates empty dicts for any missing slices.
27 Args:
28 state: Current state
29 slice_names: List of slice names to ensure exist
31 Returns:
32 New state with all specified slices present
34 Example:
35 >>> state = {}
36 >>> state = ensure_slices(state, ["request", "response", "_internal"])
37 >>> "request" in state # True
38 """
39 result = dict(state)
40 for name in slice_names:
41 if name not in result or not isinstance(result.get(name), dict):
42 result[name] = {}
43 return result
46def merge_session(
47 state: dict,
48 session_data: dict,
49 slices: list[str] | None = None,
50) -> dict:
51 """Merge session data into state (immutable).
53 For each specified slice, merges session data with current state.
54 Does not create new slices if they don't exist in session_data.
56 Args:
57 state: Current state
58 session_data: Session data to merge
59 slices: Slice names to merge (default: ["_internal", "interview", "shopping"])
61 Returns:
62 New state with session data merged
64 Example:
65 >>> state = {"interview": {"count": 1}}
66 >>> session = {"interview": {"history": ["q1"]}}
67 >>> merged = merge_session(state, session, ["interview"])
68 >>> merged["interview"] # {"count": 1, "history": ["q1"]}
69 """
70 if slices is None:
71 slices = ["_internal", "interview", "shopping"]
73 result = dict(state)
74 for slice_name in slices:
75 if slice_name in session_data:
76 current_slice = result.get(slice_name, {})
77 if not isinstance(current_slice, dict):
78 current_slice = {}
79 session_slice = session_data[slice_name]
80 if isinstance(session_slice, dict):
81 result[slice_name] = {**current_slice, **session_slice}
83 return result
86def reset_internal_flags(
87 state: dict,
88 flags: dict[str, Any] | None = None,
89 **kwargs: Any,
90) -> dict:
91 """Reset internal flags to specified values (immutable).
93 Can pass flags as a dict or as keyword arguments.
94 Only resets flags that have corresponding accessors in Internal.
96 Args:
97 state: Current state
98 flags: Dict of flag_name -> value to set
99 **kwargs: Alternative way to specify flags
101 Returns:
102 New state with flags reset
104 Example:
105 >>> state = reset_internal_flags(state,
106 ... turn_count=0,
107 ... is_first_turn=True,
108 ... )
109 """
110 all_flags = {**(flags or {}), **kwargs}
111 result = state
113 for flag_name, value in all_flags.items():
114 accessor = getattr(Internal, flag_name, None)
115 if accessor is not None and isinstance(accessor, StateAccessor):
116 result = accessor.set(result, value)
118 return result
121def create_base_state(
122 session_id: str = "",
123 action: str = "",
124 params: dict | None = None,
125 message: str | None = None,
126 image: str | None = None,
127 active_mode: str | None = None,
128) -> dict:
129 """Create a minimal base state with request and internal slices.
131 This is a simplified state factory for OSS use. Applications may
132 need to extend this with additional slices.
134 Args:
135 session_id: Session identifier
136 action: Action to perform
137 params: Optional action parameters
138 message: Optional user message
139 image: Optional base64-encoded image
140 active_mode: Optional initial mode
142 Returns:
143 New state dict with request, response, and _internal slices
145 Example:
146 >>> state = create_base_state(
147 ... session_id="abc123",
148 ... action="answer",
149 ... message="I like casual style",
150 ... )
151 """
152 state: dict = {}
154 # Request slice
155 state = Request.session_id.set(state, session_id)
156 state = Request.action.set(state, action)
157 state = Request.params.set(state, params)
158 state = Request.message.set(state, message)
159 state = Request.image.set(state, image)
161 # Response slice (empty)
162 state = reset_response(state)
164 # Internal slice
165 state = Internal.turn_count.set(state, 0)
166 state = Internal.is_first_turn.set(state, True)
167 state = Internal.active_mode.set(state, active_mode)
168 state = Internal.next_node.set(state, None)
169 state = Internal.decision.set(state, None)
170 state = Internal.error.set(state, None)
172 return state
175def copy_slice(state: dict, slice_name: str) -> dict:
176 """Get a shallow copy of a slice from state.
178 Args:
179 state: Current state
180 slice_name: Name of slice to copy
182 Returns:
183 Copy of the slice dict, or empty dict if not present
184 """
185 slice_data = state.get(slice_name, {})
186 if isinstance(slice_data, dict):
187 return dict(slice_data)
188 return {}
191def update_slice(state: dict, slice_name: str, **updates: Any) -> dict:
192 """Update multiple fields in a slice at once (immutable).
194 Args:
195 state: Current state
196 slice_name: Name of slice to update
197 **updates: Field updates as keyword arguments
199 Returns:
200 New state with slice updated
202 Example:
203 >>> state = update_slice(state, "interview",
204 ... question_count=5,
205 ... last_question={"text": "..."},
206 ... )
207 """
208 current_slice = state.get(slice_name, {})
209 if not isinstance(current_slice, dict):
210 current_slice = {}
211 new_slice = {**current_slice, **updates}
212 return {**state, slice_name: new_slice}
215def get_nested(state: dict, *keys: str, default: Any = None) -> Any:
216 """Get a nested value from state using a path of keys.
218 Args:
219 state: Current state
220 *keys: Path of keys to traverse
221 default: Default value if path not found
223 Returns:
224 Value at path, or default
226 Example:
227 >>> state = {"interview": {"collected_info": {"name": "Alice"}}}
228 >>> get_nested(state, "interview", "collected_info", "name") # "Alice"
229 >>> get_nested(state, "interview", "missing", default="unknown") # "unknown"
230 """
231 current = state
232 for key in keys:
233 if not isinstance(current, dict):
234 return default
235 current = current.get(key)
236 if current is None:
237 return default
238 return current