Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
111 changes: 79 additions & 32 deletions mellea/backends/tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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",
Expand Down
Loading
Loading