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}