diff --git a/src/fastmcp/prompts/prompt.py b/src/fastmcp/prompts/prompt.py index 4d2bc74e5..85eb3bc9d 100644 --- a/src/fastmcp/prompts/prompt.py +++ b/src/fastmcp/prompts/prompt.py @@ -13,7 +13,7 @@ from mcp.types import PromptArgument as MCPPromptArgument from pydantic import BaseModel, BeforeValidator, Field, TypeAdapter, validate_call from fastmcp.server.dependencies import get_context -from fastmcp.utilities.json_schema import prune_params +from fastmcp.utilities.json_schema import compress_schema from fastmcp.utilities.logging import get_logger from fastmcp.utilities.types import ( _convert_set_defaults, @@ -115,7 +115,11 @@ class Prompt(BaseModel): context_kwarg = find_kwarg_by_type(fn, kwarg_type=Context) if context_kwarg: - parameters = prune_params(parameters, params=[context_kwarg]) + prune_params = [context_kwarg] + else: + prune_params = None + + parameters = compress_schema(parameters, prune_params=prune_params) # Convert parameters to PromptArguments arguments: list[PromptArgument] = [] diff --git a/src/fastmcp/resources/template.py b/src/fastmcp/resources/template.py index 1335a2559..eaf0cfe7b 100644 --- a/src/fastmcp/resources/template.py +++ b/src/fastmcp/resources/template.py @@ -21,6 +21,7 @@ from pydantic import ( from fastmcp.resources.types import FunctionResource, Resource from fastmcp.server.dependencies import get_context +from fastmcp.utilities.json_schema import compress_schema from fastmcp.utilities.types import ( _convert_set_defaults, find_kwarg_by_type, @@ -150,6 +151,10 @@ class ResourceTemplate(BaseModel): # Get schema from TypeAdapter - will fail if function isn't properly typed parameters = TypeAdapter(fn).json_schema() + # compress the schema + prune_params = [context_kwarg] if context_kwarg else None + parameters = compress_schema(parameters, prune_params=prune_params) + # ensure the arguments are properly cast fn = validate_call(fn) diff --git a/src/fastmcp/tools/tool.py b/src/fastmcp/tools/tool.py index aa7c6bf63..73b84d76f 100644 --- a/src/fastmcp/tools/tool.py +++ b/src/fastmcp/tools/tool.py @@ -12,7 +12,7 @@ from pydantic import BaseModel, BeforeValidator, Field import fastmcp from fastmcp.server.dependencies import get_context -from fastmcp.utilities.json_schema import prune_params +from fastmcp.utilities.json_schema import compress_schema from fastmcp.utilities.logging import get_logger from fastmcp.utilities.types import ( Image, @@ -81,7 +81,11 @@ class Tool(BaseModel): context_kwarg = find_kwarg_by_type(fn, kwarg_type=Context) if context_kwarg: - schema = prune_params(schema, params=[context_kwarg]) + prune_params = [context_kwarg] + else: + prune_params = None + + schema = compress_schema(schema, prune_params=prune_params) return cls( fn=fn, diff --git a/tests/utilities/test_json_schema.py b/tests/utilities/test_json_schema.py index 7ae684523..cc9cc156c 100644 --- a/tests/utilities/test_json_schema.py +++ b/tests/utilities/test_json_schema.py @@ -1,110 +1,246 @@ -from fastmcp.utilities.json_schema import _prune_param, prune_params +from fastmcp.utilities.json_schema import ( + _prune_additional_properties, + _prune_param, + _prune_unused_defs, + compress_schema, +) -def test_prune_param_nonexistent(): - """Test pruning a parameter that doesn't exist.""" - schema = {"properties": {"foo": {"type": "string"}}} - result = _prune_param(schema, "bar") - assert result == schema # Schema should be unchanged +class TestPruneParam: + """Tests for the _prune_param function.""" + + def test_nonexistent(self): + """Test pruning a parameter that doesn't exist.""" + schema = {"properties": {"foo": {"type": "string"}}} + result = _prune_param(schema, "bar") + assert result == schema # Schema should be unchanged + + def test_exists(self): + """Test pruning a parameter that exists.""" + schema = {"properties": {"foo": {"type": "string"}, "bar": {"type": "integer"}}} + result = _prune_param(schema, "bar") + assert result["properties"] == {"foo": {"type": "string"}} + + def test_last_property(self): + """Test pruning the only/last parameter, should leave empty properties object.""" + schema = {"properties": {"foo": {"type": "string"}}} + result = _prune_param(schema, "foo") + assert "properties" in result + assert result["properties"] == {} + + def test_from_required(self): + """Test pruning a parameter that's in the required list.""" + schema = { + "properties": {"foo": {"type": "string"}, "bar": {"type": "integer"}}, + "required": ["foo", "bar"], + } + result = _prune_param(schema, "bar") + assert result["required"] == ["foo"] + + def test_last_required(self): + """Test pruning the last required parameter, should remove required field.""" + schema = { + "properties": {"foo": {"type": "string"}, "bar": {"type": "integer"}}, + "required": ["foo"], + } + result = _prune_param(schema, "foo") + assert "required" not in result -def test_prune_param_exists(): - """Test pruning a parameter that exists.""" - schema = {"properties": {"foo": {"type": "string"}, "bar": {"type": "integer"}}} - result = _prune_param(schema, "bar") - assert result["properties"] == {"foo": {"type": "string"}} +class TestPruneUnusedDefs: + """Tests for the _prune_unused_defs function.""" - -def test_prune_param_last_property(): - """Test pruning the only/last parameter, should leave empty properties object.""" - schema = {"properties": {"foo": {"type": "string"}}} - result = _prune_param(schema, "foo") - assert "properties" in result - assert result["properties"] == {} - - -def test_prune_param_from_required(): - """Test pruning a parameter that's in the required list.""" - schema = { - "properties": {"foo": {"type": "string"}, "bar": {"type": "integer"}}, - "required": ["foo", "bar"], - } - result = _prune_param(schema, "bar") - assert result["required"] == ["foo"] - - -def test_prune_param_last_required(): - """Test pruning the last required parameter, should remove required field.""" - schema = { - "properties": {"foo": {"type": "string"}, "bar": {"type": "integer"}}, - "required": ["foo"], - } - result = _prune_param(schema, "foo") - assert "required" not in result - - -def test_prune_param_with_refs(): - """Test pruning a parameter that has references in $defs.""" - schema = { - "properties": { - "foo": {"$ref": "#/$defs/foo_def"}, - "bar": {"$ref": "#/$defs/bar_def"}, - }, - "$defs": { - "foo_def": {"type": "string"}, - "bar_def": {"type": "integer"}, - }, - } - result = _prune_param(schema, "bar") - assert "bar_def" not in result["$defs"] - assert "foo_def" in result["$defs"] - - -def test_prune_param_all_refs(): - """Test pruning all parameters with refs, should remove $defs.""" - schema = { - "properties": { - "foo": {"$ref": "#/$defs/foo_def"}, - }, - "$defs": { - "foo_def": {"type": "string"}, - }, - } - result = _prune_param(schema, "foo") - assert "$defs" not in result - - -def test_prune_params_multiple(): - """Test pruning multiple parameters at once.""" - schema = { - "properties": { - "foo": {"type": "string"}, - "bar": {"type": "integer"}, - "baz": {"type": "boolean"}, - }, - "required": ["foo", "bar"], - } - result = prune_params(schema, ["foo", "baz"]) - assert result["properties"] == {"bar": {"type": "integer"}} - assert result["required"] == ["bar"] - - -def test_prune_params_nested_refs(): - """Test pruning with nested references.""" - schema = { - "properties": { - "foo": { - "type": "object", - "properties": {"nested": {"$ref": "#/$defs/nested_def"}}, + def test_removes_unreferenced_defs(self): + """Test that unreferenced definitions are removed.""" + schema = { + "properties": { + "foo": {"$ref": "#/$defs/foo_def"}, }, - "bar": {"$ref": "#/$defs/bar_def"}, - }, - "$defs": { - "nested_def": {"type": "string"}, - "bar_def": {"type": "integer"}, - }, - } - # Removing foo should keep nested_def as it's not referenced anymore - result = _prune_param(schema, "foo") - assert "nested_def" not in result["$defs"] - assert "bar_def" in result["$defs"] + "$defs": { + "foo_def": {"type": "string"}, + "unused_def": {"type": "integer"}, + }, + } + result = _prune_unused_defs(schema) + assert "foo_def" in result["$defs"] + assert "unused_def" not in result["$defs"] + + def test_nested_references_kept(self): + """Test that definitions referenced via nesting are kept.""" + schema = { + "properties": { + "foo": {"$ref": "#/$defs/foo_def"}, + }, + "$defs": { + "foo_def": { + "type": "object", + "properties": {"nested": {"$ref": "#/$defs/nested_def"}}, + }, + "nested_def": {"type": "string"}, + "unused_def": {"type": "integer"}, + }, + } + result = _prune_unused_defs(schema) + assert "foo_def" in result["$defs"] + assert "nested_def" in result["$defs"] + assert "unused_def" not in result["$defs"] + + def test_array_references_kept(self): + """Test that definitions referenced in array items are kept.""" + schema = { + "properties": { + "items": {"type": "array", "items": {"$ref": "#/$defs/item_def"}}, + }, + "$defs": { + "item_def": {"type": "string"}, + "unused_def": {"type": "integer"}, + }, + } + result = _prune_unused_defs(schema) + assert "item_def" in result["$defs"] + assert "unused_def" not in result["$defs"] + + def test_removes_defs_field_when_empty(self): + """Test that $defs field is removed when all definitions are unused.""" + schema = { + "properties": { + "foo": {"type": "string"}, + }, + "$defs": { + "unused_def": {"type": "integer"}, + }, + } + result = _prune_unused_defs(schema) + assert "$defs" not in result + + +class TestPruneAdditionalProperties: + """Tests for the _prune_additional_properties function.""" + + def test_removes_when_false(self): + """Test that additionalProperties is removed when it's false.""" + schema = { + "type": "object", + "properties": {"foo": {"type": "string"}}, + "additionalProperties": False, + } + result = _prune_additional_properties(schema) + assert "additionalProperties" not in result + + def test_keeps_when_true(self): + """Test that additionalProperties is kept when it's true.""" + schema = { + "type": "object", + "properties": {"foo": {"type": "string"}}, + "additionalProperties": True, + } + result = _prune_additional_properties(schema) + assert "additionalProperties" in result + assert result["additionalProperties"] is True + + def test_keeps_when_object(self): + """Test that additionalProperties is kept when it's an object schema.""" + schema = { + "type": "object", + "properties": {"foo": {"type": "string"}}, + "additionalProperties": {"type": "string"}, + } + result = _prune_additional_properties(schema) + assert "additionalProperties" in result + assert result["additionalProperties"] == {"type": "string"} + + +class TestCompressSchema: + """Tests for the compress_schema function.""" + + def test_prune_params(self): + """Test pruning parameters with compress_schema.""" + schema = { + "properties": { + "foo": {"type": "string"}, + "bar": {"type": "integer"}, + "baz": {"type": "boolean"}, + }, + "required": ["foo", "bar"], + } + result = compress_schema(schema, prune_params=["foo", "baz"]) + assert result["properties"] == {"bar": {"type": "integer"}} + assert result["required"] == ["bar"] + + def test_prune_defs(self): + """Test pruning unused definitions with compress_schema.""" + schema = { + "properties": { + "foo": {"$ref": "#/$defs/foo_def"}, + "bar": {"type": "integer"}, + }, + "$defs": { + "foo_def": {"type": "string"}, + "unused_def": {"type": "number"}, + }, + } + result = compress_schema(schema) + assert "foo_def" in result["$defs"] + assert "unused_def" not in result["$defs"] + + def test_disable_prune_defs(self): + """Test disabling pruning of unused definitions.""" + schema = { + "properties": { + "foo": {"$ref": "#/$defs/foo_def"}, + "bar": {"type": "integer"}, + }, + "$defs": { + "foo_def": {"type": "string"}, + "unused_def": {"type": "number"}, + }, + } + result = compress_schema(schema, prune_defs=False) + assert "foo_def" in result["$defs"] + assert "unused_def" in result["$defs"] + + def test_pruning_additional_properties(self): + """Test pruning additionalProperties when False.""" + schema = { + "type": "object", + "properties": {"foo": {"type": "string"}}, + "additionalProperties": False, + } + result = compress_schema(schema) + assert "additionalProperties" not in result + + def test_disable_pruning_additional_properties(self): + """Test disabling pruning of additionalProperties.""" + schema = { + "type": "object", + "properties": {"foo": {"type": "string"}}, + "additionalProperties": False, + } + result = compress_schema(schema, prune_additional_properties=False) + assert "additionalProperties" in result + assert result["additionalProperties"] is False + + def test_combined_operations(self): + """Test all pruning operations together.""" + schema = { + "type": "object", + "properties": { + "keep": {"type": "string"}, + "remove": {"$ref": "#/$defs/remove_def"}, + }, + "required": ["keep", "remove"], + "additionalProperties": False, + "$defs": { + "remove_def": {"type": "string"}, + "unused_def": {"type": "number"}, + }, + } + result = compress_schema(schema, prune_params=["remove"]) + # Check that parameter was removed + assert "remove" not in result["properties"] + # Check that required list was updated + assert result["required"] == ["keep"] + # Check that unused definitions were removed + assert "$defs" not in result # Both defs should be gone + # Check that additionalProperties was removed + assert "additionalProperties" not in result diff --git a/tests/utilities/test_typeadapter.py b/tests/utilities/test_typeadapter.py index 921858624..68f11b91d 100644 --- a/tests/utilities/test_typeadapter.py +++ b/tests/utilities/test_typeadapter.py @@ -13,7 +13,7 @@ import annotated_types import pytest from pydantic import BaseModel, Field -from fastmcp.utilities.json_schema import prune_params +from fastmcp.utilities.json_schema import compress_schema from fastmcp.utilities.types import get_cached_typeadapter @@ -175,7 +175,7 @@ def test_skip_names(): # Get schema and prune parameters type_adapter = get_cached_typeadapter(func_with_many_params) schema = type_adapter.json_schema() - pruned_schema = prune_params(schema, params=["skip_this", "also_skip"]) + pruned_schema = compress_schema(schema, prune_params=["skip_this", "also_skip"]) # Check that only the desired parameters remain assert "keep_this" in pruned_schema["properties"]