File size: 10,379 Bytes
d1ce356
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
#!/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()