diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index aa16d4bfdb7..f987324e4de 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -1212,11 +1212,14 @@ def _resolve_local_schema_ref(root: Mapping[str, object], ref: str) -> Mapping[s def _inline_root_schema_ref(schema: Mapping[str, object]) -> Mapping[str, object]: - """Replace a root-level local ``$ref`` with the schema it points at. + """Merge a root-level local ``$ref`` with the schema it points at. - The glossaries are carried over so refs inside the target still resolve, and - the target wins on any key it shares with the root. A root without a ``$ref``, - or one pointing outside the document or at a missing name, is returned as is. + Per JSON Schema, keys beside a ``$ref`` apply on top of what it references + rather than being replaced by it, so ``properties`` and ``required`` union + and the local key wins elsewhere. Dropping the siblings instead would lose + constraints 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): @@ -1224,10 +1227,11 @@ def _inline_root_schema_ref(schema: Mapping[str, object]) -> Mapping[str, object target: Final = _resolve_local_schema_ref(schema, ref) if target is None: return schema - glossaries: Final = MappingProxyType( - {container: schema[container] for _, container in _LOCAL_SCHEMA_REF_PREFIXES if container in schema} - ) - return MappingProxyType({**glossaries, **target}) + siblings: Final = MappingProxyType({key: value for key, value in schema.items() if key != "$ref"}) + properties: Final = MappingProxyType({**_schema_properties(target), **_schema_properties(siblings)}) + 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( @@ -1263,6 +1267,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], @@ -1291,7 +1307,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 b0d5a7c74b3..76dd27ee7e6 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 @@ -2010,7 +2010,7 @@ class TestSanitizeInputSchemaForAnthropic: "required": ["kind", "b"], } - def _union_root(self, combinator): + 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"]) @@ -2027,6 +2027,40 @@ class TestSanitizeInputSchemaForAnthropic: "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_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, @@ -2085,6 +2119,38 @@ class TestSanitizeInputSchemaForAnthropic: 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)}" + + def test_branches_that_resolved_to_one_object_collapse_to_one(self): + """Repeated ``$ref``s share a single memoised result, so the merge only + needs to see it once. Collapsing them is what keeps a compact schema + repeating one reference from costing work per branch.""" + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + _distinct_branches, + ) + + shared = {"type": "object", "properties": {"a": {"type": "string"}}} + other = {"type": "object", "properties": {"b": {"type": "string"}}} + + result = _distinct_branches((shared, shared, other, shared)) + + assert len(result) == 2, f"identical objects should collapse, got {len(result)} branches" + assert result[0] is shared, f"the first occurrence should be kept, got {result}" + assert result[1] is other, f"distinct branches must survive in order, got {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