Beyond_Prompt-based_Retrieval / Biomanus_upload /plan_query_with_mcp_graph.py
czty's picture
Add files using upload-large-folder tool
b2c86fd verified
Raw
History Blame Contribute Delete
10.4 kB
#!/usr/bin/env python3
from __future__ import annotations
import argparse
import json
from datetime import datetime
from pathlib import Path
from build_generated_mcp_graph import (
DEFAULT_HELP_ROOT,
DEFAULT_MCP_ROOT,
DEFAULT_OUTPUT_ROOT,
build_graph,
load_server_catalog,
)
from biomni.graph import GraphRouter, ToolGraph, ToolSchemaExtractor
from biomni.model.query_rewriter import QueryRewriter
from biomni.model.retriever import ToolRetriever
PROJECT_ROOT = Path(__file__).resolve().parent
DEFAULT_PLAN_ROOT = PROJECT_ROOT / "query_graph_plans"
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description=(
"Read a user query, rewrite it into a retrieval-friendly prompt, "
"and generate a graph-based MCP subgraph call chain."
)
)
parser.add_argument("--query", default="", help="Natural-language task query.")
parser.add_argument("--query-file", default="", help="Optional text/markdown file containing the query.")
parser.add_argument(
"--graph-dir",
default="",
help="Optional prebuilt graph directory containing server_catalog.json.",
)
parser.add_argument(
"--mcp-root",
default=str(DEFAULT_MCP_ROOT),
help="MCP root for graph building when --graph-dir is not provided.",
)
parser.add_argument(
"--help-root",
default=str(DEFAULT_HELP_ROOT),
help="Help text root for graph building when --graph-dir is not provided.",
)
parser.add_argument(
"--graph-output-root",
default=str(DEFAULT_OUTPUT_ROOT),
help="Output root for graph building when --graph-dir is not provided.",
)
parser.add_argument(
"--output-root",
default=str(DEFAULT_PLAN_ROOT),
help="Directory to store rewritten prompt and graph routing artifacts.",
)
parser.add_argument("--run-name", default="", help="Optional run folder name.")
parser.add_argument(
"--top-k-servers",
type=int,
default=8,
help="How many servers to keep in the routed subgraph.",
)
parser.add_argument(
"--candidate-pool",
type=int,
default=24,
help="Candidate server pool size before truncating to top-k.",
)
parser.add_argument(
"--top-k-tools",
type=int,
default=12,
help="How many tools to keep per selected server.",
)
parser.add_argument(
"--executable-only",
action="store_true",
help="Only include servers with non-empty server_meta.command in routing.",
)
return parser.parse_args()
def resolve_query(args: argparse.Namespace) -> str:
parts: list[str] = []
if args.query.strip():
parts.append(args.query.strip())
if args.query_file:
query_path = Path(args.query_file).expanduser().resolve()
if not query_path.exists():
raise FileNotFoundError(f"Query file does not exist: {query_path}")
parts.append(query_path.read_text(encoding="utf-8").strip())
query = "\n\n".join(part for part in parts if part)
if not query:
raise ValueError("Please provide --query or --query-file.")
return query
def prepare_graph_catalog(args: argparse.Namespace) -> tuple[Path, list[dict]]:
if args.graph_dir:
graph_dir = Path(args.graph_dir).expanduser().resolve()
if not graph_dir.exists():
raise FileNotFoundError(f"Graph directory does not exist: {graph_dir}")
else:
graph_dir = build_graph(
mcp_root=Path(args.mcp_root).expanduser().resolve(),
help_root=Path(args.help_root).expanduser().resolve(),
output_root=Path(args.graph_output_root).expanduser().resolve(),
)
server_entries = load_server_catalog(graph_dir)
if args.executable_only:
server_entries = [
entry
for entry in server_entries
if entry.get("command_available") or entry.get("server_meta", {}).get("command")
]
if not server_entries:
raise RuntimeError("No MCP server entries available for graph routing.")
return graph_dir, server_entries
def build_rewritten_prompt(query_context: dict, graph_dir: Path, server_count: int) -> str:
subtasks = query_context.get("subtasks", [])
subtasks_text = "\n".join(f"- {step}" for step in subtasks) if subtasks else "- No explicit subtasks inferred."
categories = ", ".join(query_context.get("categories", [])) or "general_biomedical_analysis"
server_hints = ", ".join(query_context.get("server_hints", [])) or "none"
return (
"You are planning a graph-guided MCP workflow from a user query.\n\n"
f"Original query:\n{query_context.get('original_query', '').strip()}\n\n"
f"Task summary:\n{query_context.get('task_summary', '').strip()}\n\n"
f"Rewritten retrieval query:\n{query_context.get('retrieval_query', '').strip()}\n\n"
f"Inferred categories: {categories}\n"
f"Server hints: {server_hints}\n\n"
"Planned subtasks:\n"
f"{subtasks_text}\n\n"
"Graph constraints:\n"
f"- Graph directory: {graph_dir}\n"
f"- Indexed server count: {server_count}\n"
"- Route through ToolGraph/GraphRouter to get a minimal MCP subgraph call chain.\n"
)
def format_call_chain(route_result: dict) -> str:
lines = ["# Graph Subgraph Call Chain", ""]
execution_plan = route_result.get("execution_plan", [])
workflow_subgraph = route_result.get("workflow_subgraph", [])
seed_nodes = route_result.get("seed_nodes", [])
selected_servers = route_result.get("selected_servers", [])
lines.append("## Seed Nodes")
if seed_nodes:
lines.extend(f"- {node_id}" for node_id in seed_nodes)
else:
lines.append("- (none)")
lines.append("")
lines.append("## Execution Plan")
if execution_plan:
for index, step in enumerate(execution_plan, start=1):
tools = ", ".join(step.get("tools", [])) or "(no tools)"
lines.append(
f"{index}. stage={step.get('stage', 'analysis')} | "
f"server={step.get('server', 'unknown')} | tools={tools}"
)
else:
lines.append("1. No staged plan was produced from the current query/graph.")
lines.append("")
lines.append("## Workflow Subgraph")
if workflow_subgraph:
for node in workflow_subgraph:
tools = ", ".join(node.get("tools", [])) or "(no tools)"
lines.append(
f"- server={node.get('server')} | category={node.get('category')} "
f"| score={node.get('score', 0):.2f} | tools={tools}"
)
else:
lines.append("- (none)")
lines.append("")
lines.append("## Server Selection Reasons")
if selected_servers:
for server in selected_servers:
reasons = server.get("graph_reason", [])
if reasons:
lines.append(f"- {server.get('name')}: {'; '.join(reasons)}")
else:
lines.append(f"- {server.get('name')}: (no explicit reasons)")
else:
lines.append("- (none)")
lines.append("")
return "\n".join(lines)
def make_run_dir(output_root: Path, run_name: str) -> Path:
if run_name.strip():
folder = run_name.strip()
else:
folder = datetime.now().strftime("query_plan_%Y%m%d_%H%M%S")
run_dir = output_root / folder
run_dir.mkdir(parents=True, exist_ok=True)
return run_dir
def main() -> None:
args = parse_args()
query = resolve_query(args)
graph_dir, server_entries = prepare_graph_catalog(args)
rewriter = QueryRewriter()
query_context = rewriter.rewrite(query, llm=None)
schema_extractor = ToolSchemaExtractor()
tool_graph = ToolGraph(schema_extractor=schema_extractor)
tool_graph.build_from_server_entries(server_entries)
router = GraphRouter(schema_extractor=schema_extractor)
route_result = router.route(
query_context=query_context,
tool_graph=tool_graph,
retriever=ToolRetriever(),
llm=None,
top_k_servers=args.top_k_servers,
candidate_pool=args.candidate_pool,
top_k_tools=args.top_k_tools,
)
rewritten_prompt = build_rewritten_prompt(
query_context=query_context,
graph_dir=graph_dir,
server_count=len(server_entries),
)
call_chain_markdown = format_call_chain(route_result)
output_root = Path(args.output_root).expanduser().resolve()
run_dir = make_run_dir(output_root, args.run_name)
rewritten_prompt_path = run_dir / "rewritten_prompt.txt"
query_context_path = run_dir / "query_context.json"
route_result_path = run_dir / "graph_route_result.json"
call_chain_path = run_dir / "graph_subgraph_call_chain.md"
run_meta_path = run_dir / "run_metadata.json"
rewritten_prompt_path.write_text(rewritten_prompt, encoding="utf-8")
query_context_path.write_text(json.dumps(query_context, ensure_ascii=False, indent=2), encoding="utf-8")
route_result_path.write_text(json.dumps(route_result, ensure_ascii=False, indent=2), encoding="utf-8")
call_chain_path.write_text(call_chain_markdown, encoding="utf-8")
run_meta_path.write_text(
json.dumps(
{
"query": query,
"graph_dir": str(graph_dir),
"server_count_for_routing": len(server_entries),
"selected_server_count": len(route_result.get("selected_servers", [])),
"seed_nodes_count": len(route_result.get("seed_nodes", [])),
"top_k_servers": args.top_k_servers,
"candidate_pool": args.candidate_pool,
"top_k_tools": args.top_k_tools,
"executable_only": args.executable_only,
},
ensure_ascii=False,
indent=2,
),
encoding="utf-8",
)
print("Query graph planning completed.")
print(f"Run dir: {run_dir}")
print(f"Rewritten prompt: {rewritten_prompt_path}")
print(f"Query context: {query_context_path}")
print(f"Graph route result: {route_result_path}")
print(f"Subgraph call chain: {call_chain_path}")
print(f"Run metadata: {run_meta_path}")
if __name__ == "__main__":
main()