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
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-09 00:42 +0900
1"""GraphBuilder - Registry-based graph construction.
3Reads registered nodes from NodeRegistry and
4automatically builds LangGraph StateGraph.
5"""
6from __future__ import annotations
8from typing import Any, Callable, Optional
9from langchain_core.runnables import RunnableConfig
11from langgraph.graph import StateGraph, END
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
20logger = get_logger("agent_contracts.graph_builder")
23class GraphBuilder:
24 """Registry-based graph construction utility.
26 Example:
27 builder = GraphBuilder(registry)
28 builder.add_supervisor("shopping", llm)
29 builder.add_supervisor("card", llm)
30 graph = builder.build()
31 """
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.
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
58 def add_supervisor(
59 self,
60 name: str,
61 llm=None,
62 **services,
63 ) -> "GraphBuilder":
64 """Add Supervisor.
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
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
87 self.logger.info(
88 f"Added supervisor: {name} ({len(self.registry.get_supervisor_nodes(name))} nodes)"
89 )
91 return self
93 def build_routing_map(self, supervisor_name: str) -> dict[str, str]:
94 """Auto-generate routing map.
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
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)
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 {}
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)
126 wrapper.__name__ = f"{node_name}_node"
127 return wrapper
129 def create_supervisor_wrapper(self, supervisor_name: str) -> Callable:
130 """Create LangGraph-compatible Supervisor wrapper."""
131 supervisor = self.supervisor_instances.get(supervisor_name)
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)
149 wrapper.__name__ = f"{supervisor_name}_supervisor"
150 return wrapper
152 def create_routing_function(self, supervisor_name: str) -> Callable:
153 """Create routing function after Supervisor.
155 Automatically routes to 'done' if response_type is terminal.
156 """
157 valid_nodes = set(self.registry.get_supervisor_nodes(supervisor_name))
159 # Get terminal types from config
160 config = get_config()
161 terminal_types = set(config.supervisor.terminal_response_types)
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"
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"
177 route.__name__ = f"route_after_{supervisor_name}_supervisor"
178 return route
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.
193 Usage:
194 from agent_contracts import build_graph_from_registry
196 registry = get_node_registry()
197 graph = build_graph_from_registry(registry, llm=llm)
198 compiled = graph.compile()
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
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 )
221 # Add supervisors
222 supervisor_list = supervisors or []
223 for sup_name in supervisor_list:
224 builder.add_supervisor(sup_name, llm=llm, **services)
226 # Create StateGraph
227 state_cls = state_class or dict
228 graph = StateGraph(state_cls)
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))
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))
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)
243 graph.add_conditional_edges(
244 f"{sup_name}_supervisor",
245 route_fn,
246 routing_map,
247 )
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
254 if contract.is_terminal:
255 graph.add_edge(node_name, END)
256 else:
257 graph.add_edge(node_name, f"{sup_name}_supervisor")
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)
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)
272 logger.info(
273 f"Graph built: {len(builder.supervisor_names)} supervisors, {len(builder.node_classes)} nodes"
274 )
276 return graph