Coverage for src / agent_contracts / graph_builder.py: 16%

122 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-01-09 00:42 +0900

1"""GraphBuilder - Registry-based graph construction. 

2 

3Reads registered nodes from NodeRegistry and 

4automatically builds LangGraph StateGraph. 

5""" 

6from __future__ import annotations 

7 

8from typing import Any, Callable, Optional 

9from langchain_core.runnables import RunnableConfig 

10 

11from langgraph.graph import StateGraph, END 

12 

13from agent_contracts.registry import NodeRegistry, get_node_registry 

14from agent_contracts.supervisor import GenericSupervisor 

15from agent_contracts.state import merge_slice_updates 

16from agent_contracts.contracts import NodeContract 

17from agent_contracts.config import get_config 

18from agent_contracts.utils.logging import get_logger 

19 

20logger = get_logger("agent_contracts.graph_builder") 

21 

22 

23class GraphBuilder: 

24 """Registry-based graph construction utility. 

25  

26 Example: 

27 builder = GraphBuilder(registry) 

28 builder.add_supervisor("shopping", llm) 

29 builder.add_supervisor("card", llm) 

30 graph = builder.build() 

31 """ 

32 

33 def __init__( 

34 self, 

35 registry: NodeRegistry | None = None, 

36 state_class: type | None = None, 

37 llm_provider: Callable[[], Any] | None = None, 

38 dependency_provider: Callable[[NodeContract], dict] | None = None, 

39 ): 

40 """Initialize. 

41  

42 Args: 

43 registry: Node registry 

44 state_class: State type (uses dict if not provided) 

45 llm_provider: Function that provides LLM instances 

46 dependency_provider: Function that provides dependencies for nodes 

47 """ 

48 self.registry = registry or get_node_registry() 

49 self.state_class = state_class 

50 self.supervisor_names: set[str] = set() 

51 self.supervisor_instances: dict[str, GenericSupervisor] = {} 

52 self.node_classes: dict[str, type] = {} 

53 self.node_instances: dict[str, Any] = {} 

54 self.llm_provider = llm_provider 

55 self.dependency_provider = dependency_provider 

56 self.logger = logger 

57 

58 def add_supervisor( 

59 self, 

60 name: str, 

61 llm=None, 

62 **services, 

63 ) -> "GraphBuilder": 

64 """Add Supervisor. 

65  

66 Related node instances are also created. 

67 """ 

68 self.supervisor_names.add(name) 

69 if self.llm_provider is None: 

70 supervisor = GenericSupervisor( 

71 supervisor_name=name, 

72 llm=llm, 

73 registry=self.registry, 

74 ) 

75 self.supervisor_instances[name] = supervisor 

76 

77 # Create related node instances 

78 for node_name in self.registry.get_supervisor_nodes(name): 

79 node_cls = self.registry.get_node_class(node_name) 

80 if node_cls is None: 

81 continue 

82 self.node_classes[node_name] = node_cls 

83 if self.dependency_provider is None and self.llm_provider is None: 

84 instance = node_cls(llm=llm, **services) 

85 self.node_instances[node_name] = instance 

86 

87 self.logger.info( 

88 f"Added supervisor: {name} ({len(self.registry.get_supervisor_nodes(name))} nodes)" 

89 ) 

90 

91 return self 

92 

93 def build_routing_map(self, supervisor_name: str) -> dict[str, str]: 

94 """Auto-generate routing map. 

95  

96 Returns: 

97 {"node_name": "node_name", "done": END} 

98 """ 

99 routing = {name: name for name in self.registry.get_supervisor_nodes(supervisor_name)} 

100 routing["done"] = END # LangGraph END constant 

101 return routing 

102 

103 def create_node_wrapper(self, node_name: str) -> Callable: 

104 """Create LangGraph-compatible node wrapper.""" 

105 node_cls = self.node_classes.get(node_name) 

106 instance = self.node_instances.get(node_name) 

107 

108 async def wrapper(state: dict, config: Optional[RunnableConfig] = None) -> dict: 

109 if node_cls is None: 

110 self.logger.error(f"Node class not found: {node_name}") 

111 return {} 

112 

113 if self.dependency_provider or self.llm_provider: 

114 contract = node_cls.CONTRACT 

115 services = self.dependency_provider(contract) if self.dependency_provider else {} 

116 llm = self.llm_provider() if (self.llm_provider and contract.requires_llm) else None 

117 node = node_cls(llm=llm, **services) 

118 updates = await node(state, config=config) 

119 else: 

120 if instance is None: 

121 self.logger.error(f"Node instance not found: {node_name}") 

122 return {} 

123 updates = await instance(state, config=config) 

124 return merge_slice_updates(state, updates) 

125 

126 wrapper.__name__ = f"{node_name}_node" 

127 return wrapper 

128 

129 def create_supervisor_wrapper(self, supervisor_name: str) -> Callable: 

130 """Create LangGraph-compatible Supervisor wrapper.""" 

131 supervisor = self.supervisor_instances.get(supervisor_name) 

132 

133 async def wrapper(state: dict, config: Optional[RunnableConfig] = None) -> dict: 

