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