czty's picture
Add files using upload-large-folder tool
b2c86fd verified
Raw
History Blame Contribute Delete
12.3 kB
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)