mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
68b4b699bc
commit
9bb707a8f4
3 changed files with 199 additions and 12 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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 ""
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue