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