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

1"""State Accessors - Type-safe state access pattern. 

2 

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. 

6 

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 

15 

16from typing import TypeVar, Generic, overload 

17 

18T = TypeVar("T") 

19 

20 

21class StateAccessor(Generic[T]): 

22 """Type-safe accessor for state fields. 

23  

24 Provides get/set methods that work immutably on state dictionaries. 

25 The accessor knows its slice name, field name, and default value. 

26  

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 

31  

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

39 

40 __slots__ = ("slice_name", "field_name", "default") 

41 

42 def __init__(self, slice_name: str, field_name: str, default: T) -> None: 

43 """Initialize the accessor. 

44  

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 

53 

54 def get(self, state: dict) -> T: 

55 """Get the field value from state. 

56  

57 Args: 

58 state: The state dictionary 

59  

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) 

67 

68 def set(self, state: dict, value: T) -> dict: 

69 """Set the field value and return a new state (immutable). 

70  

71 Args: 

72 state: The current state dictionary 

73 value: The new value to set 

74  

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} 

83 

84 def update(self, state: dict, func) -> dict: 

85 """Update the field using a function (immutable). 

86  

87 Args: 

88 state: The current state dictionary 

89 func: A function that takes the current value and returns new value 

90  

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) 

97 

98 def __repr__(self) -> str: 

99 return f"StateAccessor({self.slice_name!r}, {self.field_name!r}, default={self.default!r})" 

100 

101 

102# ============================================================================= 

103# Standard Accessors: _internal slice 

104# ============================================================================= 

105 

106class Internal: 

107 """Accessors for _internal slice fields. 

108  

109 These fields are used by the supervisor/router and should not be 

110 accessed directly by nodes. 

111 """ 

112 

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) 

120 

121 

122# ============================================================================= 

123# Standard Accessors: request slice 

124# ============================================================================= 

125 

126class Request: 

127 """Accessors for request slice fields. 

128  

129 Request slice contains information from the API request. 

130 Typically read-only for nodes. 

131 """ 

132 

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) 

138 

139 

140# ============================================================================= 

141# Standard Accessors: response slice 

142# ============================================================================= 

143 

144class Response: 

145 """Accessors for response slice fields. 

146  

147 Response slice contains information to return as API response. 

148 """ 

149 

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) 

153 

154 

155# ============================================================================= 

156# Convenience Functions 

157# ============================================================================= 

158 

159def reset_response(state: dict) -> dict: 

160 """Reset the response slice to empty values (immutable). 

161  

162 Args: 

163 state: Current state 

164  

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 

172 

173 

174def increment_turn(state: dict) -> dict: 

175 """Increment turn count and set is_first_turn to False (immutable). 

176  

177 Args: 

178 state: Current state 

179  

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 

186 

187 

188def set_error(state: dict, error: str) -> dict: 

189 """Set error in state (immutable). 

190  

191 Args: 

192 state: Current state 

193 error: Error message 

194  

195 Returns: 

196 New state with error set 

197 """ 

198 return Internal.error.set(state, error) 

199 

200 

201def clear_error(state: dict) -> dict: 

202 """Clear error in state (immutable). 

203  

204 Args: 

205 state: Current state 

206  

207 Returns: 

208 New state with error cleared 

209 """ 

210 return Internal.error.set(state, None)