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

1"""State Operations - Higher-level state manipulation helpers. 

2 

3Provides immutable helper functions for common state operations like 

4merging sessions, resetting flags, and creating initial state. 

5 

6All functions follow immutable patterns - they return new state dictionaries 

7rather than mutating the input. 

8""" 

9from __future__ import annotations 

10 

11from typing import Any 

12 

13from agent_contracts.state_accessors import ( 

14 StateAccessor, 

15 Internal, 

16 Request, 

17 Response, 

18 reset_response, 

19) 

20 

21 

22def ensure_slices(state: dict, slice_names: list[str]) -> dict: 

23 """Ensure specified slices exist in state (immutable). 

24  

25 Creates empty dicts for any missing slices. 

26  

27 Args: 

28 state: Current state 

29 slice_names: List of slice names to ensure exist 

30  

31 Returns: 

32 New state with all specified slices present 

33  

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 

44 

45 

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

52  

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. 

55  

56 Args: 

57 state: Current state 

58 session_data: Session data to merge 

59 slices: Slice names to merge (default: ["_internal", "interview", "shopping"]) 

60  

61 Returns: 

62 New state with session data merged 

63  

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

72 

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} 

82 

83 return result 

84 

85 

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

92  

93 Can pass flags as a dict or as keyword arguments. 

94 Only resets flags that have corresponding accessors in Internal. 

95  

96 Args: 

97 state: Current state 

98 flags: Dict of flag_name -> value to set 

99 **kwargs: Alternative way to specify flags 

100  

101 Returns: 

102 New state with flags reset 

103  

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 

112 

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) 

117 

118 return result 

119 

120 

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. 

130  

131 This is a simplified state factory for OSS use. Applications may 

132 need to extend this with additional slices. 

133  

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 

141  

142 Returns: 

143 New state dict with request, response, and _internal slices 

144  

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 = {} 

153 

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) 

160 

161 # Response slice (empty) 

162 state = reset_response(state) 

163 

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) 

171 

172 return state 

173 

174 

175def copy_slice(state: dict, slice_name: str) -> dict: 

176 """Get a shallow copy of a slice from state. 

177  

178 Args: 

179 state: Current state 

180 slice_name: Name of slice to copy 

181  

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 {} 

189 

190 

191def update_slice(state: dict, slice_name: str, **updates: Any) -> dict: 

192 """Update multiple fields in a slice at once (immutable). 

193  

194 Args: 

195 state: Current state 

196 slice_name: Name of slice to update 

197 **updates: Field updates as keyword arguments 

198  

199 Returns: 

200 New state with slice updated 

201  

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} 

213 

214 

215def get_nested(state: dict, *keys: str, default: Any = None) -> Any: 

216 """Get a nested value from state using a path of keys. 

217  

218 Args: 

219 state: Current state 

220 *keys: Path of keys to traverse 

221 default: Default value if path not found 

222  

223 Returns: 

224 Value at path, or default 

225  

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