diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index e555d7e8ec0..da2d75a0a38 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -11,7 +11,9 @@ from itertools import groupby, islice from os import PathLike from pathlib import Path from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, cast +from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, TypeVar, cast + +from typing_extensions import ReadOnly import litellm from litellm import verbose_logger @@ -1140,10 +1142,17 @@ def sanitize_input_schema_for_anthropic(input_schema: dict) -> "AnthropicInputSc (``AnthropicConfig._map_tool_helper``) and Anthropic Messages MCP paths run a schema through here so an external MCP schema cannot succeed on one route and 400 on the other. + + A root that is a local ``$ref`` or a top-level combinator is resolved into an + object first. Without that, ``type: "object"`` and an empty ``properties`` + are stamped over a schema whose fields live one level down, and the + allowlist then drops the combinator: a ``TypeAdapter(Union[A, B])`` tool + reaches Anthropic taking no arguments at all. An unresolvable root (external + or missing ``$ref``, a branch that cannot merge) is left alone. """ from litellm.types.llms.anthropic import AnthropicInputSchema - normalized = dict(input_schema) if input_schema else {} + normalized = dict(flatten_top_level_schema_combinators(_inline_root_schema_ref(input_schema or {}))) if normalized.get("type") != "object": normalized["type"] = "object" if "properties" not in normalized: @@ -1204,6 +1213,53 @@ def _resolve_local_schema_ref(root: Mapping[str, object], ref: str) -> Mapping[s return target if isinstance(target, dict) else None +class _AllOfSchema(TypedDict): + allOf: ReadOnly[Sequence[object]] + + +def _both_constraints_apply(target_property: object, sibling_property: object) -> object: + """One property schema satisfying both sides of a ``$ref`` merge. + + A plain dict rather than a ``MappingProxyType``, because this value is + serialised straight into the provider request and ``json`` cannot encode a + mapping proxy. + """ + if target_property is None or target_property == sibling_property: + return sibling_property + both: Final = [target_property, sibling_property] # mutable-ok: allOf is a JSON array; a tuple is not one + combined: Final[_AllOfSchema] = {"allOf": both} + return combined + + +def _inline_root_schema_ref(schema: Mapping[str, object]) -> Mapping[str, object]: + """Merge a root-level local ``$ref`` with the schema it points at. + + Per JSON Schema, keys beside a ``$ref`` apply on top of what it references + rather than being replaced by it. ``properties`` and ``required`` therefore + union, and a property both sides declare becomes an ``allOf`` of the two so + neither one's constraints are lost; Anthropic accepts a combinator nested + inside a property, only a root one is a problem. Replacing the root outright + would drop what the caller stated here, such as ``additionalProperties: + false``. A root without a ``$ref``, or one pointing outside the document or + at a missing name, is returned as is. + """ + ref: Final = schema.get("$ref") + if not isinstance(ref, str): + return schema + target: Final = _resolve_local_schema_ref(schema, ref) + if target is None: + return schema + siblings: Final = MappingProxyType({key: value for key, value in schema.items() if key != "$ref"}) + target_properties: Final = _schema_properties(target) + sibling_wins: Final = MappingProxyType({**target_properties, **_schema_properties(siblings)}) + properties: Final = { # mutable-ok: tool parameters are JSON dicts, and a mapping proxy will not serialise + name: _both_constraints_apply(target_properties.get(name), value) for name, value in sibling_wins.items() + } + required: Final = sorted(_schema_required_names(target) | _schema_required_names(siblings)) + required_update: Final = MappingProxyType({"required": required}) if required else _EMPTY_SCHEMA + return MappingProxyType({**target, **siblings, "properties": properties, **required_update}) + + def _mergeable_branch( root: Mapping[str, object], branch: object, @@ -1237,6 +1293,18 @@ def _is_object_schema(schema: Mapping[str, object]) -> bool: return schema.get("type") == "object" or ("type" not in schema and "properties" in schema) +def _distinct_branches(branches: tuple[Mapping[str, object], ...]) -> tuple[Mapping[str, object], ...]: + """Collapse branches that resolved to the same object. + + Repeated ``$ref``s share one memoised result, so a schema declaring the same + reference thousands of times would otherwise re-merge its properties once per + branch. Merging a branch again cannot change the outcome: property merges are + idempotent, and so are the union and intersection used for ``required``. + """ + by_identity: Final = MappingProxyType({id(branch): branch for branch in branches}) + return tuple(by_identity.values()) + + def _flatten_schema_against_root( schema: Mapping[str, object], root: Mapping[str, object], @@ -1265,7 +1333,8 @@ def _flatten_schema_against_root( if any(branch is None for _, group in raw_branch_groups for branch in group): return schema branch_groups: Final = tuple( - (combinator, tuple(branch for branch in group if branch is not None)) for combinator, group in raw_branch_groups + (combinator, _distinct_branches(tuple(branch for branch in group if branch is not None))) + for combinator, group in raw_branch_groups ) branches: Final = tuple(branch for _, group in branch_groups for branch in group) is_object_schema: Final = _is_object_schema(schema) or ( diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py index 45fc93f04c1..138998db16a 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py @@ -1990,3 +1990,236 @@ class TestMergeConsecutiveSystemMessages: ) assert merged == [{"role": "system"}, {"role": "user", "content": "Hi"}] + + +class TestSanitizeInputSchemaForAnthropic: + """Anthropic rejects a tool schema whose root is a union or a ``$ref``, so the + sanitizer coerces one into an object. Stamping ``type: "object"`` and an empty + ``properties`` over a root whose fields live one level down, then dropping the + combinator with the allowlist, sent a ``TypeAdapter(Union[A, B])`` tool to + Anthropic taking no arguments at all (#43157).""" + + A: Final = { + "type": "object", + "properties": {"kind": {"type": "string"}, "a": {"type": "string"}}, + "required": ["kind", "a"], + } + B: Final = { + "type": "object", + "properties": {"kind": {"type": "string"}, "b": {"type": "integer"}}, + "required": ["kind", "b"], + } + + def _union_root(self, combinator: str) -> dict[str, object]: + return {"$defs": {"A": self.A, "B": self.B}, combinator: [{"$ref": "#/$defs/A"}, {"$ref": "#/$defs/B"}]} + + @pytest.mark.parametrize("combinator", ["anyOf", "oneOf"]) + def test_union_root_keeps_every_branch_field(self, combinator): + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + sanitize_input_schema_for_anthropic, + ) + + result = sanitize_input_schema_for_anthropic(self._union_root(combinator)) + + assert sorted(result["properties"]) == [ + "a", + "b", + "kind", + ], f"a {combinator} root must keep its branches' fields, got {dict(result)}" + + def test_ref_root_keeps_the_keys_declared_beside_it(self): + """JSON Schema applies keys beside a ``$ref`` on top of what it references. + + Replacing the root with the target instead would silently drop constraints + the caller stated here, such as ``additionalProperties``. + """ + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + sanitize_input_schema_for_anthropic, + ) + + result = sanitize_input_schema_for_anthropic( + { + "$defs": {"A": self.A}, + "$ref": "#/$defs/A", + "properties": {"sibling": {"type": "string"}}, + "required": ["sibling"], + "additionalProperties": False, + } + ) + + assert sorted(result["properties"]) == [ + "a", + "kind", + "sibling", + ], f"the target's fields and the sibling's must both survive, got {dict(result)}" + assert sorted(result["required"]) == [ + "a", + "kind", + "sibling", + ], f"both required lists must survive, got {dict(result)}" + assert result.get("additionalProperties") is False, ( + f"a constraint stated beside the $ref must not be dropped, got {dict(result)}" + ) + + def test_a_property_declared_on_both_sides_of_a_ref_keeps_both_constraints(self): + """Keys beside a ``$ref`` apply on top of it, so where both declare the + same property neither constraint may be dropped. Anthropic accepts a + combinator nested inside a property, so both can be carried.""" + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + sanitize_input_schema_for_anthropic, + ) + + result = sanitize_input_schema_for_anthropic( + { + "$defs": {"A": {"type": "object", "properties": {"a": {"type": "string", "minLength": 5}}}}, + "$ref": "#/$defs/A", + "properties": {"a": {"type": "string", "maxLength": 9}}, + } + ) + + assert result["properties"]["a"] == { + "allOf": [{"type": "string", "minLength": 5}, {"type": "string", "maxLength": 9}] + }, f"both sides' constraints on a shared property must survive, got {dict(result)}" + + def test_a_property_identical_on_both_sides_is_not_wrapped(self): + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + sanitize_input_schema_for_anthropic, + ) + + shared = {"type": "string"} + + result = sanitize_input_schema_for_anthropic( + { + "$defs": {"A": {"type": "object", "properties": {"a": shared}}}, + "$ref": "#/$defs/A", + "properties": {"a": shared}, + } + ) + + assert result["properties"]["a"] == shared, f"identical declarations need no allOf wrapper, got {dict(result)}" + + def test_ref_root_resolves_to_the_schema_it_points_at(self): + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + sanitize_input_schema_for_anthropic, + ) + + result = sanitize_input_schema_for_anthropic({"$defs": {"A": self.A}, "$ref": "#/$defs/A"}) + + assert sorted(result["properties"]) == ["a", "kind"], f"a $ref root must inline its target, got {dict(result)}" + assert sorted(result["required"]) == ["a", "kind"], ( + f"the target's required list must survive, got {dict(result)}" + ) + + def test_union_required_narrows_to_the_fields_every_branch_demands(self): + """Merging branches is lossy by design: only what all of them require stays + required, which is the tradeoff OpenAI and Azure already accept via the + shared flattening helper.""" + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + sanitize_input_schema_for_anthropic, + ) + + result = sanitize_input_schema_for_anthropic(self._union_root("anyOf")) + + assert sorted(result.get("required", [])) == ["kind"], ( + f"only the field both branches require stays required, got {dict(result)}" + ) + + @pytest.mark.parametrize( + "name, schema, expected_properties", + [ + ("plain object", {"type": "object", "properties": {"x": {"type": "string"}}}, ["x"]), + ("no-arg tool", {"type": "object", "properties": {}}, []), + ("empty schema", {}, []), + ("external $ref root", {"$ref": "https://example.com/schema.json"}, []), + ("missing $ref target", {"$defs": {}, "$ref": "#/$defs/Nope"}, []), + ], + ) + def test_shapes_the_sanitizer_already_handled_are_untouched(self, name, schema, expected_properties): + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + sanitize_input_schema_for_anthropic, + ) + + result = sanitize_input_schema_for_anthropic(schema) + + assert sorted(result["properties"]) == expected_properties, f"{name} should be unchanged, got {dict(result)}" + assert result["type"] == "object", f"{name} must still be coerced to an object, got {dict(result)}" + + def test_a_combinator_nested_inside_a_property_is_left_alone(self): + """Anthropic accepts nested combinators; only the root shape is a problem.""" + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + sanitize_input_schema_for_anthropic, + ) + + nested = {"anyOf": [{"type": "string"}, {"type": "integer"}]} + + result = sanitize_input_schema_for_anthropic({"type": "object", "properties": {"x": nested}}) + + assert result["properties"]["x"] == nested, f"a nested union must survive verbatim, got {dict(result)}" + + def test_repeating_one_ref_branch_costs_no_more_than_declaring_it_once(self): + """Repeated ``$ref``s share one memoised result, so merging them again + cannot change the outcome. A schema repeating a reference thousands of + times must therefore agree with the same schema declaring it once, which + is what stops a compact payload amplifying into per-branch merge work.""" + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + sanitize_input_schema_for_anthropic, + ) + + defs = {"$defs": {"A": self.A}} + once = sanitize_input_schema_for_anthropic({**defs, "anyOf": [{"$ref": "#/$defs/A"}]}) + many = sanitize_input_schema_for_anthropic({**defs, "anyOf": [{"$ref": "#/$defs/A"}] * 2000}) + + assert dict(many) == dict(once), f"repeating a branch must not change the result, got {dict(many)}" + + @pytest.mark.parametrize( + "name, root", + [ + ("ref root", {"$ref": "#/$defs/A"}), + ("union root", {"anyOf": [{"$ref": "#/$defs/A"}, {"$ref": "#/$defs/B"}]}), + ], + ) + def test_the_sanitized_schema_can_be_serialized_onto_the_wire(self, name, root): + """The result is JSON-encoded into the provider request, so every value in + it has to be encodable. A mapping proxy is not.""" + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + sanitize_input_schema_for_anthropic, + ) + + result = sanitize_input_schema_for_anthropic({"$defs": {"A": self.A, "B": self.B}, **root}) + + assert json.loads(json.dumps(result))["properties"], ( + f"{name} must survive a JSON round trip, got {dict(result)}" + ) + + def test_a_pydantic_union_tool_reaches_anthropic_with_its_arguments(self): + """The reporter's path: Pydantic emits a root ``anyOf`` over ``$defs``.""" + from typing import Literal + + from pydantic import BaseModel, TypeAdapter + + from litellm.llms.anthropic.chat.transformation import AnthropicConfig + + class Alpha(BaseModel): + kind: Literal["a"] + a: str + + class Beta(BaseModel): + kind: Literal["b"] + b: int + + tool = { + "type": "function", + "function": { + "name": "record", + "description": "record an event", + "parameters": TypeAdapter(Alpha | Beta).json_schema(), + }, + } + + mapped, _ = AnthropicConfig()._map_tools([tool]) + + assert sorted(mapped[0]["input_schema"]["properties"]) == [ + "a", + "b", + "kind", + ], f"the tool must reach Anthropic with arguments, got {mapped[0]['input_schema']}"