diff --git a/src/fastmcp/cli/generate.py b/src/fastmcp/cli/generate.py index 49f8ad2d9..2c9bb86fa 100644 --- a/src/fastmcp/cli/generate.py +++ b/src/fastmcp/cli/generate.py @@ -23,28 +23,89 @@ 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", -} +_SIMPLE_TYPES = {"string", "integer", "number", "boolean", "null"} -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") +def _is_simple_type(schema: dict[str, Any]) -> bool: + """Check if a schema represents a simple (non-complex) type.""" + schema_type = schema.get("type") if isinstance(schema_type, list): - return " | ".join(_JSON_SCHEMA_TYPE_MAP.get(t, "str") for t in schema_type) + # Union of types - simple only if all are simple + return all(t in _SIMPLE_TYPES for t in schema_type) + return schema_type in _SIMPLE_TYPES - return _JSON_SCHEMA_TYPE_MAP.get(schema_type, "str") + +def _is_simple_array(schema: dict[str, Any]) -> tuple[bool, str | None]: + """Check if schema is an array of simple types. + + Returns (is_simple_array, item_type_str). + """ + if schema.get("type") != "array": + return False, None + + items = schema.get("items", {}) + if not _is_simple_type(items): + return False, None + + # Map JSON Schema type to Python type + item_type = items.get("type", "string") + type_map = { + "string": "str", + "integer": "int", + "number": "float", + "boolean": "bool", + } + return True, type_map.get(item_type, "str") + + +def _schema_to_python_type(schema: dict[str, Any]) -> tuple[str, bool]: + """Convert a JSON Schema to a Python type annotation. + + Returns (type_annotation, needs_json_parsing). + """ + # Check for simple array first + is_simple_arr, item_type = _is_simple_array(schema) + if is_simple_arr: + return f"list[{item_type}]", False + + # Check for simple type + if _is_simple_type(schema): + schema_type = schema.get("type", "string") + if isinstance(schema_type, list): + # Union of simple types + type_map = { + "string": "str", + "integer": "int", + "number": "float", + "boolean": "bool", + "null": "None", + } + parts = [type_map.get(t, "str") for t in schema_type] + return " | ".join(parts), False + + type_map = { + "string": "str", + "integer": "int", + "number": "float", + "boolean": "bool", + "null": "None", + } + return type_map.get(schema_type, "str"), False + + # Complex type - needs JSON parsing + return "str", True + + +def _format_schema_for_help(schema: dict[str, Any]) -> str: + """Format a JSON schema for display in help text.""" + import json + + # Pretty print the schema, indented for help text + schema_str = json.dumps(schema, indent=2) + # Indent each line for help text alignment + lines = schema_str.split("\n") + indented = "\n ".join(lines) + return f"JSON Schema: {indented}" # --------------------------------------------------------------------------- @@ -99,19 +160,29 @@ def _tool_function_source(tool: mcp.types.Tool) -> str: properties: dict[str, Any] = schema.get("properties", {}) required = set(schema.get("required", [])) - # Build parameter lines + # Build parameter lines and track which need JSON parsing param_lines: list[str] = [] call_args: list[str] = [] + json_params: list[tuple[str, str]] = [] # (prop_name, safe_name) for prop_name, prop_schema in properties.items(): - py_type = _schema_type_to_python(prop_schema) + py_type, needs_json = _schema_to_python_type(prop_schema) help_text = prop_schema.get("description", "") is_required = prop_name in required safe_name = _to_python_identifier(prop_name) - # Escape quotes in help text - help_escaped = help_text.replace("\\", "\\\\").replace('"', '\\"') + # For complex types, add schema to help text + if needs_json: + schema_help = _format_schema_for_help(prop_schema) + help_text = f"{help_text}\\n{schema_help}" if help_text else schema_help + json_params.append((prop_name, safe_name)) + # Escape special characters in help text + help_escaped = ( + help_text.replace("\\", "\\\\").replace('"', '\\"').replace("\n", "\\n") + ) + + # Build parameter annotation if is_required: annotation = ( f'Annotated[{py_type}, cyclopts.Parameter(help="{help_escaped}")]' @@ -125,8 +196,12 @@ def _tool_function_source(tool: mcp.types.Tool) -> str: ) param_lines.append(f" {safe_name}: {annotation} = {default!r},") else: - annotation = f'Annotated[{py_type} | None, cyclopts.Parameter(help="{help_escaped}")]' - param_lines.append(f" {safe_name}: {annotation} = None,") + # For list types, default to empty list; others default to None + if py_type.startswith("list["): + param_lines.append(f" {safe_name}: {py_type} = [],") + else: + annotation = f'Annotated[{py_type} | None, cyclopts.Parameter(help="{help_escaped}")]' + param_lines.append(f" {safe_name}: {annotation} = None,") call_args.append(f"{prop_name!r}: {safe_name}") @@ -149,7 +224,26 @@ def _tool_function_source(tool: mcp.types.Tool) -> str: lines.append(") -> None:") lines.append(f" '''{description}'''") - dict_items = ", ".join(call_args) + + # Add JSON parsing for complex parameters + if json_params: + lines.append(" # Parse JSON parameters") + for _prop_name, safe_name in json_params: + lines.append( + f" {safe_name}_parsed = json.loads({safe_name}) if {safe_name} else None" + ) + lines.append("") + + # Build call arguments, using parsed versions for JSON params + call_arg_parts = [] + for prop_name, _ in properties.items(): + safe_name = _to_python_identifier(prop_name) + if any(pn == prop_name for pn, _ in json_params): + call_arg_parts.append(f"{prop_name!r}: {safe_name}_parsed") + else: + call_arg_parts.append(f"{prop_name!r}: {safe_name}") + + dict_items = ", ".join(call_arg_parts) lines.append(f" await _call_tool({tool.name!r}, {{{dict_items}}})") lines.append("") @@ -250,7 +344,12 @@ def generate_cli_script( async def _call_tool(tool_name: str, arguments: dict) -> None: - filtered = {k: v for k, v in arguments.items() if v is not None} + # Filter out None values and empty lists (defaults for optional array params) + filtered = { + k: v + for k, v in arguments.items() + if v is not None and (not isinstance(v, list) or len(v) > 0) + } async with Client(CLIENT_SPEC) as client: result = await client.call_tool(tool_name, filtered, raise_on_error=False) _print_tool_result(result) diff --git a/tests/cli/test_generate_cli.py b/tests/cli/test_generate_cli.py index dfae49a08..f86e2460c 100644 --- a/tests/cli/test_generate_cli.py +++ b/tests/cli/test_generate_cli.py @@ -13,7 +13,7 @@ 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, + _schema_to_python_type, _to_python_identifier, _tool_function_source, generate_cli_command, @@ -23,47 +23,61 @@ from fastmcp.cli.generate import ( from fastmcp.client.transports.stdio import StdioTransport # --------------------------------------------------------------------------- -# _schema_type_to_python +# _schema_to_python_type # --------------------------------------------------------------------------- -class TestSchemaTypeToPython: - def test_string(self): - assert _schema_type_to_python({"type": "string"}) == "str" +class TestSchemaToPythonType: + def test_simple_string(self): + py_type, needs_json = _schema_to_python_type({"type": "string"}) + assert py_type == "str" + assert needs_json is False - def test_integer(self): - assert _schema_type_to_python({"type": "integer"}) == "int" + def test_simple_integer(self): + py_type, needs_json = _schema_to_python_type({"type": "integer"}) + assert py_type == "int" + assert needs_json is False - def test_number(self): - assert _schema_type_to_python({"type": "number"}) == "float" + def test_simple_number(self): + py_type, needs_json = _schema_to_python_type({"type": "number"}) + assert py_type == "float" + assert needs_json is False - def test_boolean(self): - assert _schema_type_to_python({"type": "boolean"}) == "bool" + def test_simple_boolean(self): + py_type, needs_json = _schema_to_python_type({"type": "boolean"}) + assert py_type == "bool" + assert needs_json is False - 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"}]} + def test_array_of_strings(self): + py_type, needs_json = _schema_to_python_type( + {"type": "array", "items": {"type": "string"}} ) - assert result == "str | int" + assert py_type == "list[str]" + assert needs_json is False - def test_type_list(self): - result = _schema_type_to_python({"type": ["string", "null"]}) - assert result == "str | None" + def test_array_of_integers(self): + py_type, needs_json = _schema_to_python_type( + {"type": "array", "items": {"type": "integer"}} + ) + assert py_type == "list[int]" + assert needs_json is False + + def test_complex_object(self): + py_type, needs_json = _schema_to_python_type({"type": "object"}) + assert py_type == "str" + assert needs_json is True + + def test_complex_nested_array(self): + py_type, needs_json = _schema_to_python_type( + {"type": "array", "items": {"type": "object"}} + ) + assert py_type == "str" + assert needs_json is True + + def test_union_of_simple_types(self): + py_type, needs_json = _schema_to_python_type({"type": ["string", "null"]}) + assert py_type == "str | None" + assert needs_json is False # --------------------------------------------------------------------------- @@ -247,6 +261,79 @@ class TestToolFunctionSource: # Generated code should compile compile(source, "", "exec") + def test_array_of_strings_parameter(self): + tool = mcp.types.Tool( + name="tag_items", + description="Tag multiple items.", + inputSchema={ + "properties": { + "item_id": {"type": "string"}, + "tags": {"type": "array", "items": {"type": "string"}}, + }, + "required": ["item_id"], + }, + ) + source = _tool_function_source(tool) + # Should use list[str] type + assert "tags: list[str] = []" in source + # Should not have JSON parsing for simple arrays + assert "json.loads" not in source + compile(source, "", "exec") + + def test_complex_object_parameter(self): + tool = mcp.types.Tool( + name="create_user", + description="Create a user.", + inputSchema={ + "properties": { + "name": {"type": "string"}, + "metadata": { + "type": "object", + "properties": { + "role": {"type": "string"}, + "dept": {"type": "string"}, + }, + }, + }, + "required": ["name"], + }, + ) + source = _tool_function_source(tool) + # Should use str type for complex object + assert "metadata: Annotated[str | None" in source + # Should include JSON schema in help (with escaped quotes) + assert "JSON Schema:" in source + assert '\\"type\\": \\"object\\"' in source + # Should have JSON parsing + assert "metadata_parsed = json.loads(metadata) if metadata else None" in source + # Should use parsed version in call + assert "'metadata': metadata_parsed" in source + compile(source, "", "exec") + + def test_nested_array_parameter(self): + tool = mcp.types.Tool( + name="batch_process", + description="Process batches.", + inputSchema={ + "properties": { + "batches": { + "type": "array", + "items": { + "type": "object", + "properties": {"id": {"type": "string"}}, + }, + }, + }, + "required": ["batches"], + }, + ) + source = _tool_function_source(tool) + # Nested arrays need JSON parsing + assert "batches: Annotated[str" in source + assert "JSON Schema:" in source + assert "batches_parsed = json.loads(batches)" in source + compile(source, "", "exec") + # --------------------------------------------------------------------------- # _derive_server_name