diff --git a/src/fastmcp/cli/cli.py b/src/fastmcp/cli/cli.py index 85b9d9ff5..147d9dc2d 100644 --- a/src/fastmcp/cli/cli.py +++ b/src/fastmcp/cli/cli.py @@ -20,6 +20,7 @@ from rich.table import Table import fastmcp from fastmcp.cli import run as run_module from fastmcp.cli.client import call_command, discover_command, list_command +from fastmcp.cli.generate import generate_cli_command from fastmcp.cli.install import install_app from fastmcp.cli.tasks import tasks_app from fastmcp.utilities.cli import is_already_in_uv_subprocess, load_and_merge_config @@ -957,6 +958,7 @@ app.command(tasks_app) app.command(list_command, name="list") app.command(call_command, name="call") app.command(discover_command, name="discover") +app.command(generate_cli_command, name="generate-cli") if __name__ == "__main__": diff --git a/src/fastmcp/cli/generate.py b/src/fastmcp/cli/generate.py new file mode 100644 index 000000000..ac8a446c9 --- /dev/null +++ b/src/fastmcp/cli/generate.py @@ -0,0 +1,497 @@ +"""Generate a standalone CLI script from an MCP server's capabilities.""" + +import sys +import textwrap +from pathlib import Path +from typing import Annotated, Any + +import cyclopts +import mcp.types +from rich.console import Console + +from fastmcp.cli.client import _build_client, resolve_server_spec +from fastmcp.client.transports.base import ClientTransport +from fastmcp.client.transports.stdio import StdioTransport +from fastmcp.utilities.logging import get_logger + +logger = get_logger("cli.generate") +console = Console() + +# --------------------------------------------------------------------------- +# JSON Schema type → Python type string +# --------------------------------------------------------------------------- + +_JSON_SCHEMA_TYPE_MAP: dict[str, str] = { + "string": "str", + "integer": "int", + "number": "float", + "boolean": "bool", + "array": "list", + "object": "dict", + "null": "None", +} + + +def _schema_type_to_python(schema: dict[str, Any]) -> str: + """Convert a JSON Schema type fragment to a Python type annotation string.""" + if "anyOf" in schema: + parts = [_schema_type_to_python(s) for s in schema["anyOf"]] + return " | ".join(parts) + + schema_type = schema.get("type", "string") + if isinstance(schema_type, list): + return " | ".join(_JSON_SCHEMA_TYPE_MAP.get(t, "str") for t in schema_type) + + return _JSON_SCHEMA_TYPE_MAP.get(schema_type, "str") + + +# --------------------------------------------------------------------------- +# Transport serialization +# --------------------------------------------------------------------------- + + +def serialize_transport( + resolved: str | dict[str, Any] | ClientTransport, +) -> tuple[str, set[str]]: + """Serialize a resolved transport to a Python expression string. + + Returns ``(expression, extra_imports)`` where *extra_imports* is a set of + import lines needed by the expression. + """ + if isinstance(resolved, str): + return repr(resolved), set() + + if isinstance(resolved, StdioTransport): + parts = [f"command={resolved.command!r}", f"args={resolved.args!r}"] + if resolved.env: + parts.append(f"env={resolved.env!r}") + if resolved.cwd: + parts.append(f"cwd={resolved.cwd!r}") + expr = f"StdioTransport({', '.join(parts)})" + imports = {"from fastmcp.client.transports import StdioTransport"} + return expr, imports + + if isinstance(resolved, dict): + return repr(resolved), set() + + # Fallback: try repr + return repr(resolved), set() + + +# --------------------------------------------------------------------------- +# Per-tool code generation +# --------------------------------------------------------------------------- + + +def _tool_function_source(tool: mcp.types.Tool) -> str: + """Generate the source for a single ``@call_tool_app.command`` function.""" + schema = tool.inputSchema + properties: dict[str, Any] = schema.get("properties", {}) + required = set(schema.get("required", [])) + + # Build parameter lines + param_lines: list[str] = [] + call_args: list[str] = [] + + for prop_name, prop_schema in properties.items(): + py_type = _schema_type_to_python(prop_schema) + help_text = prop_schema.get("description", "") + is_required = prop_name in required + + # Escape quotes in help text + help_escaped = help_text.replace("\\", "\\\\").replace('"', '\\"') + + if is_required: + annotation = ( + f'Annotated[{py_type}, cyclopts.Parameter(help="{help_escaped}")]' + ) + param_lines.append(f" {prop_name}: {annotation},") + else: + default = prop_schema.get("default") + if default is not None: + annotation = ( + f'Annotated[{py_type}, cyclopts.Parameter(help="{help_escaped}")]' + ) + param_lines.append(f" {prop_name}: {annotation} = {default!r},") + else: + annotation = f'Annotated[{py_type} | None, cyclopts.Parameter(help="{help_escaped}")]' + param_lines.append(f" {prop_name}: {annotation} = None,") + + call_args.append(f"{prop_name!r}: {prop_name}") + + # Function name: use tool name directly (preserve underscores) + fn_name = tool.name.replace("-", "_") + + # Docstring + description = (tool.description or "").replace('"""', '\\"\\"\\"') + + lines = [] + lines.append("") + # Always pass name= to preserve the original tool name (cyclopts + # would otherwise convert underscores to hyphens). + lines.append(f"@call_tool_app.command(name={tool.name!r})") + lines.append(f"async def {fn_name}(") + + if param_lines: + lines.append(" *,") + lines.extend(param_lines) + + lines.append(") -> None:") + lines.append(f' """{description}"""') + dict_items = ", ".join(call_args) + lines.append(f" await _call_tool({tool.name!r}, {{{dict_items}}})") + lines.append("") + + return "\n".join(lines) + + +# --------------------------------------------------------------------------- +# Full script generation +# --------------------------------------------------------------------------- + + +def generate_cli_script( + server_name: str, + server_spec: str, + transport_code: str, + extra_imports: set[str], + tools: list[mcp.types.Tool], +) -> str: + """Generate the full CLI script source code.""" + + # Determine app name from server_name + app_name = server_name.replace(" ", "-").lower() + + # --- Header --- + lines: list[str] = [] + lines.append("#!/usr/bin/env python3") + lines.append(f'"""CLI for {server_name} MCP server.') + lines.append("") + lines.append(f"Generated by: fastmcp generate-cli {server_spec}") + lines.append('"""') + lines.append("") + + # --- Imports --- + lines.append("import json") + lines.append("import sys") + lines.append("from typing import Annotated") + lines.append("") + lines.append("import cyclopts") + lines.append("import mcp.types") + lines.append("from rich.console import Console") + lines.append("") + lines.append("from fastmcp import Client") + for imp in sorted(extra_imports): + lines.append(imp) + lines.append("") + + # --- Transport config --- + lines.append("# Modify this to change how the CLI connects to the MCP server.") + lines.append(f"CLIENT_SPEC = {transport_code}") + lines.append("") + + # --- App setup --- + lines.append( + f'app = cyclopts.App(name="{app_name}", help="CLI for {server_name} MCP server")' + ) + lines.append( + 'call_tool_app = cyclopts.App(name="call-tool", help="Call a tool on the server")' + ) + lines.append("app.command(call_tool_app)") + lines.append("") + lines.append("console = Console()") + lines.append("") + lines.append("") + + # --- Shared helpers --- + lines.append( + textwrap.dedent("""\ + # --------------------------------------------------------------------------- + # Helpers + # --------------------------------------------------------------------------- + + + def _print_tool_result(result): + if result.is_error: + for block in result.content: + if isinstance(block, mcp.types.TextContent): + console.print(f"[bold red]Error:[/bold red] {block.text}") + else: + console.print(f"[bold red]Error:[/bold red] {block}") + sys.exit(1) + + if result.structured_content is not None: + console.print_json(json.dumps(result.structured_content)) + return + + for block in result.content: + if isinstance(block, mcp.types.TextContent): + console.print(block.text) + elif isinstance(block, mcp.types.ImageContent): + size = len(block.data) * 3 // 4 + console.print(f"[dim][Image: {block.mimeType}, ~{size} bytes][/dim]") + elif isinstance(block, mcp.types.AudioContent): + size = len(block.data) * 3 // 4 + console.print(f"[dim][Audio: {block.mimeType}, ~{size} bytes][/dim]") + + + async def _call_tool(tool_name: str, arguments: dict) -> None: + filtered = {k: v for k, v in arguments.items() if v is not None} + async with Client(CLIENT_SPEC) as client: + result = await client.call_tool(tool_name, filtered, raise_on_error=False) + _print_tool_result(result) + if result.is_error: + sys.exit(1)""") + ) + lines.append("") + lines.append("") + + # --- Generic commands --- + lines.append( + textwrap.dedent("""\ + # --------------------------------------------------------------------------- + # List / read commands + # --------------------------------------------------------------------------- + + + @app.command + async def list_tools() -> None: + \"\"\"List available tools.\"\"\" + async with Client(CLIENT_SPEC) as client: + tools = await client.list_tools() + if not tools: + console.print("[dim]No tools found.[/dim]") + return + for tool in tools: + sig_parts = [] + props = tool.inputSchema.get("properties", {}) + required = set(tool.inputSchema.get("required", [])) + for pname, pschema in props.items(): + ptype = pschema.get("type", "string") + if pname in required: + sig_parts.append(f"{pname}: {ptype}") + else: + sig_parts.append(f"{pname}: {ptype} = ...") + sig = f"{tool.name}({', '.join(sig_parts)})" + console.print(f" [cyan]{sig}[/cyan]") + if tool.description: + console.print(f" {tool.description}") + console.print() + + + @app.command + async def list_resources() -> None: + \"\"\"List available resources.\"\"\" + async with Client(CLIENT_SPEC) as client: + resources = await client.list_resources() + if not resources: + console.print("[dim]No resources found.[/dim]") + return + for r in resources: + console.print(f" [cyan]{r.uri}[/cyan]") + desc_parts = [r.name or "", r.description or ""] + desc = " — ".join(p for p in desc_parts if p) + if desc: + console.print(f" {desc}") + console.print() + + + @app.command + async def read_resource(uri: Annotated[str, cyclopts.Parameter(help="Resource URI")]) -> None: + \"\"\"Read a resource by URI.\"\"\" + async with Client(CLIENT_SPEC) as client: + contents = await client.read_resource(uri) + for block in contents: + if isinstance(block, mcp.types.TextResourceContents): + console.print(block.text) + elif isinstance(block, mcp.types.BlobResourceContents): + size = len(block.blob) * 3 // 4 + console.print(f"[dim][Blob: {block.mimeType}, ~{size} bytes][/dim]") + + + @app.command + async def list_prompts() -> None: + \"\"\"List available prompts.\"\"\" + async with Client(CLIENT_SPEC) as client: + prompts = await client.list_prompts() + if not prompts: + console.print("[dim]No prompts found.[/dim]") + return + for p in prompts: + args_str = "" + if p.arguments: + parts = [a.name for a in p.arguments] + args_str = f"({', '.join(parts)})" + console.print(f" [cyan]{p.name}{args_str}[/cyan]") + if p.description: + console.print(f" {p.description}") + console.print() + + + @app.command + async def get_prompt( + name: Annotated[str, cyclopts.Parameter(help="Prompt name")], + *arguments: str, + ) -> None: + \"\"\"Get a prompt by name. Pass arguments as key=value pairs.\"\"\" + parsed: dict[str, str] = {} + for arg in arguments: + if "=" not in arg: + console.print(f"[bold red]Error:[/bold red] Invalid argument {arg!r} — expected key=value") + sys.exit(1) + key, value = arg.split("=", 1) + parsed[key] = value + + async with Client(CLIENT_SPEC) as client: + result = await client.get_prompt(name, parsed or None) + for msg in result.messages: + console.print(f"[bold]{msg.role}:[/bold]") + if isinstance(msg.content, mcp.types.TextContent): + console.print(f" {msg.content.text}") + elif isinstance(msg.content, mcp.types.ImageContent): + size = len(msg.content.data) * 3 // 4 + console.print(f" [dim][Image: {msg.content.mimeType}, ~{size} bytes][/dim]") + else: + console.print(f" {msg.content}") + console.print()""") + ) + lines.append("") + lines.append("") + + # --- Generated tool commands --- + if tools: + lines.append( + "# ---------------------------------------------------------------------------" + ) + lines.append("# Tool commands (generated from server schema)") + lines.append( + "# ---------------------------------------------------------------------------" + ) + + for tool in tools: + lines.append(_tool_function_source(tool)) + + # --- Entry point --- + lines.append("") + lines.append('if __name__ == "__main__":') + lines.append(" app()") + lines.append("") + + return "\n".join(lines) + + +# --------------------------------------------------------------------------- +# CLI command +# --------------------------------------------------------------------------- + + +async def generate_cli_command( + server_spec: Annotated[ + str, + cyclopts.Parameter( + help="Server URL, Python file, MCPConfig JSON, discovered name, or .js file", + ), + ], + output: Annotated[ + str, + cyclopts.Parameter( + help="Output file path (default: cli.py)", + ), + ] = "cli.py", + *, + force: Annotated[ + bool, + cyclopts.Parameter( + name=["-f", "--force"], + help="Overwrite output file if it exists", + ), + ] = False, + timeout: Annotated[ + float | None, + cyclopts.Parameter("--timeout", help="Connection timeout in seconds"), + ] = None, + auth: Annotated[ + str | None, + cyclopts.Parameter( + "--auth", + help="Auth method: 'oauth', a bearer token string, or 'none' to disable", + ), + ] = None, +) -> None: + """Generate a standalone CLI script from an MCP server. + + Connects to the server, reads its tools/resources/prompts, and writes + a Python script that can invoke them directly. + + Examples: + fastmcp generate-cli weather + fastmcp generate-cli weather my_cli.py + fastmcp generate-cli http://localhost:8000/mcp + fastmcp generate-cli server.py output.py -f + """ + output_path = Path(output) + if output_path.exists() and not force: + console.print( + f"[bold red]Error:[/bold red] [cyan]{output_path}[/cyan] already exists. " + f"Use [cyan]-f[/cyan] to overwrite." + ) + sys.exit(1) + + # Resolve the server spec to a transport + resolved = resolve_server_spec(server_spec) + transport_code, extra_imports = serialize_transport(resolved) + + # Derive a human-friendly server name from the spec + server_name = _derive_server_name(server_spec) + + # Connect and discover capabilities + client = _build_client(resolved, timeout=timeout, auth=auth) + + try: + async with client: + tools = await client.list_tools() + console.print( + f"[dim]Discovered {len(tools)} tool(s) from {server_spec}[/dim]" + ) + + except Exception as exc: + console.print(f"[bold red]Error:[/bold red] Could not connect: {exc}") + sys.exit(1) + + # Generate and write the script + script = generate_cli_script( + server_name=server_name, + server_spec=server_spec, + transport_code=transport_code, + extra_imports=extra_imports, + tools=tools, + ) + + output_path.write_text(script) + output_path.chmod(output_path.stat().st_mode | 0o111) # make executable + + console.print( + f"[green]✓[/green] Wrote [cyan]{output_path}[/cyan] " + f"with {len(tools)} tool command(s)" + ) + console.print(f"[dim]Run: python {output_path} --help[/dim]") + + +def _derive_server_name(server_spec: str) -> str: + """Derive a human-friendly name from a server spec.""" + # URL — use hostname + if server_spec.startswith(("http://", "https://")): + from urllib.parse import urlparse + + parsed = urlparse(server_spec) + return parsed.hostname or "server" + + # File path — use stem + if server_spec.endswith((".py", ".js", ".json")): + return Path(server_spec).stem + + # Bare name or qualified name + if ":" in server_spec: + return server_spec.split(":", 1)[1] + + return server_spec diff --git a/tests/cli/test_generate_cli.py b/tests/cli/test_generate_cli.py new file mode 100644 index 000000000..919dc8ee0 --- /dev/null +++ b/tests/cli/test_generate_cli.py @@ -0,0 +1,403 @@ +"""Tests for fastmcp generate-cli command.""" + +from pathlib import Path +from typing import Any +from unittest.mock import patch + +import mcp.types +import pytest + +from fastmcp import FastMCP +from fastmcp.cli import generate as generate_module +from fastmcp.cli.client import Client +from fastmcp.cli.generate import ( + _derive_server_name, + _schema_type_to_python, + _tool_function_source, + generate_cli_command, + generate_cli_script, + serialize_transport, +) +from fastmcp.client.transports.stdio import StdioTransport + +# --------------------------------------------------------------------------- +# _schema_type_to_python +# --------------------------------------------------------------------------- + + +class TestSchemaTypeToPython: + def test_string(self): + assert _schema_type_to_python({"type": "string"}) == "str" + + def test_integer(self): + assert _schema_type_to_python({"type": "integer"}) == "int" + + def test_number(self): + assert _schema_type_to_python({"type": "number"}) == "float" + + def test_boolean(self): + assert _schema_type_to_python({"type": "boolean"}) == "bool" + + def test_array(self): + assert _schema_type_to_python({"type": "array"}) == "list" + + def test_object(self): + assert _schema_type_to_python({"type": "object"}) == "dict" + + def test_null(self): + assert _schema_type_to_python({"type": "null"}) == "None" + + def test_unknown_defaults_to_str(self): + assert _schema_type_to_python({"type": "foobar"}) == "str" + + def test_missing_type_defaults_to_str(self): + assert _schema_type_to_python({}) == "str" + + def test_any_of(self): + result = _schema_type_to_python( + {"anyOf": [{"type": "string"}, {"type": "integer"}]} + ) + assert result == "str | int" + + def test_type_list(self): + result = _schema_type_to_python({"type": ["string", "null"]}) + assert result == "str | None" + + +# --------------------------------------------------------------------------- +# serialize_transport +# --------------------------------------------------------------------------- + + +class TestSerializeTransport: + def test_url_string(self): + code, imports = serialize_transport("http://localhost:8000/mcp") + assert code == "'http://localhost:8000/mcp'" + assert imports == set() + + def test_stdio_transport_basic(self): + transport = StdioTransport(command="fastmcp", args=["run", "server.py"]) + code, imports = serialize_transport(transport) + assert "StdioTransport" in code + assert "command='fastmcp'" in code + assert "args=['run', 'server.py']" in code + assert "from fastmcp.client.transports import StdioTransport" in imports + + def test_stdio_transport_with_env(self): + transport = StdioTransport( + command="python", args=["-m", "myserver"], env={"KEY": "val"} + ) + code, imports = serialize_transport(transport) + assert "env={'KEY': 'val'}" in code + + def test_dict_passthrough(self): + d: dict[str, Any] = {"mcpServers": {"test": {"url": "http://localhost"}}} + code, imports = serialize_transport(d) + assert "mcpServers" in code + assert imports == set() + + +# --------------------------------------------------------------------------- +# _tool_function_source +# --------------------------------------------------------------------------- + + +class TestToolFunctionSource: + def test_required_param(self): + tool = mcp.types.Tool( + name="greet", + inputSchema={ + "properties": {"name": {"type": "string", "description": "Who"}}, + "required": ["name"], + }, + ) + source = _tool_function_source(tool) + assert "async def greet(" in source + assert "name: Annotated[str" in source + assert "= None" not in source + assert "_call_tool('greet', {'name': name})" in source + + def test_optional_param(self): + tool = mcp.types.Tool( + name="search", + inputSchema={ + "properties": { + "query": {"type": "string", "description": "Search query"}, + "limit": {"type": "integer", "description": "Max results"}, + }, + "required": ["query"], + }, + ) + source = _tool_function_source(tool) + assert "query: Annotated[str" in source + assert "limit: Annotated[int | None" in source + assert "= None" in source + + def test_param_with_default(self): + tool = mcp.types.Tool( + name="fetch", + inputSchema={ + "properties": { + "url": {"type": "string", "description": "URL"}, + "timeout": { + "type": "integer", + "description": "Timeout", + "default": 30, + }, + }, + "required": ["url"], + }, + ) + source = _tool_function_source(tool) + assert "timeout: Annotated[int" in source + assert "= 30" in source + + def test_no_params(self): + tool = mcp.types.Tool( + name="ping", + inputSchema={"properties": {}}, + ) + source = _tool_function_source(tool) + assert "async def ping(" in source + assert "_call_tool('ping', {})" in source + + def test_preserves_underscores(self): + tool = mcp.types.Tool( + name="get_forecast", + inputSchema={ + "properties": {"city": {"type": "string"}}, + "required": ["city"], + }, + ) + source = _tool_function_source(tool) + assert "async def get_forecast(" in source + + def test_description_in_docstring(self): + tool = mcp.types.Tool( + name="greet", + description="Say hello to someone.", + inputSchema={ + "properties": {"name": {"type": "string"}}, + "required": ["name"], + }, + ) + source = _tool_function_source(tool) + assert '"""Say hello to someone."""' in source + + +# --------------------------------------------------------------------------- +# _derive_server_name +# --------------------------------------------------------------------------- + + +class TestDeriveServerName: + def test_bare_name(self): + assert _derive_server_name("weather") == "weather" + + def test_qualified_name(self): + assert _derive_server_name("cursor:weather") == "weather" + + def test_python_file(self): + assert _derive_server_name("server.py") == "server" + + def test_url(self): + assert _derive_server_name("http://localhost:8000/mcp") == "localhost" + + +# --------------------------------------------------------------------------- +# generate_cli_script — produces compilable Python +# --------------------------------------------------------------------------- + + +class TestGenerateCliScript: + def _make_tools(self) -> list[mcp.types.Tool]: + return [ + mcp.types.Tool( + name="greet", + description="Say hello", + inputSchema={ + "properties": { + "name": {"type": "string", "description": "Who to greet"}, + }, + "required": ["name"], + }, + ), + mcp.types.Tool( + name="add_numbers", + description="Add two numbers", + inputSchema={ + "properties": { + "a": {"type": "integer", "description": "First number"}, + "b": {"type": "integer", "description": "Second number"}, + }, + "required": ["a", "b"], + }, + ), + ] + + def test_compiles(self): + script = generate_cli_script( + server_name="test", + server_spec="test", + transport_code='"http://localhost:8000/mcp"', + extra_imports=set(), + tools=self._make_tools(), + ) + compile(script, "", "exec") + + def test_contains_tool_functions(self): + script = generate_cli_script( + server_name="test", + server_spec="test", + transport_code='"http://localhost:8000/mcp"', + extra_imports=set(), + tools=self._make_tools(), + ) + assert "async def greet(" in script + assert "async def add_numbers(" in script + + def test_contains_generic_commands(self): + script = generate_cli_script( + server_name="test", + server_spec="test", + transport_code='"http://localhost:8000/mcp"', + extra_imports=set(), + tools=[], + ) + assert "async def list_tools(" in script + assert "async def list_resources(" in script + assert "async def list_prompts(" in script + assert "async def read_resource(" in script + assert "async def get_prompt(" in script + + def test_embeds_transport(self): + script = generate_cli_script( + server_name="test", + server_spec="test", + transport_code="StdioTransport(command='fastmcp', args=['run', 'x.py'])", + extra_imports={"from fastmcp.client.transports import StdioTransport"}, + tools=[], + ) + assert "StdioTransport(command='fastmcp'" in script + assert "from fastmcp.client.transports import StdioTransport" in script + + def test_no_tools_still_valid(self): + script = generate_cli_script( + server_name="empty", + server_spec="empty", + transport_code='"http://localhost"', + extra_imports=set(), + tools=[], + ) + compile(script, "", "exec") + assert "call_tool_app" in script + + def test_compiles_with_stdio_transport(self): + transport = StdioTransport(command="fastmcp", args=["run", "server.py"]) + transport_code, extra_imports = serialize_transport(transport) + script = generate_cli_script( + server_name="test", + server_spec="server.py", + transport_code=transport_code, + extra_imports=extra_imports, + tools=self._make_tools(), + ) + compile(script, "", "exec") + + +# --------------------------------------------------------------------------- +# generate_cli_command — integration tests +# --------------------------------------------------------------------------- + + +def _build_test_server() -> FastMCP: + """Create a minimal FastMCP server for integration tests.""" + server = FastMCP("TestServer") + + @server.tool + def greet(name: str) -> str: + """Say hello to someone.""" + return f"Hello, {name}!" + + @server.tool + def add(a: int, b: int) -> int: + """Add two numbers.""" + return a + b + + @server.resource("test://greeting") + def greeting_resource() -> str: + """A static greeting resource.""" + return "Hello from resource!" + + @server.prompt + def ask(topic: str) -> str: + """Ask about a topic.""" + return f"Tell me about {topic}" + + return server + + +@pytest.fixture() +def _patch_client(): + """Patch resolve_server_spec and _build_client to use an in-process server.""" + server = _build_test_server() + + def fake_resolve(server_spec: Any, **kwargs: Any) -> str: + return "fake://server" + + def fake_build_client(resolved: Any, **kwargs: Any) -> Client: + return Client(server) + + with ( + patch.object(generate_module, "resolve_server_spec", side_effect=fake_resolve), + patch.object(generate_module, "_build_client", side_effect=fake_build_client), + ): + yield + + +class TestGenerateCliCommand: + @pytest.mark.usefixtures("_patch_client") + async def test_writes_file(self, tmp_path: Path): + output = tmp_path / "cli.py" + await generate_cli_command("test-server", str(output)) + assert output.exists() + content = output.read_text() + compile(content, str(output), "exec") + + @pytest.mark.usefixtures("_patch_client") + async def test_contains_tools(self, tmp_path: Path): + output = tmp_path / "cli.py" + await generate_cli_command("test-server", str(output)) + content = output.read_text() + assert "async def greet(" in content + assert "async def add(" in content + + @pytest.mark.usefixtures("_patch_client") + async def test_default_output_path( + self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch + ): + monkeypatch.chdir(tmp_path) + await generate_cli_command("test-server") + assert (tmp_path / "cli.py").exists() + + @pytest.mark.usefixtures("_patch_client") + async def test_error_if_exists(self, tmp_path: Path): + output = tmp_path / "cli.py" + output.write_text("existing") + with pytest.raises(SystemExit): + await generate_cli_command("test-server", str(output)) + + @pytest.mark.usefixtures("_patch_client") + async def test_force_overwrites(self, tmp_path: Path): + output = tmp_path / "cli.py" + output.write_text("existing") + await generate_cli_command("test-server", str(output), force=True) + content = output.read_text() + assert content != "existing" + assert "async def greet(" in content + + @pytest.mark.usefixtures("_patch_client") + async def test_file_is_executable(self, tmp_path: Path): + output = tmp_path / "cli.py" + await generate_cli_command("test-server", str(output)) + assert output.stat().st_mode & 0o111