134 if self.llm_provider: 

135 llm = self.llm_provider() 

136 current = GenericSupervisor( 

137 supervisor_name=supervisor_name, 

138 llm=llm, 

139 registry=self.registry, 

140 ) 

141 updates = await current.run(state, config=config) 

142 else: 

143 if supervisor is None: 

144 self.logger.error(f"Supervisor not found: {supervisor_name}") 

145 return {} 

146 updates = await supervisor.run(state, config=config) 

147 return merge_slice_updates(state, updates) 

148 

149 wrapper.__name__ = f"{supervisor_name}_supervisor" 

150 return wrapper 

151 

152 def create_routing_function(self, supervisor_name: str) -> Callable: 

153 """Create routing function after Supervisor. 

154  

155 Automatically routes to 'done' if response_type is terminal. 

156 """ 

157 valid_nodes = set(self.registry.get_supervisor_nodes(supervisor_name)) 

158 

159 # Get terminal types from config 

160 config = get_config() 

161 terminal_types = set(config.supervisor.terminal_response_types) 

162 

163 def route(state: dict) -> str: 

164 # First check response_type (termination signal from node) 

165 response = state.get("response", {}) 

166 response_type = response.get("response_type") 

167 if response_type in terminal_types: 

168 return "done" 

169 

170 # Then check decision 

171 internal = state.get("_internal", {}) 

172 decision = internal.get("decision", "done") 

173 if decision in valid_nodes: 

174 return decision 

175 return "done" 

176 

177 route.__name__ = f"route_after_{supervisor_name}_supervisor" 

178 return route 

179 

180 

181def build_graph_from_registry( 

182 registry: NodeRegistry | None = None, 

183 llm=None, 

184 llm_provider: Callable[[], Any] | None = None, 

185 dependency_provider: Callable[[NodeContract], dict] | None = None, 

186 entrypoint: tuple[str, Callable, Callable] | None = None, 

187 supervisors: list[str] | None = None, 

188 state_class: type | None = None, 

189 **services, 

190) -> StateGraph: 

191 """Auto-build graph from registry. 

192  

193 Usage: 

194 from agent_contracts import build_graph_from_registry 

195  

196 registry = get_node_registry() 

197 graph = build_graph_from_registry(registry, llm=llm) 

198 compiled = graph.compile() 

199  

200 Args: 

201 registry: Node registry 

202 llm: LLM instance (for all supervisors) 

203 llm_provider: Function to get LLM instances 

204 dependency_provider: Function to get dependencies for nodes 

205 entrypoint: (name, node_func, route_func) tuple for entry point 

206 supervisors: List of supervisor names to add 

207 state_class: State class for StateGraph 

208 **services: Services to inject into nodes 

209  

210 Returns: 

211 StateGraph (not compiled) 

212 """ 

213 reg = registry or get_node_registry() 

214 builder = GraphBuilder( 

215 registry=reg, 

216 state_class=state_class, 

217 llm_provider=llm_provider, 

218 dependency_provider=dependency_provider, 

219 ) 

220 

221 # Add supervisors 

222 supervisor_list = supervisors or [] 

223 for sup_name in supervisor_list: 

224 builder.add_supervisor(sup_name, llm=llm, **services) 

225 

226 # Create StateGraph 

227 state_cls = state_class or dict 

228 graph = StateGraph(state_cls) 

229 

230 # Add Supervisor nodes 

231 for sup_name in builder.supervisor_names: 

232 graph.add_node(f"{sup_name}_supervisor", builder.create_supervisor_wrapper(sup_name)) 

233 

234 # Add worker nodes 

235 for node_name in builder.node_classes.keys(): 

236 graph.add_node(node_name, builder.create_node_wrapper(node_name)) 

237 

238 # Supervisor -> workers conditional edges 

239 for sup_name in builder.supervisor_names: 

240 route_fn = builder.create_routing_function(sup_name) 

241 routing_map = builder.build_routing_map(sup_name) 

242 

243 graph.add_conditional_edges( 

244 f"{sup_name}_supervisor", 

245 route_fn, 

246 routing_map, 

247 ) 

248 

249 # Worker -> Supervisor return edges 

250 for node_name, node_cls in builder.node_classes.items(): 

251 contract = node_cls.CONTRACT 

252 sup_name = contract.supervisor 

253 

254 if contract.is_terminal: 

255 graph.add_edge(node_name, END) 

256 else: 

257 graph.add_edge(node_name, f"{sup_name}_supervisor") 

258 

259 # Entry point 

260 if entrypoint: 

261 entry_name, entry_node, route_fn = entrypoint 

262 graph.add_node(entry_name, entry_node) 

263 graph.set_entry_point(entry_name) 

264 

265 routing_map = { 

266 f"{sup_name}_supervisor": f"{sup_name}_supervisor" 

267 for sup_name in builder.supervisor_names 

268 } 

269 routing_map[END] = END 

270 graph.add_conditional_edges(entry_name, route_fn, routing_map) 

271 

272 logger.info( 

273 f"Graph built: {len(builder.supervisor_names)} supervisors, {len(builder.node_classes)} nodes" 

274 ) 

275 

276 return graph