diff --git a/src/fastmcp/tools/tool_transform.py b/src/fastmcp/tools/tool_transform.py index 063c32361..56e0bfd82 100644 --- a/src/fastmcp/tools/tool_transform.py +++ b/src/fastmcp/tools/tool_transform.py @@ -841,12 +841,30 @@ class TransformedTool(Tool): if "default" in param_schema: final_required.discard(param_name) - return { + # Merge $defs from both schemas, with override taking precedence + merged_defs = base_schema.get("$defs", {}).copy() + override_defs = override_schema.get("$defs", {}) + + for def_name, def_schema in override_defs.items(): + if def_name in merged_defs: + base_def = merged_defs[def_name].copy() + base_def.update(def_schema) + merged_defs[def_name] = base_def + else: + merged_defs[def_name] = def_schema.copy() + + result = { "type": "object", "properties": merged_props, "required": list(final_required), } + if merged_defs: + result["$defs"] = merged_defs + result = compress_schema(result, prune_defs=True) + + return result + @staticmethod def _function_has_kwargs(fn: Callable[..., Any]) -> bool: """Check if function accepts **kwargs. diff --git a/tests/tools/test_tool_transform.py b/tests/tools/test_tool_transform.py index f0b21bded..d01cf666d 100644 --- a/tests/tools/test_tool_transform.py +++ b/tests/tools/test_tool_transform.py @@ -4,6 +4,7 @@ from typing import Annotated, Any import pytest from dirty_equals import IsList +from inline_snapshot import snapshot from mcp.types import TextContent from pydantic import BaseModel, Field, TypeAdapter from typing_extensions import TypedDict @@ -1055,38 +1056,6 @@ class TestEnableDisable: await client.call_tool("new_add", {"x": 1, "y": 2}) -def test_arg_transform_examples_in_schema(add_tool): - # Simple example - new_tool = Tool.from_tool( - add_tool, - transform_args={ - "old_x": ArgTransform(examples=[1, 2, 3]), - }, - ) - prop = get_property(new_tool, "old_x") - assert prop["examples"] == [1, 2, 3] - - # Nested example (e.g., for array type) - new_tool2 = Tool.from_tool( - add_tool, - transform_args={ - "old_x": ArgTransform(examples=[["a", "b"], ["c", "d"]]), - }, - ) - prop2 = get_property(new_tool2, "old_x") - assert prop2["examples"] == [["a", "b"], ["c", "d"]] - - # If not set, should not be present - new_tool3 = Tool.from_tool( - add_tool, - transform_args={ - "old_x": ArgTransform(), - }, - ) - prop3 = get_property(new_tool3, "old_x") - assert "examples" not in prop3 - - class TestTransformToolOutputSchema: """Test output schema handling in transformed tools.""" @@ -1502,3 +1471,239 @@ def test_tool_transform_config_removes_meta(sample_tool): config = ToolTransformConfig(name="config_tool", meta=None) transformed = config.apply(sample_tool) assert transformed.meta is None + + +class TestInputSchema: + """Test schema definition handling and reference finding.""" + + def test_arg_transform_examples_in_schema(self, add_tool: Tool): + # Simple example + new_tool = Tool.from_tool( + add_tool, + transform_args={ + "old_x": ArgTransform(examples=[1, 2, 3]), + }, + ) + prop = get_property(new_tool, "old_x") + assert prop["examples"] == [1, 2, 3] + + # Nested example (e.g., for array type) + new_tool2 = Tool.from_tool( + add_tool, + transform_args={ + "old_x": ArgTransform(examples=[["a", "b"], ["c", "d"]]), + }, + ) + prop2 = get_property(new_tool2, "old_x") + assert prop2["examples"] == [["a", "b"], ["c", "d"]] + + # If not set, should not be present + new_tool3 = Tool.from_tool( + add_tool, + transform_args={ + "old_x": ArgTransform(), + }, + ) + prop3 = get_property(new_tool3, "old_x") + assert "examples" not in prop3 + + def test_merge_schema_with_defs_precedence(self): + """Test _merge_schema_with_precedence merges $defs correctly.""" + base_schema = { + "type": "object", + "properties": {"field1": {"$ref": "#/$defs/BaseType"}}, + "$defs": { + "BaseType": {"type": "string", "description": "base"}, + "SharedType": {"type": "integer", "minimum": 0}, + }, + } + + override_schema = { + "type": "object", + "properties": {"field2": {"$ref": "#/$defs/OverrideType"}}, + "$defs": { + "OverrideType": {"type": "boolean"}, + "SharedType": {"type": "integer", "minimum": 10}, # Override + }, + } + + transformed_tool_schema = TransformedTool._merge_schema_with_precedence( + base_schema, override_schema + ) + + # SharedType should no longer be present on the schema + assert "SharedType" not in transformed_tool_schema["$defs"] + + assert transformed_tool_schema == snapshot( + { + "type": "object", + "properties": { + "field1": {"$ref": "#/$defs/BaseType"}, + "field2": {"$ref": "#/$defs/OverrideType"}, + }, + "required": [], + "$defs": { + "BaseType": {"type": "string", "description": "base"}, + "OverrideType": {"type": "boolean"}, + }, + } + ) + + def test_transform_tool_with_complex_defs_pruning(self): + """Test that tool transformation properly prunes unused $defs.""" + + class UsedType(BaseModel): + value: str + + class UnusedType(BaseModel): + other: int + + @Tool.from_function + def complex_tool( + used_param: UsedType, unused_param: UnusedType | None = None + ) -> str: + return used_param.value + + # Transform to hide unused_param + transformed_tool: TransformedTool = Tool.from_tool( + complex_tool, transform_args={"unused_param": ArgTransform(hide=True)} + ) + + assert "UnusedType" not in transformed_tool.parameters["$defs"] + + assert transformed_tool.parameters == snapshot( + { + "type": "object", + "properties": { + "used_param": {"$ref": "#/$defs/UsedType", "title": "Used Param"} + }, + "required": ["used_param"], + "$defs": { + "UsedType": { + "properties": {"value": {"title": "Value", "type": "string"}}, + "required": ["value"], + "title": "UsedType", + "type": "object", + } + }, + } + ) + + def test_transform_with_custom_function_preserves_needed_defs(self): + """Test that custom transform functions preserve necessary $defs.""" + + class InputType(BaseModel): + data: str + + class OutputType(BaseModel): + result: str + + @Tool.from_function + def base_tool(input_data: InputType) -> OutputType: + return OutputType(result=input_data.data.upper()) + + async def transform_function(renamed_input: InputType): + return await forward(renamed_input=renamed_input) + + # Transform with custom function and argument rename + transformed = Tool.from_tool( + base_tool, + transform_fn=transform_function, + transform_args={"input_data": ArgTransform(name="renamed_input")}, + ) + + assert transformed.parameters == snapshot( + { + "type": "object", + "properties": { + "renamed_input": { + "$ref": "#/$defs/InputType", + "title": "Input Data", + } + }, + "required": ["renamed_input"], + "$defs": { + "InputType": { + "properties": {"data": {"title": "Data", "type": "string"}}, + "required": ["data"], + "title": "InputType", + "type": "object", + } + }, + } + ) + + def test_chained_transforms_preserve_correct_defs(self): + """Test that chained transformations preserve correct $defs.""" + + class TypeA(BaseModel): + a: str + + class TypeB(BaseModel): + b: int + + class TypeC(BaseModel): + c: bool + + @Tool.from_function + def base_tool(param_a: TypeA, param_b: TypeB, param_c: TypeC) -> str: + return f"{param_a.a}-{param_b.b}-{param_c.c}" + + # First transform: hide param_c + transform1 = Tool.from_tool( + base_tool, + transform_args={"param_c": ArgTransform(hide=True, default=TypeC(c=True))}, + ) + + assert transform1.parameters == snapshot( + { + "type": "object", + "properties": { + "param_a": {"$ref": "#/$defs/TypeA", "title": "Param A"}, + "param_b": {"$ref": "#/$defs/TypeB", "title": "Param B"}, + }, + "required": IsList("param_b", "param_a", check_order=False), + "$defs": { + "TypeA": { + "properties": {"a": {"title": "A", "type": "string"}}, + "required": ["a"], + "title": "TypeA", + "type": "object", + }, + "TypeB": { + "properties": {"b": {"title": "B", "type": "integer"}}, + "required": ["b"], + "title": "TypeB", + "type": "object", + }, + }, + } + ) + + assert "TypeA" in transform1.parameters["$defs"] + + # Second transform: hide param_b + transform2 = Tool.from_tool( + transform1, + transform_args={"param_b": ArgTransform(hide=True, default=TypeB(b=42))}, + ) + + assert "TypeB" not in transform2.parameters["$defs"] + + assert transform2.parameters == snapshot( + { + "type": "object", + "properties": { + "param_a": {"$ref": "#/$defs/TypeA", "title": "Param A"} + }, + "required": ["param_a"], + "$defs": { + "TypeA": { + "properties": {"a": {"title": "A", "type": "string"}}, + "required": ["a"], + "title": "TypeA", + "type": "object", + } + }, + } + )