fix(oci): port schema/type utilities from langchain-oracle reference impl

- Add resolve_oci_schema_refs: inline $ref/$defs — OCI rejects JSON Schema refs
- Add resolve_oci_schema_anyof: flatten Optional[T] anyOf (Pydantic v2 emits these)
- Add sanitize_oci_schema: strip title, normalise null types, ensure array items
- Add OCI_JSON_TO_PYTHON_TYPES: Cohere expects Python type names (str/int/float),
  not JSON Schema names (string/integer/number)
- Add enrich_cohere_param_description: embed enum/format/range/pattern constraints
  into description since CohereParameterDefinition has no dedicated fields
- Apply all of the above in adapt_tool_definitions_to_cohere_standard and
  adapt_tool_definition_to_oci_standard
- Fix toolChoice conversion: map OpenAI string ('auto','none','required') to OCI
  dict form ({"type":"AUTO"} etc.) — the API rejects plain strings
- Update unit test expectations to match correct Python type names and enriched
  descriptions
This commit is contained in:
Federico Kamelhar 2026-04-05 12:31:11 -04:00
parent 68b4b699bc
commit 9bb707a8f4
3 changed files with 199 additions and 12 deletions

View file

@ -26,10 +26,15 @@ from litellm.llms.custom_httpx.http_handler import (
version,
)
from litellm.llms.oci.common_utils import (
OCI_JSON_TO_PYTHON_TYPES,
OCIError,
OCIRequestWrapper, # re-exported for backwards compatibility
enrich_cohere_param_description,
get_oci_base_url,
resolve_oci_credentials,
resolve_oci_schema_anyof,
resolve_oci_schema_refs,
sanitize_oci_schema,
sign_oci_request,
validate_oci_environment,
)
@ -324,6 +329,21 @@ class OCIChatConfig(BaseConfig):
selected_params["tools"], vendor # type: ignore[arg-type]
)
# Convert toolChoice from OpenAI string form ("auto", "none", "required") to the
# OCI dict form ({"type": "AUTO"} etc.) — the API rejects plain strings.
if "toolChoice" in selected_params:
tc = selected_params["toolChoice"]
if isinstance(tc, str):
tc_map = {
"auto": {"type": "AUTO"},
"none": {"type": "NONE"},
"required": {"type": "REQUIRED"},
"any": {"type": "REQUIRED"},
}
selected_params["toolChoice"] = tc_map.get(
tc.lower(), {"type": "FUNCTION", "name": tc}
)
# Transform response_format type to OCI uppercase format
if "responseFormat" in selected_params:
rf = selected_params["responseFormat"]
@ -458,18 +478,34 @@ class OCIChatConfig(BaseConfig):
def adapt_tool_definitions_to_cohere_standard(
self, tools: List[Dict[str, Any]]
) -> List[CohereTool]:
"""Adapt tool definitions to Cohere format."""
"""Adapt tool definitions to the OCI Cohere format.
- Resolves ``$ref``/``$defs`` and ``anyOf`` patterns that OCI rejects.
- Maps JSON Schema type names to Python type names (``"string"`` → ``"str"``).
- Embeds unsupported constraints (enum, format, range, pattern) into the
parameter description so the model can still see them.
"""
cohere_tools = []
for tool in tools:
function_def = tool.get("function", {})
parameters = function_def.get("parameters", {}).get("properties", {})
required = function_def.get("parameters", {}).get("required", [])
raw_params = function_def.get("parameters", {})
# Resolve schema extensions that OCI Cohere doesn't support
resolved = sanitize_oci_schema(
resolve_oci_schema_anyof(resolve_oci_schema_refs(raw_params))
)
properties = resolved.get("properties", {})
required = resolved.get("required", [])
parameter_definitions = {}
for param_name, param_schema in parameters.items():
for param_name, param_schema in properties.items():
json_type = param_schema.get("type", "string")
python_type = OCI_JSON_TO_PYTHON_TYPES.get(json_type, json_type)
parameter_definitions[param_name] = CohereParameterDefinition(
description=param_schema.get("description", ""),
type=param_schema.get("type", "string"),
description=enrich_cohere_param_description(
param_schema.get("description", ""), param_schema
),
type=python_type,
isRequired=param_name in required,
)
@ -1026,11 +1062,17 @@ def adapt_tool_definition_to_oci_standard(tools: List[Dict], vendor: OCIVendors)
if not isinstance(tool_function, dict):
raise OCIError(status_code=400, message="Tool `function` must be a dictionary")
# Resolve $ref/$defs and anyOf that OCI GenericChatRequest doesn't support
raw_params = tool_function.get("parameters", {})
resolved_params = sanitize_oci_schema(
resolve_oci_schema_anyof(resolve_oci_schema_refs(raw_params))
)
new_tool = OCIToolDefinition(
type="FUNCTION",
name=tool_function.get("name"),
description=tool_function.get("description", ""),
parameters=tool_function.get("parameters", {}),
parameters=resolved_params,
)
new_tools.append(new_tool)

View file

