diff --git a/mellea/backends/tools.py b/mellea/backends/tools.py index ae461dfb6..c3936e09f 100644 --- a/mellea/backends/tools.py +++ b/mellea/backends/tools.py @@ -632,9 +632,15 @@ def _build_pydantic_type_from_schema(schema: dict[str, Any]) -> Any: **nested_fields, ) - # Handle arrays + # Handle arrays. An `items` schema naming no type (missing, `{}`, + # description-only, or `true`) leaves elements unconstrained; the + # `string` default would coerce numbers to strings. if json_type == "array": - item_schema = schema.get("items", {}) + item_schema = schema.get("items") + if not isinstance(item_schema, dict) or not any( + key in item_schema for key in ("type", "anyOf", "enum", "const") + ): + return list[Any] item_type = _build_pydantic_type_from_schema(item_schema) return list[item_type] # type: ignore @@ -708,41 +714,46 @@ def _build_pydantic_type_from_schema(schema: dict[str, Any]) -> Any: # Simple type mapping return JSON_TYPE_TO_PYTHON.get(json_type, Any) - # Build Pydantic model from JSON schema - field_definitions: dict[str, Any] = {} + # Build inside the try so an unbuildable schema takes the lenient fallback. + try: + # Build Pydantic model from JSON schema + field_definitions: dict[str, Any] = {} - for param_name, param_schema in properties.items(): - param_type = _build_pydantic_type_from_schema(param_schema) + for param_name, param_schema in properties.items(): + param_type = _build_pydantic_type_from_schema(param_schema) - # Determine if parameter is required - if param_name in required_fields: - # Required parameter - field_definitions[param_name] = (param_type, ...) + # Determine if parameter is required + if param_name in required_fields: + # Required parameter + field_definitions[param_name] = (param_type, ...) + else: + # Optional parameter (default to None). The schema drops `null` + # from simple Optional types, so a field with no non-null default + # is taken as nullable; `page_size: int = 10` still rejects None. + if param_schema.get("default") is None: + param_type = param_type | None + field_definitions[param_name] = (param_type, None) + + # Configure model for type coercion if requested + if coerce_types: + model_config = ConfigDict( + str_strip_whitespace=True, + strict=False, # Allow type coercion + extra="forbid" if strict else "allow", # Handle extra fields + # Enable coercion modes for common LLM output issues + coerce_numbers_to_str=True, # Allow int/float -> str + ) else: - # Optional parameter (default to None) - field_definitions[param_name] = (param_type, None) - - # Configure model for type coercion if requested - if coerce_types: - model_config = ConfigDict( - str_strip_whitespace=True, - strict=False, # Allow type coercion - extra="forbid" if strict else "allow", # Handle extra fields - # Enable coercion modes for common LLM output issues - coerce_numbers_to_str=True, # Allow int/float -> str - ) - else: - model_config = ConfigDict( - strict=True, # No coercion - extra="forbid" if strict else "allow", - ) + model_config = ConfigDict( + strict=True, # No coercion + extra="forbid" if strict else "allow", + ) - # Create dynamic Pydantic model for validation - ValidatorModel = create_model( - f"{tool_name}_Validator", __config__=model_config, **field_definitions - ) + # Create dynamic Pydantic model for validation + ValidatorModel = create_model( + f"{tool_name}_Validator", __config__=model_config, **field_definitions + ) - try: # Validate using Pydantic validated_model = ValidatorModel(**args) # Only emit fields the model actually sent. A bare model_dump() would @@ -1346,6 +1357,28 @@ def _inline(branch: dict) -> dict: return out +# JSON Schema keywords copied onto a rebuilt simple property, so the model and +# `validate_tool_arguments` both see element types and constraints. +_CARRIED_SCHEMA_KEYWORDS = ( + "items", + "additionalProperties", + "minItems", + "maxItems", + "uniqueItems", + "minProperties", + "maxProperties", + "minLength", + "maxLength", + "pattern", + "format", + "minimum", + "maximum", + "exclusiveMinimum", + "exclusiveMaximum", + "multipleOf", +) + + # https://github.com/ollama/ollama-python/blob/60e7b2f9ce710eeb57ef2986c46ea612ae7516af/ollama/_utils.py#L56-L90 def convert_function_to_ollama_tool( func: Callable, name: str | None = None @@ -1493,12 +1526,26 @@ def convert_function_to_ollama_tool( # from scratch would otherwise drop it. if "default" in v: simple_prop["default"] = v["default"] + # Carry element types and constraints. For Optional, Pydantic puts + # them on the single non-null anyOf branch; with several branches + # they can't be merged onto one property. + if "anyOf" in v: + non_null = [s for s in v["anyOf"] if s.get("type") != "null"] + keyword_source = non_null[0] if len(non_null) == 1 else {} + else: + keyword_source = v + for keyword in _CARRIED_SCHEMA_KEYWORDS: + if keyword in keyword_source: + simple_prop[keyword] = keyword_source[keyword] schema["properties"][k] = simple_prop # Final pass: recursively inline all remaining $refs at any depth. # This catches dangling references in nested model properties that weren't # caught by the earlier single-level ref-inlining passes. _recursively_inline_refs(schema, defs) + # Inlining can expose discriminated unions the pre-pass missed + # (`list[Pet]`, a nested model's field); flatten those too. + _recursively_flatten_in_properties(schema, defs) tool = OllamaTool( type="function", diff --git a/test/backends/test_tool_collection_params_unit.py b/test/backends/test_tool_collection_params_unit.py new file mode 100644 index 000000000..b28416d17 --- /dev/null +++ b/test/backends/test_tool_collection_params_unit.py @@ -0,0 +1,491 @@ +# Copyright IBM Corp. All Rights Reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for tool schemas and validation of list, dict and constrained parameters. + +Regression tests for the bug where the schema rebuild for simple parameters +kept only `description`, `type`, `enum` and `default`, dropping `items`, +`additionalProperties` and constraints. The model was never told a list's +element type, and `validate_tool_arguments` (which builds its validator from +the same schema) fell back to `string` elements, so `list[int]` arguments +reached the tool as strings. The schema is shared by every backend via +`convert_tools_to_json`, and the validator runs on every backend. + +See: https://github.com/generative-computing/mellea/issues/1693 +""" + +import json +from typing import Annotated, Any, Literal + +import pytest +from pydantic import BaseModel, Field, ValidationError + +from mellea.backends.tools import MelleaTool, validate_tool_arguments + + +class Point(BaseModel): + x: int + y: int + + +class Cat(BaseModel): + kind: Literal["cat"] + meow: int + + +class Dog(BaseModel): + kind: Literal["dog"] + bark: int + + +Pet = Annotated[Cat | Dog, Field(discriminator="kind")] + + +class Owner(BaseModel): + name: str + pet: Pet + + +def pets_tool(pets: list[Pet]) -> int: + """Count pets. + + Args: + pets: the pets + """ + return len(pets) + + +def optional_pets_tool(pets: list[Pet] | None = None) -> int: + """Count pets, if given. + + Args: + pets: the pets + """ + return len(pets or []) + + +def named_pets_tool(pets: dict[str, Pet]) -> int: + """Count named pets. + + Args: + pets: pets by name + """ + return len(pets) + + +def owners_tool(owners: list[Owner]) -> int: + """Count owners. + + Args: + owners: the owners + """ + return len(owners) + + +def owner_tool(owner: Owner) -> str: + """Name an owner. + + Args: + owner: the owner + """ + return owner.name + + +def ints_tool(nums: list[int]) -> int: + """Sum some numbers. + + Args: + nums: the numbers + """ + return sum(nums) + + +def floats_tool(nums: list[float]) -> float: + """Sum some floats. + + Args: + nums: the numbers + """ + return sum(nums) + + +def optional_ints_tool(nums: list[int] | None = None) -> int: + """Sum some numbers, if given. + + Args: + nums: the numbers + """ + return sum(nums or []) + + +def bools_tool(flags: list[bool]) -> int: + """Count true flags. + + Args: + flags: the flags + """ + return sum(flags) + + +def nested_ints_tool(rows: list[list[int]]) -> int: + """Sum a matrix. + + Args: + rows: the rows + """ + return sum(sum(row) for row in rows) + + +def points_tool(pts: list[Point]) -> int: + """Count points. + + Args: + pts: the points + """ + return len(pts) + + +def optional_points_tool(pts: list[Point] | None = None) -> int: + """Count points, if given. + + Args: + pts: the points + """ + return len(pts or []) + + +def counts_tool(counts: dict[str, int]) -> int: + """Total some counts. + + Args: + counts: counts by name + """ + return sum(counts.values()) + + +def bounded_tool(score: Annotated[int, Field(ge=0, le=10)]) -> int: + """Record a score. + + Args: + score: the score + """ + return score + + +def optional_bounded_tool(score: Annotated[int, Field(ge=0)] | None = None) -> int: + """Record a score, if given. + + Args: + score: the score + """ + return score or 0 + + +def non_empty_tool(nums: Annotated[list[int], Field(min_length=1)]) -> int: + """Sum at least one number. + + Args: + nums: the numbers + """ + return sum(nums) + + +def bounded_str_tool( + code: Annotated[str, Field(max_length=5, pattern="^[A-Z]+$")], +) -> str: + """Look up a code. + + Args: + code: the code + """ + return code + + +def any_list_tool(values: list[Any]) -> int: + """Count values. + + Args: + values: the values + """ + return len(values) + + +def described_any_list_tool( + values: list[Annotated[Any, Field(description="a value")]], +) -> int: + """Count values. + + Args: + values: the values + """ + return len(values) + + +def strings_tool(words: list[str]) -> int: + """Count words. + + Args: + words: the words + """ + return len(words) + + +def _prop(func, param): + """Return the serialized schema for one parameter of a callable's tool.""" + return MelleaTool.from_callable(func).as_json_tool["function"]["parameters"][ + "properties" + ][param] + + +def _required(func): + """Return the `required` list of a callable's tool schema.""" + params = MelleaTool.from_callable(func).as_json_tool["function"]["parameters"] + return params.get("required", []) + + +# ============================================================================ +# Schema generation: what every backend sends to the model +# ============================================================================ + + +class TestElementTypesInSchema: + """Element types of lists and dicts must reach the model.""" + + def test_list_int_property_shape(self): + assert _prop(ints_tool, "nums") == { + "description": "the numbers", + "type": "array", + "items": {"type": "integer"}, + } + + def test_list_float_items(self): + assert _prop(floats_tool, "nums")["items"] == {"type": "number"} + + def test_list_bool_items(self): + assert _prop(bools_tool, "flags")["items"] == {"type": "boolean"} + + def test_list_str_items(self): + assert _prop(strings_tool, "words")["items"] == {"type": "string"} + + def test_nested_list_items(self): + assert _prop(nested_ints_tool, "rows")["items"] == { + "type": "array", + "items": {"type": "integer"}, + } + + def test_optional_list_items_taken_from_array_branch(self): + prop = _prop(optional_ints_tool, "nums") + assert prop["type"] == "array" + assert prop["items"] == {"type": "integer"} + assert "nums" not in _required(optional_ints_tool) + + def test_list_of_models_items_inlined(self): + items = _prop(points_tool, "pts")["items"] + assert "$ref" not in items + assert items["type"] == "object" + assert set(items["properties"]) == {"x", "y"} + assert items["properties"]["x"]["type"] == "integer" + assert set(items["required"]) == {"x", "y"} + + def test_optional_list_of_models_items_inlined(self): + prop = _prop(optional_points_tool, "pts") + assert prop["type"] == "array" + assert "$ref" not in prop["items"] + assert set(prop["items"]["properties"]) == {"x", "y"} + assert "pts" not in _required(optional_points_tool) + + def test_dict_value_type(self): + prop = _prop(counts_tool, "counts") + assert prop["type"] == "object" + assert prop["additionalProperties"] == {"type": "integer"} + + def test_any_list_keeps_unconstrained_items(self): + assert _prop(any_list_tool, "values")["items"] == {} + + +class TestDiscriminatedUnionsInCollections: + """Discriminated unions in carried keywords are flattened for tool APIs.""" + + @pytest.mark.parametrize( + ("func", "param"), + [ + (pets_tool, "pets"), + (optional_pets_tool, "pets"), + (named_pets_tool, "pets"), + (owners_tool, "owners"), + (owner_tool, "owner"), + ], + ids=["list", "optional_list", "dict", "list_of_models", "model_field"], + ) + def test_no_oneof_or_discriminator(self, func, param): + rendered = json.dumps(_prop(func, param)) + assert "oneOf" not in rendered + assert "discriminator" not in rendered + assert "$ref" not in rendered + + def test_list_items_become_inlined_anyof(self): + items = _prop(pets_tool, "pets")["items"] + kinds = {b["properties"]["kind"]["const"] for b in items["anyOf"]} + assert kinds == {"cat", "dog"} + + def test_dict_values_become_inlined_anyof(self): + values = _prop(named_pets_tool, "pets")["additionalProperties"] + assert len(values["anyOf"]) == 2 + + +class TestConstraintsInSchema: + """`Field` constraints must reach the model.""" + + def test_numeric_bounds(self): + prop = _prop(bounded_tool, "score") + assert prop["type"] == "integer" + assert prop["minimum"] == 0 + assert prop["maximum"] == 10 + + def test_optional_numeric_bound_taken_from_branch(self): + prop = _prop(optional_bounded_tool, "score") + assert prop["type"] == "integer" + assert prop["minimum"] == 0 + + def test_list_length(self): + prop = _prop(non_empty_tool, "nums") + assert prop["minItems"] == 1 + assert prop["items"] == {"type": "integer"} + + def test_string_length_and_pattern(self): + prop = _prop(bounded_str_tool, "code") + assert prop["maxLength"] == 5 + assert prop["pattern"] == "^[A-Z]+$" + + def test_title_not_carried(self): + """Pydantic's per-property `title` stays stripped, as before.""" + assert "title" not in _prop(bounded_tool, "score") + assert "title" not in _prop(ints_tool, "nums") + + +# ============================================================================ +# Validation: what the tool actually receives +# ============================================================================ + + +class TestListValidation: + """Correct list arguments must pass validation with their element types intact.""" + + def test_list_int_not_coerced_to_str(self): + tool = MelleaTool.from_callable(ints_tool) + validated = validate_tool_arguments(tool, {"nums": [1, 2]}, strict=True) + assert validated == {"nums": [1, 2]} + assert all(type(n) is int for n in validated["nums"]) + + def test_list_int_tool_runs(self): + """The reproducer from the issue: the tool must not raise TypeError.""" + tool = MelleaTool.from_callable(ints_tool) + validated = validate_tool_arguments(tool, {"nums": [1, 2]}, strict=True) + assert tool.run(**validated) == 3 + + def test_list_int_elements_coerced_from_str(self): + tool = MelleaTool.from_callable(ints_tool) + validated = validate_tool_arguments(tool, {"nums": ["1", "2"]}) + assert validated == {"nums": [1, 2]} + + def test_list_float(self): + tool = MelleaTool.from_callable(floats_tool) + validated = validate_tool_arguments(tool, {"nums": [1.5]}, strict=True) + assert validated == {"nums": [1.5]} + assert type(validated["nums"][0]) is float + + def test_optional_list_int(self): + tool = MelleaTool.from_callable(optional_ints_tool) + validated = validate_tool_arguments(tool, {"nums": [1, 2]}, strict=True) + assert validated == {"nums": [1, 2]} + assert all(type(n) is int for n in validated["nums"]) + + def test_optional_list_int_none(self): + tool = MelleaTool.from_callable(optional_ints_tool) + validated = validate_tool_arguments(tool, {"nums": None}, strict=True) + assert validated == {"nums": None} + + @pytest.mark.parametrize( + ("func", "args"), + [ + (bools_tool, {"flags": [True, False]}), + (nested_ints_tool, {"rows": [[1], [2, 3]]}), + (points_tool, {"pts": [{"x": 1, "y": 2}]}), + (optional_points_tool, {"pts": [{"x": 1, "y": 2}]}), + ], + ids=["list_bool", "list_list_int", "list_model", "optional_list_model"], + ) + def test_previously_rejected_types_validate(self, func, args): + """These used to skip validation; `strict=True` proves they now pass it.""" + tool = MelleaTool.from_callable(func) + assert validate_tool_arguments(tool, args, strict=True) == args + + def test_list_of_models_rejects_missing_field(self): + """Element validation is real: a point missing `y` is rejected.""" + tool = MelleaTool.from_callable(points_tool) + with pytest.raises(ValidationError, match=r"pts\.0\.y"): + validate_tool_arguments(tool, {"pts": [{"x": 1}]}, strict=True) + + def test_list_str_still_coerces_numbers(self): + """Existing `list[str]` coercion of numbers to strings is unchanged.""" + tool = MelleaTool.from_callable(strings_tool) + validated = validate_tool_arguments(tool, {"words": ["a", 1]}) + assert validated == {"words": ["a", "1"]} + + def test_any_list_elements_untouched(self): + tool = MelleaTool.from_callable(any_list_tool) + validated = validate_tool_arguments( + tool, {"values": [1, "a", True]}, strict=True + ) + assert validated == {"values": [1, "a", True]} + assert type(validated["values"][0]) is int + + def test_list_of_discriminated_union_validates(self): + """Elements are validated against the branch the tag selects.""" + tool = MelleaTool.from_callable(pets_tool) + validated = validate_tool_arguments( + tool, {"pets": [{"kind": "cat", "meow": "1"}]}, strict=True + ) + assert validated == {"pets": [{"kind": "cat", "meow": 1}]} + + def test_list_of_discriminated_union_rejects_wrong_branch_fields(self): + tool = MelleaTool.from_callable(pets_tool) + with pytest.raises(ValidationError, match=r"pets\.0\.cat\.meow"): + validate_tool_arguments( + tool, {"pets": [{"kind": "cat", "bark": 1}]}, strict=True + ) + + def test_described_any_list_elements_untouched(self): + """An annotation-only items schema constrains nothing either.""" + tool = MelleaTool.from_callable(described_any_list_tool) + validated = validate_tool_arguments(tool, {"values": [1, 2]}, strict=True) + assert validated == {"values": [1, 2]} + assert type(validated["values"][0]) is int + + @pytest.mark.parametrize( + "array_schema", + [ + {"type": "array"}, + {"type": "array", "items": True}, + {"type": "array", "items": {"description": "a value"}}, + ], + ids=["no_items", "items_true", "annotation_only_items"], + ) + def test_external_unconstrained_array_elements_untouched(self, array_schema): + """Unconstrained external `items` must not coerce elements to strings.""" + as_json_tool = { + "type": "function", + "function": { + "name": "external", + "description": "An externally defined tool.", + "parameters": { + "type": "object", + "properties": {"values": array_schema}, + "required": ["values"], + }, + }, + } + tool = MelleaTool("external", lambda values: values, as_json_tool) + validated = validate_tool_arguments(tool, {"values": [1, 2.5]}, strict=True) + assert validated == {"values": [1, 2.5]} + assert type(validated["values"][0]) is int diff --git a/test/backends/test_tool_validation_integration.py b/test/backends/test_tool_validation_integration.py index a2dcadf33..609788ac3 100644 --- a/test/backends/test_tool_validation_integration.py +++ b/test/backends/test_tool_validation_integration.py @@ -51,6 +51,26 @@ def optional_tool(required: str, optional: str | None = None) -> str: return f"{required}:{optional or 'none'}" +def limit_tool(count: int, limit: int | None = None) -> int: + """Tool with a required int and an optional int. + + Args: + count: How many items to return + limit: An optional upper bound + """ + return count if limit is None else min(count, limit) + + +def paged_tool(query: str, page_size: int = 10) -> str: + """Tool with a non-nullable defaulted parameter. + + Args: + query: The search query + page_size: Results per page + """ + return f"{query}:{page_size}" + + def union_tool(value: str | int) -> str: """Tool with union type parameter. @@ -252,13 +272,36 @@ def test_optional_param_omitted(self): assert "optional" not in validated def test_optional_param_none(self): - """Test validation when optional parameter is explicitly None.""" + """Explicit None for `str | None` validates (strict, so no fallback).""" args = {"required": "value1", "optional": None} tool = MelleaTool.from_callable(optional_tool) + validated = validate_tool_arguments(tool, args, strict=True) + + assert validated == {"required": "value1", "optional": None} + + def test_optional_int_none_strict(self): + """An explicit None for `int | None` validates.""" + args = {"count": 3, "limit": None} + tool = MelleaTool.from_callable(limit_tool) + validated = validate_tool_arguments(tool, args, strict=True) + + assert validated == {"count": 3, "limit": None} + + def test_null_for_non_nullable_default_rejected_strict(self): + """A real default (`page_size: int = 10`) still rejects None.""" + args = {"query": "q", "page_size": None} + tool = MelleaTool.from_callable(paged_tool) + with pytest.raises(ValidationError, match="page_size"): + validate_tool_arguments(tool, args, strict=True) + + def test_optional_none_keeps_other_coercions(self): + """A None must not make lenient mode drop the other coercions.""" + args = {"count": "3", "limit": None} + tool = MelleaTool.from_callable(limit_tool) validated = validate_tool_arguments(tool, args) - assert validated["required"] == "value1" - assert validated["optional"] is None + assert validated == {"count": 3, "limit": None} + assert type(validated["count"]) is int class TestDefaultedParameters: @@ -375,9 +418,48 @@ def test_union_with_string_number(self): assert validated["value"] in ["42", 42] +def _external_tool(value_schema: Any) -> MelleaTool: + """Build a tool whose single `value` parameter uses a raw JSON schema.""" + as_json_tool = { + "type": "function", + "function": { + "name": "external", + "description": "An externally defined tool.", + "parameters": { + "type": "object", + "properties": {"value": value_schema}, + "required": ["value"], + }, + }, + } + return MelleaTool("external", lambda value: value, as_json_tool) + + +# Boolean sub-schemas are valid JSON Schema (e.g. from MCP) but broke the validator. +BOOLEAN_SUBSCHEMAS = [ + pytest.param(True, 1, id="property_true"), + pytest.param({"type": "object", "properties": {"a": True}}, {"a": 1}, id="nested"), + pytest.param({"anyOf": [True, {"type": "null"}]}, 1, id="anyof"), +] + + class TestEdgeCases: """Test edge cases.""" + @pytest.mark.parametrize(("value_schema", "value"), BOOLEAN_SUBSCHEMAS) + def test_unsupported_schema_lenient_returns_original_args( + self, value_schema, value + ): + """An unbuildable schema falls back to the original args, not a raise.""" + tool = _external_tool(value_schema) + assert validate_tool_arguments(tool, {"value": value}) == {"value": value} + + @pytest.mark.parametrize(("value_schema", "value"), BOOLEAN_SUBSCHEMAS) + def test_unsupported_schema_strict_raises(self, value_schema, value): + tool = _external_tool(value_schema) + with pytest.raises((TypeError, AttributeError)): + validate_tool_arguments(tool, {"value": value}, strict=True) + def test_no_parameters_tool(self): """Test validation with no-parameter tool.""" args = {}