diff --git a/src/anthropic/lib/_parse/_transform.py b/src/anthropic/lib/_parse/_transform.py index ce0c83ac9..9214ae902 100644 --- a/src/anthropic/lib/_parse/_transform.py +++ b/src/anthropic/lib/_parse/_transform.py @@ -31,6 +31,44 @@ "uuid", } +_SUPPORTED_TYPES: frozenset[str] = frozenset(SupportedTypes.__args__) + +# The keywords each type actually consumes below. When a type array is split +# into branches, only these travel to the matching branch; anything else stays +# on the parent, so it is described once rather than repeated in every branch. +# Keep in step with the per-type handling in transform_schema(). +_TYPE_SPECIFIC_KEYS: dict[str, tuple[str, ...]] = { + "object": ("properties", "additionalProperties", "required"), + "string": ("format",), + "array": ("items", "minItems"), +} + + +def _validate_type_array(type_: list[Any]) -> list[str]: + """Validate a JSON Schema `type` array and return its members. + + JSON Schema (draft 4 onward) allows `type` to be an array of type names, + which is what `z.string().nullable()` and `Optional[str]` emit. The array + must be non-empty and its members must be unique type names. + """ + if not type_: + raise ValueError("Schema 'type' array must not be empty.") + + seen: set[str] = set() + for member in type_: + if not isinstance(member, str): + raise ValueError(f"Schema 'type' array must contain strings, got {member!r}.") + if member not in _SUPPORTED_TYPES: + raise ValueError( + f"Unsupported schema type {member!r} in 'type' array. Supported types: " + f"{', '.join(sorted(_SUPPORTED_TYPES))}." + ) + if member in seen: + raise ValueError(f"Schema 'type' array must not repeat {member!r}.") + seen.add(member) + + return cast("list[str]", type_) + def get_transformed_string( schema: dict[str, Any], @@ -96,21 +134,52 @@ def transform_schema( strict_schema["$ref"] = ref return strict_schema - type_: Optional[SupportedTypes] = json_schema.pop("type", None) + raw_type: Any = json_schema.pop("type", None) any_of = json_schema.pop("anyOf", None) one_of = json_schema.pop("oneOf", None) all_of = json_schema.pop("allOf", None) - if is_list(any_of): + # stays None when the schema is a combinator or a type array, so the + # per-type handling further down is skipped for those + type_: Optional[SupportedTypes] = None + + if is_list(raw_type): + # `{"type": ["string", "null"]}` is valid JSON Schema and is what + # `z.string().nullable()` and `Optional[str]` emit. It means the same + # thing as the `anyOf` spelling, so rewrite it that way -- the same + # move this function already makes for `oneOf` just below. + members = _validate_type_array(cast("list[Any]", raw_type)) + + # Give each branch only the keywords its own type consumes. Handing the + # whole schema to every branch would, for `["object", "null"]`, stringify + # the entire `properties` dict into the null branch's description. + # Whatever no branch claims stays on the parent and is described once. + branches: list[dict[str, Any]] = [] + for member in members: + branch: dict[str, Any] = {"type": member} + for key in _TYPE_SPECIFIC_KEYS.get(member, ()): + if key in json_schema: + branch[key] = json_schema[key] + branches.append(branch) + + for key in {key for member in members for key in _TYPE_SPECIFIC_KEYS.get(member, ())}: + json_schema.pop(key, None) + + strict_schema["anyOf"] = [transform_schema(branch) for branch in branches] + elif is_list(any_of): strict_schema["anyOf"] = [transform_schema(cast("dict[str, Any]", variant)) for variant in any_of] elif is_list(one_of): strict_schema["anyOf"] = [transform_schema(cast("dict[str, Any]", variant)) for variant in one_of] elif is_list(all_of): strict_schema["allOf"] = [transform_schema(cast("dict[str, Any]", variant)) for variant in all_of] else: - if type_ is None: + if raw_type is None: raise ValueError("Schema must have a 'type', 'anyOf', 'oneOf', or 'allOf' field.") + # An unrecognised type *name* still reaches assert_never below, which + # test_unsupported_type_asserts pins deliberately. That path is left + # alone: a type array is not a bad type name, it is a different shape. + type_ = cast("SupportedTypes", raw_type) strict_schema["type"] = type_ enum = json_schema.pop("enum", None) diff --git a/tests/lib/_parse/test_transform.py b/tests/lib/_parse/test_transform.py index 7a2799dce..d806fc2be 100644 --- a/tests/lib/_parse/test_transform.py +++ b/tests/lib/_parse/test_transform.py @@ -229,3 +229,117 @@ def test_original_schema_not_mutated(): transform_schema(original_schema) assert original_schema == original_schema_backup + + +def test_type_array_nullable_string(): + # what z.string().nullable() and Optional[str] emit + schema = {"type": ["string", "null"]} + result = transform_schema(schema) + assert result == snapshot({"anyOf": [{"type": "string"}, {"type": "null"}]}) + + +def test_type_array_matches_the_any_of_spelling(): + # the two spellings mean the same thing, so they should transform alike + assert transform_schema({"type": ["string", "null"]}) == transform_schema( + {"anyOf": [{"type": "string"}, {"type": "null"}]} + ) + + +def test_type_array_union_without_null(): + schema = {"type": ["string", "number"]} + result = transform_schema(schema) + assert result == snapshot({"anyOf": [{"type": "string"}, {"type": "number"}]}) + + +def test_type_array_keeps_description_on_the_parent(): + # description belongs to the schema, not to one branch of it; an + # unsupported sibling keyword is described once rather than per branch + schema = {"type": ["string", "null"], "description": "A query", "minLength": 2} + result = transform_schema(schema) + assert result == snapshot( + { + "anyOf": [{"type": "string"}, {"type": "null"}], + "description": "A query\n\n{minLength: 2}", + } + ) + + +def test_type_array_object_keeps_properties_out_of_the_null_branch(): + schema = { + "type": ["object", "null"], + "properties": {"name": {"type": "string"}}, + "required": ["name"], + } + result = transform_schema(schema) + assert result == snapshot( + { + "anyOf": [ + { + "type": "object", + "properties": {"name": {"type": "string"}}, + "additionalProperties": False, + "required": ["name"], + }, + {"type": "null"}, + ] + } + ) + + +def test_type_array_array_branch_keeps_items(): + schema = {"type": ["array", "null"], "items": {"type": "integer"}, "minItems": 1} + result = transform_schema(schema) + assert result == snapshot( + { + "anyOf": [ + {"type": "array", "items": {"type": "integer"}, "minItems": 1}, + {"type": "null"}, + ] + } + ) + + +def test_type_array_nested_in_a_property(): + schema = { + "type": "object", + "properties": {"cursor": {"type": ["string", "null"]}}, + "required": ["cursor"], + } + result = transform_schema(schema) + assert result == snapshot( + { + "type": "object", + "properties": {"cursor": {"anyOf": [{"type": "string"}, {"type": "null"}]}}, + "additionalProperties": False, + "required": ["cursor"], + } + ) + + +@pytest.mark.parametrize( + "type_", + [ + pytest.param([], id="empty"), + pytest.param(["string", 7], id="non-string member"), + pytest.param(["string", "banana"], id="unknown type name"), + pytest.param(["string", "string"], id="repeated member"), + ], +) +def test_type_array_malformed_raises_value_error(type_: object) -> None: + # a malformed array is the caller's input being wrong, so it should read as + # a ValueError rather than as an SDK invariant breaking + with pytest.raises(ValueError): + transform_schema({"type": type_}) + + +def test_type_array_does_not_mutate_the_original_schema(): + original_schema = { + "type": "object", + "properties": {"cursor": {"type": ["string", "null"], "minLength": 2}}, + "required": ["cursor"], + } + original_schema_backup = deepcopy(original_schema) + + transform_schema(original_schema) + + assert original_schema == original_schema_backup