@ -4,7 +4,7 @@ import hashlib
import json
import os
from dataclasses import dataclass
from typing import Any, Dict, Optional, Protocol, Tuple
from typing import Any, Dict, List, Optional, Protocol, Tuple
from urllib.parse import urlparse
import httpx
@ -382,3 +382,148 @@ def validate_oci_environment(
headers.setdefault("content-type", "application/json")
headers.setdefault("user-agent", f"litellm/{version}")
return headers
# ---------------------------------------------------------------------------
# JSON schema utilities for OCI tool definitions
#
# OCI Generative AI does not support JSON Schema extensions ($ref, $defs,
# anyOf). Pydantic v2 emits all three for models with Optional fields or
# nested schemas. The helpers below are ported from the official
# langchain-oracle reference implementation so that tool schemas are always
# valid before they reach the OCI endpoint.
# ---------------------------------------------------------------------------
# Mapping from JSON Schema type names to Python type names, as expected by
# the OCI Cohere API's CohereParameterDefinition.type field.
OCI_JSON_TO_PYTHON_TYPES: Dict[str, str] = {
"string": "str",
"number": "float",
"boolean": "bool",
"integer": "int",
"array": "List",
"object": "Dict",
"any": "any",
}
def resolve_oci_schema_refs(schema: Dict[str, Any]) -> Dict[str, Any]:
"""Inline all ``$ref``/``$defs`` references — OCI does not support JSON Schema ``$ref``."""
defs = schema.get("$defs", {})
resolving_stack: set = set()
def _resolve(obj: Any) -> Any:
if isinstance(obj, dict):
if "$ref" in obj:
ref = obj["$ref"]
if ref.startswith("#/$defs/"):
key = ref.split("/")[-1]
if key in resolving_stack:
return {"type": "object"} # break cycles
resolving_stack.add(key)
try:
return _resolve(defs.get(key, obj))
finally:
resolving_stack.discard(key)
return obj # external $ref — leave unchanged
return {k: _resolve(v) for k, v in obj.items()}
if isinstance(obj, list):
return [_resolve(item) for item in obj]
return obj
resolved = _resolve(schema)
if isinstance(resolved, dict):
resolved.pop("$defs", None)
return resolved
def resolve_oci_schema_anyof(obj: Any) -> Any:
"""Resolve Pydantic v2 ``Optional[T]`` → ``anyOf`` patterns.
Pydantic v2 emits ``{"anyOf": [{"type": "T"}, {"type": "null"}]}`` for
``Optional[T]``. OCI models don't understand ``anyOf``, so we pick the
first non-null branch and merge top-level metadata into it.
"""
if isinstance(obj, dict):
if "anyOf" in obj and "type" not in obj:
non_null = [
t for t in obj["anyOf"]
if not (isinstance(t, dict) and t.get("type") == "null")
]
if non_null:
resolved = {**obj, **non_null[0]}
resolved.pop("anyOf", None)
return resolve_oci_schema_anyof(resolved)
return {k: resolve_oci_schema_anyof(v) for k, v in obj.items()}
if isinstance(obj, list):
return [resolve_oci_schema_anyof(item) for item in obj]
return obj
def sanitize_oci_schema(schema: Any) -> Any:
"""Recursively remove OCI-incompatible fields from a JSON schema.
Strips ``title`` keys, removes ``None``-valued ``default`` entries,
normalises ``type: [T, "null"]`` list types, and ensures arrays carry an
``items`` definition.
"""
if isinstance(schema, list):
return [sanitize_oci_schema(item) for item in schema]
if not isinstance(schema, dict):
return schema
sanitized: Dict[str, Any] = {}
for key, value in schema.items():
if key == "title":
continue
if key == "default" and value is None:
continue
if key == "type":
if value == "any":
sanitized[key] = "object"
continue
if isinstance(value, list):
non_null = [t for t in value if t != "null"]
sanitized[key] = non_null[0] if non_null else "string"
continue
sanitized[key] = sanitize_oci_schema(value)
if sanitized.get("type") == "array" and "items" not in sanitized:
sanitized["items"] = {"type": "object"}
required = sanitized.get("required")
properties = sanitized.get("properties")
if "required" in sanitized:
if isinstance(required, list) and isinstance(properties, dict):
sanitized["required"] = [
f for f in required if isinstance(f, str) and f in properties
]
elif not isinstance(required, list):
sanitized["required"] = []
return sanitized
def enrich_cohere_param_description(description: str, param_schema: Dict[str, Any]) -> str:
"""Embed schema constraints into a Cohere parameter description.
``CohereParameterDefinition`` only has ``type``, ``description``, and
``isRequired``. Rich constraints (``enum``, ``format``, ``minimum``,
``maximum``, ``pattern``) are appended to the description string so the
model can still see and respect them.
"""
parts = [description] if description else []
if "enum" in param_schema:
parts.append(f"Allowed values: {param_schema['enum']}")
if "format" in param_schema:
parts.append(f"Format: {param_schema['format']}")
if "minimum" in param_schema or "maximum" in param_schema:
range_parts = []
if "minimum" in param_schema:
range_parts.append(f"min={param_schema['minimum']}")
if "maximum" in param_schema:
range_parts.append(f"max={param_schema['maximum']}")
parts.append(f"Range: {', '.join(range_parts)}")
if "pattern" in param_schema:
parts.append(f"Pattern: {param_schema['pattern']}")
return ". ".join(parts) if parts else ""

View file

@ -90,13 +90,13 @@ class TestOCICohereToolCalls:
# Check location parameter
location_param = weather_tool.parameterDefinitions["location"]
assert location_param.description == "The city or location to get weather for"
assert location_param.type == "string"
assert location_param.type == "str"
assert location_param.isRequired == True
# Check unit parameter
unit_param = weather_tool.parameterDefinitions["unit"]
assert unit_param.description == "Temperature unit (celsius or fahrenheit)"
assert unit_param.type == "string"
assert unit_param.description == "Temperature unit (celsius or fahrenheit). Allowed values: ['celsius', 'fahrenheit']"
assert unit_param.type == "str"
assert unit_param.isRequired == False
# Check second tool
@ -107,7 +107,7 @@ class TestOCICohereToolCalls:
expression_param = calc_tool.parameterDefinitions["expression"]
assert expression_param.description == "Mathematical expression to evaluate"
assert expression_param.type == "string"
assert expression_param.type == "str"
assert expression_param.isRequired == True
def test_cohere_request_with_tools(self):