czty's picture
Add files using upload-large-folder tool
d1ce356 verified
Raw
History Blame Contribute Delete
12.6 kB
from __future__ import annotations
import re
from biomni.graph.tool_graph import ToolGraph
class MCPGraphRAGRetriever:
"""Generic GraphRAG retriever over the full MCP tool graph.
This class intentionally avoids benchmark- or task-specific allow/deny
policies. It only retrieves a compact graph context from semantic overlap
between the query and graph nodes; the agentic planner decides how to use
the retrieved graph.
"""
TOKEN_PATTERN = re.compile(r"[a-z0-9_+-]+")
STAGE_RANK = {
"input_acquisition": 0,
"preprocessing": 1,
"analysis": 2,
"downstream": 3,
"reporting": 4,
}
GENERIC_TOKENS = {
"answer",
"based",
"biological",
"case",
"choice",
"choices",
"data",
"following",
"question",
"result",
"task",
"tool",
"using",
"which",
}
def retrieve(
self,
*,
query_context: dict,
query_semantics: dict,
tool_graph: ToolGraph,
candidate_pool: int = 24,
top_k_tools: int = 12,
) -> dict:
query_tokens = self._query_tokens(query_context, query_semantics)
query_operations = set(query_semantics.get("operations", []))
query_capabilities = set(query_semantics.get("capabilities", []))
query_datatypes = set(query_semantics.get("datatypes", []))
query_stages = set(query_semantics.get("stages", []))
candidates = []
for server_name, payload in tool_graph.server_index.items():
server_entry = tool_graph.get_server_entry(server_name) or {}
server_text = self._server_text(server_name, server_entry, payload)
server_tokens = self._tokenize(server_text)
server_overlap = query_tokens & server_tokens
for tool in payload.get("tools", []):
score, reasons = self._score_tool(
server_name=server_name,
server_tokens=server_tokens,
server_overlap=server_overlap,
tool=tool,
query_tokens=query_tokens,
query_operations=query_operations,
query_capabilities=query_capabilities,
query_datatypes=query_datatypes,
query_stages=query_stages,
tool_graph=tool_graph,
)
if score <= 0:
continue
candidates.append(
{
"server": server_name,
"tool": tool.get("name"),
"node_id": tool.get("node_id"),
"stage": tool.get("stage") or "analysis",
"operations": tool.get("operations", []),
"capabilities": tool.get("capabilities", []),
"consumes": tool.get("consumes", []),
"produces": tool.get("produces", []),
"score": round(score, 3),
"reason": "; ".join(reasons[:6]),
"graph_context": self._expand_tool_context(tool, tool_graph),
}
)
candidates.sort(key=lambda item: (item["score"], item["server"], item.get("tool") or ""), reverse=True)
shortlist = candidates[: max(candidate_pool, top_k_tools)]
compact = self._compact_subgraph(shortlist[:top_k_tools], tool_graph)
return {
"retriever": "mcp_graph_rag",
"task_profile": {
"categories": [str(item).lower() for item in query_context.get("categories", [])],
"operations": sorted(query_operations),
"capabilities": sorted(query_capabilities),
"datatypes": sorted(query_datatypes),
"stages": sorted(query_stages),
},
"query_tokens": sorted(query_tokens),
"candidate_tools": shortlist,
"compact_subgraph": compact,
"excluded_servers": [],
"operation_path_hint": self._operation_path_hint(shortlist[:top_k_tools]),
"context_text": self._format_context(shortlist[:top_k_tools], compact),
}
def _query_tokens(self, query_context: dict, query_semantics: dict) -> set[str]:
parts = [
query_context.get("question", ""),
query_context.get("original_query", ""),
query_context.get("rewritten_query", ""),
query_context.get("retrieval_query", ""),
" ".join(query_context.get("categories", [])),
" ".join(query_semantics.get("keywords", [])),
" ".join(query_semantics.get("operations", [])),
" ".join(query_semantics.get("capabilities", [])),
" ".join(query_semantics.get("datatypes", [])),
]
tokens = self._tokenize(" ".join(str(part) for part in parts if part))
return {token for token in tokens if token not in self.GENERIC_TOKENS and len(token) >= 3}
def _score_tool(
self,
*,
server_name: str,
server_tokens: set[str],
server_overlap: set[str],
tool: dict,
query_tokens: set[str],
query_operations: set[str],
query_capabilities: set[str],
query_datatypes: set[str],
query_stages: set[str],
tool_graph: ToolGraph,
) -> tuple[float, list[str]]:
text = " ".join(
[
server_name,
str(tool.get("name", "")),
str(tool.get("description", "")),
" ".join(tool.get("keywords", [])),
" ".join(tool.get("operations", [])),
" ".join(tool.get("capabilities", [])),
" ".join(tool.get("consumes", [])),
" ".join(tool.get("produces", [])),
]
).lower()
tool_tokens = self._tokenize(text)
score = 0.0
reasons = []
tool_overlap = query_tokens & tool_tokens
if tool_overlap:
score += min(len(tool_overlap), 10) * 0.8
reasons.append("tool_token_overlap=" + ",".join(sorted(list(tool_overlap))[:6]))
if server_overlap:
score += min(len(server_overlap), 6) * 0.35
reasons.append("server_token_overlap=" + ",".join(sorted(list(server_overlap))[:4]))
tool_operations = set(tool.get("operations", []))
operation_overlap = query_operations & tool_operations
if operation_overlap:
score += len(operation_overlap) * 4.0
score += len(operation_overlap) * (1.5 / max(1, len(tool_operations)))
reasons.append("operation_overlap=" + ",".join(sorted(operation_overlap)))
if len(tool_operations) > 3:
score -= min(len(tool_operations) - 3, 6) * 0.35
reasons.append("broad_operation_penalty")
capability_overlap = query_capabilities & set(tool.get("capabilities", []))
if capability_overlap:
score += len(capability_overlap) * 2.5
reasons.append("capability_overlap=" + ",".join(sorted(capability_overlap)))
tool_datatypes = set(tool.get("consumes", [])) | set(tool.get("produces", []))
datatype_overlap = query_datatypes & tool_datatypes
if datatype_overlap:
score += len(datatype_overlap) * 2.0
reasons.append("datatype_overlap=" + ",".join(sorted(datatype_overlap)))
stage = tool.get("stage")
if stage and stage in query_stages:
score += 1.0
reasons.append("stage_overlap=" + stage)
graph_bonus = self._graph_relevance_bonus(tool, query_operations, query_capabilities, query_datatypes, tool_graph)
if graph_bonus:
score += graph_bonus
reasons.append("graph_neighbor_overlap")
if not reasons and query_tokens & server_tokens:
score += 0.5
reasons.append("weak_server_match")
return score, reasons
def _graph_relevance_bonus(
self,
tool: dict,
query_operations: set[str],
query_capabilities: set[str],
query_datatypes: set[str],
tool_graph: ToolGraph,
) -> float:
node_id = tool.get("node_id")
if not node_id:
return 0.0
bonus = 0.0
for edge in tool_graph.adjacency.get(node_id, []):
target = tool_graph.nodes.get(edge.get("target"), {})
label = str(target.get("label", "")).lower()
node_type = target.get("type")
if node_type == "operation" and label in query_operations:
bonus += 2.0
elif node_type == "capability" and label in query_capabilities:
bonus += 1.5
elif node_type == "datatype" and label in query_datatypes:
bonus += 1.0
return min(bonus, 5.0)
def _compact_subgraph(self, candidates: list[dict], tool_graph: ToolGraph) -> dict:
nodes = {}
edges = []
candidate_node_ids = {item.get("node_id") for item in candidates if item.get("node_id")}
server_names = {item.get("server") for item in candidates if item.get("server")}
for server_name in server_names:
server_node = f"server:{server_name}"
if server_node in tool_graph.nodes:
nodes[server_node] = tool_graph.nodes[server_node]
for node_id in candidate_node_ids:
if node_id in tool_graph.nodes:
nodes[node_id] = tool_graph.nodes[node_id]
for edge in tool_graph.adjacency.get(node_id, []):
target = edge.get("target")
if target in tool_graph.nodes and tool_graph.nodes[target].get("type") in {
"operation",
"capability",
"datatype",
"stage",
"constraint",
}:
nodes[target] = tool_graph.nodes[target]
edges.append(edge)
return {"nodes": list(nodes.values()), "edges": edges[:120]}
def _expand_tool_context(self, tool: dict, tool_graph: ToolGraph) -> list[dict]:
node_id = tool.get("node_id")
if not node_id:
return []
context = []
for edge in tool_graph.adjacency.get(node_id, [])[:20]:
target = tool_graph.nodes.get(edge.get("target"), {})
if target:
context.append({"edge": edge.get("type"), "node_type": target.get("type"), "label": target.get("label")})
return context
def _operation_path_hint(self, candidates: list[dict]) -> list[str]:
operations = []
ordered = sorted(
candidates,
key=lambda item: (self.STAGE_RANK.get(item.get("stage"), 9), -float(item.get("score", 0.0))),
)
for item in ordered:
for operation in item.get("operations", []):
if operation not in operations:
operations.append(operation)
return operations[:8]
def _format_context(self, candidates: list[dict], compact: dict) -> str:
lines = [
f"Retrieved compact MCP graph: {len(compact.get('nodes', []))} nodes, {len(compact.get('edges', []))} edges",
"Top candidate tool bindings:",
]
for item in candidates[:12]:
lines.append(
"- "
+ f"{item['server']}.{item.get('tool')} score={item['score']} "
+ f"stage={item.get('stage')} ops={','.join(item.get('operations', [])[:4])}; {item.get('reason', '')}"
)
return "\n".join(lines)
def _server_text(self, server_name: str, server_entry: dict, payload: dict) -> str:
semantics = payload.get("semantics", {})
return " ".join(
[
server_name,
str(server_entry.get("summary", "")),
" ".join(server_entry.get("keywords", [])),
" ".join(semantics.get("keywords", [])),
" ".join(semantics.get("operations", [])),
" ".join(semantics.get("capabilities", [])),
" ".join(semantics.get("datatypes", [])),
]
).lower()
def _tokenize(self, text: str) -> set[str]:
normalized = (text or "").lower().replace("-", " ").replace("/", " ").replace("_", " ")
return {token for token in self.TOKEN_PATTERN.findall(normalized) if token}