from __future__ import annotations from collections import defaultdict from biomni.graph.schema_extractor import ToolSchemaExtractor class ToolGraph: """A lightweight heterogeneous graph over MCP servers, tools, capabilities, and data types.""" def __init__(self, schema_extractor: ToolSchemaExtractor | None = None): self.schema_extractor = schema_extractor or ToolSchemaExtractor() self.clear() def clear(self) -> None: self.nodes: dict[str, dict] = {} self.edges: list[dict] = [] self.adjacency: dict[str, list[dict]] = defaultdict(list) self.server_entries: dict[str, dict] = {} self.server_index: dict[str, dict] = {} self.tool_index: dict[str, dict] = {} self.operation_index: dict[str, dict] = {} self.datatype_to_operations: dict[str, list[str]] = defaultdict(list) def add_node(self, node_id: str, node_type: str, label: str, **attributes) -> None: payload = {"id": node_id, "type": node_type, "label": label} payload.update(attributes) self.nodes[node_id] = payload def add_edge(self, source: str, target: str, edge_type: str, weight: float = 1.0, **attributes) -> None: payload = {"source": source, "target": target, "type": edge_type, "weight": weight} payload.update(attributes) self.edges.append(payload) self.adjacency[source].append(payload) def build_from_server_entries(self, server_entries: list[dict]) -> None: self.clear() for entry in server_entries: self._add_server(entry) self._add_cross_server_workflow_edges() def get_server_entry(self, server_name: str) -> dict | None: return self.server_entries.get(server_name) def get_server_names(self) -> list[str]: return list(self.server_entries.keys()) def get_tool_nodes_for_server(self, server_name: str) -> list[dict]: return self.server_index.get(server_name, {}).get("tools", []) def get_server_semantics(self, server_name: str) -> dict: return self.server_index.get(server_name, {}).get("semantics", {}) def get_server_neighbors(self, server_name: str, max_hops: int = 2) -> list[dict]: start_node = f"server:{server_name}" if start_node not in self.nodes: return [] visited = {start_node} frontier = {start_node} collected = [] for _ in range(max_hops): next_frontier = set() for node_id in frontier: for edge in self.adjacency.get(node_id, []): if edge["target"] in visited: continue visited.add(edge["target"]) next_frontier.add(edge["target"]) collected.append(self.nodes[edge["target"]]) frontier = next_frontier if not frontier: break return collected def _add_server(self, entry: dict) -> None: server_name = entry["name"] server_node = f"server:{server_name}" semantics = self.schema_extractor.extract_server_semantics(entry) self.add_node( server_node, "server", server_name, category=entry.get("category", "general"), summary=entry.get("summary", ""), keywords=semantics.get("keywords", []), ) self.server_entries[server_name] = entry self.server_index[server_name] = {"entry": entry, "semantics": semantics, "tools": []} category = entry.get("category") if category: category_node = f"category:{category}" self.add_node(category_node, "category", category) self.add_edge(server_node, category_node, "categorized_as", weight=1.0) for tool, tool_semantics in zip(entry.get("tools", []), semantics.get("tool_semantics", []), strict=False): tool_name = tool_semantics["name"] tool_node = f"tool:{server_name}:{tool_name}" self.add_node( tool_node, "tool", tool_name, description=tool_semantics.get("description", ""), keywords=tool_semantics.get("keywords", []), module=entry.get("module"), ) self.add_edge(server_node, tool_node, "hosts", weight=3.0) self.server_index[server_name]["tools"].append({"node_id": tool_node, **tool, **tool_semantics}) self.tool_index[tool_node] = {"server_name": server_name, "tool": tool, "semantics": tool_semantics} stage = tool_semantics.get("stage") if stage: stage_node = f"stage:{stage}" self.add_node(stage_node, "stage", stage) self.add_edge(tool_node, stage_node, "belongs_to_stage", weight=1.0) for capability in tool_semantics.get("capabilities", []): capability_node = f"capability:{capability}" self.add_node(capability_node, "capability", capability) self.add_edge(tool_node, capability_node, "implements", weight=2.0) for operation in tool_semantics.get("operations", []): self._add_operation_binding(operation, tool_node, tool_semantics) for data_type in tool_semantics.get("consumes", []): datatype_node = f"datatype:{data_type}" self.add_node(datatype_node, "datatype", data_type) self.add_edge(tool_node, datatype_node, "consumes", weight=2.0) for data_type in tool_semantics.get("produces", []): datatype_node = f"datatype:{data_type}" self.add_node(datatype_node, "datatype", data_type) self.add_edge(tool_node, datatype_node, "produces", weight=2.0) for constraint in tool_semantics.get("constraints", []): constraint_node = f"constraint:{constraint}" self.add_node(constraint_node, "constraint", constraint) self.add_edge(tool_node, constraint_node, "supports", weight=1.5) self._add_workflow_edges(server_name) def _add_operation_binding(self, operation: str, tool_node: str, tool_semantics: dict) -> None: operation_node = f"operation:{operation}" spec = self.schema_extractor.OPERATION_SPECS.get(operation, {}) accepts = spec.get("accepts", []) or tool_semantics.get("consumes", []) produces = spec.get("produces", []) or tool_semantics.get("produces", []) constraints = spec.get("constraints", []) or tool_semantics.get("constraints", []) self.add_node( operation_node, "operation", operation, stage=spec.get("stage") or tool_semantics.get("stage"), accepts=accepts, produces=produces, constraints=constraints, ) self.add_edge(tool_node, operation_node, "implements_operation", weight=3.0) entry = self.operation_index.setdefault( operation, { "operation": operation, "node_id": operation_node, "accepts": [], "produces": [], "constraints": [], "tools": [], "stage": spec.get("stage") or tool_semantics.get("stage") or "analysis", }, ) entry["accepts"] = self._merge_labels(entry["accepts"], accepts) entry["produces"] = self._merge_labels(entry["produces"], produces) entry["constraints"] = self._merge_labels(entry["constraints"], constraints) entry["tools"].append({**tool_semantics, "node_id": tool_node}) for data_type in accepts: datatype_node = f"datatype:{data_type}" self.add_node(datatype_node, "datatype", data_type) self.add_edge(operation_node, datatype_node, "accepts", weight=2.0) if operation not in self.datatype_to_operations[data_type]: self.datatype_to_operations[data_type].append(operation) for data_type in produces: datatype_node = f"datatype:{data_type}" self.add_node(datatype_node, "datatype", data_type) self.add_edge(operation_node, datatype_node, "produces", weight=2.0) for constraint in constraints: constraint_node = f"constraint:{constraint}" self.add_node(constraint_node, "constraint", constraint) self.add_edge(operation_node, constraint_node, "requires", weight=1.5) def _merge_labels(self, primary: list[str], additions: list[str]) -> list[str]: merged = list(primary or []) for item in additions or []: if item and item not in merged: merged.append(item) return merged def _add_workflow_edges(self, server_name: str) -> None: tools = self.server_index.get(server_name, {}).get("tools", []) for source in tools: source_produced = set(source.get("produces", [])) source_stage = source.get("stage") for target in tools: if source["name"] == target["name"]: continue target_consumed = set(target.get("consumes", [])) if source_produced and target_consumed and source_produced & target_consumed: self.add_edge( source["node_id"], target["node_id"], "follows", weight=1.5, shared_datatypes=sorted(source_produced & target_consumed), ) elif source_stage and target.get("stage") and self._stage_distance(source_stage, target["stage"]) == 1: self.add_edge( source["node_id"], target["node_id"], "adjacent_stage", weight=0.5, ) def _add_cross_server_workflow_edges(self) -> None: consume_index: dict[str, list[dict]] = defaultdict(list) all_tools = [] for server_name in self.server_index: for tool in self.server_index[server_name].get("tools", []): all_tools.append(tool) for data_type in tool.get("consumes", []): consume_index[data_type].append(tool) generic_types = {"text", "image", "json", "csv"} for source in all_tools: produced_types = set(source.get("produces", [])) - generic_types if not produced_types: continue candidates = [] for data_type in produced_types: for target in consume_index.get(data_type, []): if source["node_id"] == target["node_id"]: continue if source["node_id"].split(":")[1] == target["node_id"].split(":")[1]: continue stage_delta = self._stage_distance(source.get("stage"), target.get("stage")) if stage_delta < 0 or stage_delta > 2: continue candidates.append((stage_delta, data_type, target)) candidates.sort(key=lambda item: (item[0], item[2].get("name", ""))) seen_targets = set() for stage_delta, data_type, target in candidates[:12]: if target["node_id"] in seen_targets: continue seen_targets.add(target["node_id"]) self.add_edge( source["node_id"], target["node_id"], "typed_flow", weight=2.5 if stage_delta <= 1 else 1.5, shared_datatypes=[data_type], cross_server=True, ) def _stage_distance(self, source_stage: str, target_stage: str) -> int: order = ["input_acquisition", "preprocessing", "analysis", "downstream", "reporting"] if source_stage not in order or target_stage not in order: return 99 return order.index(target_stage) - order.index(source_stage)