Beyond_Prompt-based_Retrieval / Biomanus_upload /patch_generated_mcp_shims.py
czty's picture
Add files using upload-large-folder tool
b2c86fd verified
Raw
History Blame Contribute Delete
3.55 kB
#!/usr/bin/env python3
from __future__ import annotations
import argparse
import re
from pathlib import Path
PROJECT_ROOT = Path(__file__).resolve().parent
DEFAULT_MCP_ROOT = PROJECT_ROOT / "biomni_web" / "backend" / "data" / "mcp_generated"
SOURCE_RE = re.compile(r"SOURCE_SERVER\s*=\s*Path\((['\"])(.*?)\1\)")
SERVER_RE = re.compile(r"SERVER_NAME\s*=\s*(['\"])(.*?)\1")
SHIM_TEMPLATE = """#!/usr/bin/env python3
from __future__ import annotations
import ast
from pathlib import Path
from mcp.server.fastmcp import FastMCP
SOURCE_SERVER = Path({source_server_path!r})
LOCAL_SERVER = Path(__file__).with_name(SOURCE_SERVER.name)
SERVER_NAME = {server_name!r}
class _ShimMCP:
@staticmethod
def tool(*args, **kwargs):
if args and callable(args[0]) and len(args) == 1 and not kwargs:
return args[0]
def _decorator(fn):
return fn
return _decorator
def _resolve_source_server():
if LOCAL_SERVER.exists() and LOCAL_SERVER.name != Path(__file__).name:
return LOCAL_SERVER
return SOURCE_SERVER
def _load_functions():
source_server = _resolve_source_server()
code = source_server.read_text(encoding="utf-8")
tree = ast.parse(code, filename=str(source_server))
function_names = [n.name for n in tree.body if isinstance(n, ast.FunctionDef) and not n.name.startswith("_")]
namespace = {{
"__name__": "__mcp_source__",
"mcp": _ShimMCP(),
}}
exec(compile(code, str(source_server), "exec"), namespace, namespace)
loaded = []
for name in function_names:
fn = namespace.get(name)
if callable(fn):
loaded.append(fn)
return loaded
mcp = FastMCP(SERVER_NAME)
for _fn in _load_functions():
mcp.tool()(_fn)
if __name__ == "__main__":
mcp.run(transport="stdio")
"""
def extract_metadata(shim_path: Path) -> tuple[str, str] | None:
text = shim_path.read_text(encoding="utf-8")
source_match = SOURCE_RE.search(text)
server_match = SERVER_RE.search(text)
if not source_match or not server_match:
return None
return source_match.group(2), server_match.group(2)
def rewrite_shim(shim_path: Path, dry_run: bool = False) -> bool:
metadata = extract_metadata(shim_path)
if metadata is None:
return False
source_server_path, server_name = metadata
new_text = SHIM_TEMPLATE.format(source_server_path=source_server_path, server_name=server_name)
current = shim_path.read_text(encoding="utf-8")
if current == new_text:
return False
if not dry_run:
shim_path.write_text(new_text, encoding="utf-8")
return True
def main() -> None:
parser = argparse.ArgumentParser(
description="Rewrite generated MCP shim servers to prefer local patched raw servers."
)
parser.add_argument(
"--mcp-root",
default=str(DEFAULT_MCP_ROOT),
help="Root directory containing mcp_* server dirs.",
)
parser.add_argument("--dry-run", action="store_true")
args = parser.parse_args()
mcp_root = Path(args.mcp_root).resolve()
if not mcp_root.is_dir():
raise SystemExit(f"MCP root does not exist or is not a directory: {mcp_root}")
scanned = 0
changed = 0
for shim_path in sorted(mcp_root.glob("mcp_*/app/*_shim_server.py")):
scanned += 1
if rewrite_shim(shim_path, dry_run=args.dry_run):
changed += 1
print(f"Scanned {scanned} shim files.")
print(f"Rewrote {changed} shim files.")
if __name__ == "__main__":
main()