mirror of
https://github.com/strands-agents/harness-sdk.git
synced 2026-10-02 02:44:48 +08:00
fix(structured_output): preserve anyOf for nullable enums and respect model_json_schema overrides (#4610)
This commit is contained in:
@@ -123,6 +123,15 @@ def _process_property(
|
||||
# For Optional fields, we mark as nullable but copy all properties from the non-null option
|
||||
result = non_null_type.copy() if isinstance(non_null_type, dict) else {}
|
||||
|
||||
# Preserve anyOf when enum/const is present — a null type with non-null enum/const values is invalid.
|
||||
if "enum" in result or "const" in result:
|
||||
result = {"anyOf": [non_null_type, {"type": "null"}]}
|
||||
# Carry over top-level metadata from the original property
|
||||
for key, value in prop.items():
|
||||
if key != "anyOf":
|
||||
result[key] = value
|
||||
return result
|
||||
|
||||
# For type, ensure it includes "null"
|
||||
if "type" in result and isinstance(result["type"], str):
|
||||
result["type"] = [result["type"], "null"]
|
||||
@@ -159,7 +168,7 @@ def _process_property(
|
||||
if key not in ["$ref", "anyOf"]:
|
||||
if isinstance(value, dict):
|
||||
result[key] = _process_nested_dict(value, defs)
|
||||
elif key == "type" and not is_required and not is_nullable:
|
||||
elif key == "type" and not is_required and not is_nullable and "enum" not in prop and "const" not in prop:
|
||||
# For non-required fields, ensure type is a list with "null"
|
||||
if isinstance(value, str):
|
||||
result[key] = [value, "null"]
|
||||
@@ -326,6 +335,10 @@ def _expand_nested_properties(schema: dict[str, Any], model: type[BaseModel]) ->
|
||||
|
||||
# If this is a BaseModel field, expand its properties with full details
|
||||
if isinstance(field_type, type) and issubclass(field_type, BaseModel):
|
||||
# Skip properties already expanded inline (e.g. via a model_json_schema override).
|
||||
if "properties" in prop_info:
|
||||
continue
|
||||
|
||||
# Get the nested model's schema with all its properties
|
||||
nested_model_schema = field_type.model_json_schema()
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from typing import Literal, Optional
|
||||
from enum import Enum
|
||||
from typing import Any, Literal, Optional
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel, Field
|
||||
@@ -423,3 +424,187 @@ def test_convert_pydantic_with_refs():
|
||||
"name": "Person",
|
||||
}
|
||||
assert tool_spec == expected_spec
|
||||
|
||||
|
||||
def test_convert_pydantic_with_nullable_enum():
|
||||
"""Test that nullable Literal/enum fields preserve anyOf instead of creating invalid type/enum combo."""
|
||||
|
||||
class NullableEnum(BaseModel):
|
||||
sentiment: Literal["positive", "negative"] | None
|
||||
flag: Literal["x"] | None = None
|
||||
|
||||
tool_spec = convert_pydantic_to_tool_spec(NullableEnum)
|
||||
|
||||
expected_spec = {
|
||||
"name": "NullableEnum",
|
||||
"description": "NullableEnum structured output tool",
|
||||
"inputSchema": {
|
||||
"json": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"sentiment": {
|
||||
"anyOf": [
|
||||
{"enum": ["positive", "negative"], "type": "string"},
|
||||
{"type": "null"},
|
||||
],
|
||||
"title": "Sentiment",
|
||||
},
|
||||
"flag": {
|
||||
"anyOf": [
|
||||
{"const": "x", "type": "string"},
|
||||
{"type": "null"},
|
||||
],
|
||||
"default": None,
|
||||
"title": "Flag",
|
||||
},
|
||||
},
|
||||
"title": "NullableEnum",
|
||||
"required": ["sentiment"],
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
assert tool_spec == expected_spec
|
||||
|
||||
# Verify we can construct a valid ToolSpec
|
||||
tool_spec_obj = ToolSpec(**tool_spec)
|
||||
assert tool_spec_obj is not None
|
||||
|
||||
|
||||
def test_convert_pydantic_with_model_json_schema_override():
|
||||
"""Test that a model_json_schema override returning a dereferenced schema is respected."""
|
||||
|
||||
filters_prop = {
|
||||
"filters": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {"name": {"type": "string"}},
|
||||
"required": ["name"],
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
class Filter(BaseModel):
|
||||
name: str
|
||||
|
||||
class FilterGroup(BaseModel):
|
||||
filters: tuple[Filter, ...]
|
||||
|
||||
class Scope(BaseModel):
|
||||
group: FilterGroup
|
||||
|
||||
@classmethod
|
||||
def model_json_schema(cls, *args: Any, **kwargs: Any) -> dict[str, Any]:
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"group": {
|
||||
"type": "object",
|
||||
"properties": filters_prop,
|
||||
"required": ["filters"],
|
||||
}
|
||||
},
|
||||
"required": ["group"],
|
||||
}
|
||||
|
||||
# Should not raise ValueError: Missing reference: Filter
|
||||
tool_spec = convert_pydantic_to_tool_spec(Scope)
|
||||
|
||||
expected_spec = {
|
||||
"name": "Scope",
|
||||
"description": "Scope structured output tool",
|
||||
"inputSchema": {
|
||||
"json": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"group": {
|
||||
"type": "object",
|
||||
"description": "",
|
||||
"properties": filters_prop,
|
||||
"required": ["filters"],
|
||||
}
|
||||
},
|
||||
"required": ["group"],
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
assert tool_spec == expected_spec
|
||||
|
||||
# Verify we can construct a valid ToolSpec
|
||||
tool_spec_obj = ToolSpec(**tool_spec)
|
||||
assert tool_spec_obj is not None
|
||||
|
||||
|
||||
def test_convert_pydantic_with_nullable_str_enum():
|
||||
"""Test that nullable str Enum fields (resolved via $ref) preserve anyOf."""
|
||||
|
||||
class Color(str, Enum):
|
||||
RED = "red"
|
||||
GREEN = "green"
|
||||
BLUE = "blue"
|
||||
|
||||
class Palette(BaseModel):
|
||||
favorite: Color | None = None
|
||||
|
||||
tool_spec = convert_pydantic_to_tool_spec(Palette)
|
||||
|
||||
expected_spec = {
|
||||
"name": "Palette",
|
||||
"description": "Palette structured output tool",
|
||||
"inputSchema": {
|
||||
"json": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"favorite": {
|
||||
"anyOf": [
|
||||
{"enum": ["red", "green", "blue"], "title": "Color", "type": "string"},
|
||||
{"type": "null"},
|
||||
],
|
||||
"default": None,
|
||||
}
|
||||
},
|
||||
"title": "Palette",
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
assert tool_spec == expected_spec
|
||||
|
||||
|
||||
def test_convert_pydantic_with_non_nullable_literal_default():
|
||||
"""Test that a non-nullable Literal with a default does not get null injected into its type."""
|
||||
|
||||
class Config(BaseModel):
|
||||
mode: Literal["fast", "slow"] = "fast"
|
||||
flag: Literal["on"] = "on"
|
||||
|
||||
tool_spec = convert_pydantic_to_tool_spec(Config)
|
||||
|
||||
expected_spec = {
|
||||
"name": "Config",
|
||||
"description": "Config structured output tool",
|
||||
"inputSchema": {
|
||||
"json": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"mode": {
|
||||
"default": "fast",
|
||||
"enum": ["fast", "slow"],
|
||||
"title": "Mode",
|
||||
"type": "string",
|
||||
},
|
||||
"flag": {
|
||||
"const": "on",
|
||||
"default": "on",
|
||||
"title": "Flag",
|
||||
"type": "string",
|
||||
},
|
||||
},
|
||||
"title": "Config",
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
assert tool_spec == expected_spec
|
||||
|
||||
Reference in New Issue
Block a user