From 9bb707a8f4de83b2821cb27ba9e3a33c15a2572c Mon Sep 17 00:00:00 2001 From: Federico Kamelhar Date: Sun, 5 Apr 2026 12:31:11 -0400 Subject: [PATCH] fix(oci): port schema/type utilities from langchain-oracle reference impl MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 --- litellm/llms/oci/chat/transformation.py | 56 ++++++- litellm/llms/oci/common_utils.py | 147 +++++++++++++++++- .../oci/chat/test_oci_cohere_tool_calls.py | 8 +- 3 files changed, 199 insertions(+), 12 deletions(-) diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index 06518feb446..3d949afb91a 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -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) diff --git a/litellm/llms/oci/common_utils.py b/litellm/llms/oci/common_utils.py index 2c95f212576..931fb344ddc 100644 --- a/litellm/llms/oci/common_utils.py +++ b/litellm/llms/oci/common_utils.py @@ -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 "" diff --git a/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py b/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py index ee30f336b6a..3a6d945fc54 100644 --- a/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py +++ b/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py @@ -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):