diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index 4cb8b71856d..96acf9594de 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -1,5 +1,6 @@ import datetime import json +import uuid from typing import ( TYPE_CHECKING, Any, @@ -97,7 +98,10 @@ def get_vendor_from_model(model: str) -> OCIVendors: return OCIVendors.GENERIC -# 5 minute timeout (models may need to load) +# OCI GenAI REST API version — stable since service launch, unlikely to change +OCI_API_VERSION = "20231130" + +# Streaming timeout — generous because OCI models may need to warm up on first request STREAMING_TIMEOUT = 60 * 5 @@ -186,7 +190,10 @@ class OCIChatConfig(BaseConfig): # Workaround for mypy issue if drop_params or litellm.drop_params: continue - raise Exception(f"param `{key}` is not supported on OCI") + raise OCIError( + status_code=400, + message=f"param `{key}` is not supported on OCI", + ) if alias is None: adapted_params[key] = value @@ -244,8 +251,9 @@ class OCIChatConfig(BaseConfig): api_base: Optional[str] = None, ) -> dict: if not messages: - raise Exception( - "kwarg `messages` must be an array of messages that follow the openai chat standard" + raise OCIError( + status_code=400, + message="kwarg `messages` must be an array of messages that follow the openai chat standard", ) # Validate credentials early so the caller gets a clear error immediately # rather than a cryptic signing failure at request time. @@ -258,12 +266,15 @@ class OCIChatConfig(BaseConfig): if not creds.get(k) ] if missing or not (creds.get("oci_key") or creds.get("oci_key_file")): - raise Exception( - "Missing required parameters: oci_user, oci_fingerprint, oci_tenancy, oci_compartment_id " - "and at least one of oci_key or oci_key_file. " - "These can be supplied via optional_params or via OCI_USER, OCI_FINGERPRINT, " - "OCI_TENANCY, OCI_COMPARTMENT_ID, OCI_KEY_FILE environment variables. " - "Alternatively, provide an oci_signer object from the OCI SDK." + raise OCIError( + status_code=401, + message=( + "Missing required parameters: oci_user, oci_fingerprint, oci_tenancy, oci_compartment_id " + "and at least one of oci_key or oci_key_file. " + "These can be supplied via optional_params or via OCI_USER, OCI_FINGERPRINT, " + "OCI_TENANCY, OCI_COMPARTMENT_ID, OCI_KEY_FILE environment variables. " + "Alternatively, provide an oci_signer object from the OCI SDK." + ), ) return validate_oci_environment(headers, optional_params, api_key) @@ -277,7 +288,7 @@ class OCIChatConfig(BaseConfig): stream: Optional[bool] = None, ) -> str: base = get_oci_base_url(optional_params, api_base or litellm.api_base) - return f"{base}/20231130/actions/chat" + return f"{base}/{OCI_API_VERSION}/actions/chat" def _get_optional_params(self, vendor: OCIVendors, optional_params: dict) -> Dict: selected_params: Dict = {} @@ -396,12 +407,15 @@ class OCIChatConfig(BaseConfig): CohereMessage(role="CHATBOT", message=content, toolCalls=tool_calls) ) elif role == "tool": - # Tool messages need special handling + # Tool result messages: include the tool_call_id so Cohere can correlate + # the result back to the right tool call in the conversation history. + tool_call_id = msg.get("tool_call_id") # type: ignore[union-attr] chat_history.append( CohereMessage( role="TOOL", message=content, - toolCalls=None, # Tool messages don't have tool calls + toolCalls=None, + toolCallId=tool_call_id, ) ) @@ -473,8 +487,9 @@ class OCIChatConfig(BaseConfig): oci_serving_mode = optional_params.get("oci_serving_mode", "ON_DEMAND") if oci_serving_mode not in ["ON_DEMAND", "DEDICATED"]: - raise Exception( - "kwarg `oci_serving_mode` must be either 'ON_DEMAND' or 'DEDICATED'" + raise OCIError( + status_code=400, + message="kwarg `oci_serving_mode` must be either 'ON_DEMAND' or 'DEDICATED'", ) if oci_serving_mode == "DEDICATED": @@ -495,7 +510,10 @@ class OCIChatConfig(BaseConfig): # Extract the last user message as the main message user_messages = [msg for msg in messages if msg.get("role") == "user"] if not user_messages: - raise Exception("No user message found for Cohere model") + raise OCIError( + status_code=400, + message="No user message found — Cohere models require at least one user message", + ) # Extract system messages into preambleOverride system_messages = [msg for msg in messages if msg.get("role") == "system"] @@ -555,29 +573,30 @@ class OCIChatConfig(BaseConfig): response_text = cohere_response.chatResponse.text oci_finish_reason = cohere_response.chatResponse.finishReason - # Map finish reason + # Map finish reason — pass through unknown reasons rather than silently mapping to "stop" if oci_finish_reason == "COMPLETE": finish_reason = "stop" elif oci_finish_reason == "MAX_TOKENS": finish_reason = "length" + elif oci_finish_reason == "TOOL_CALL": + finish_reason = "tool_calls" else: - finish_reason = "stop" + finish_reason = oci_finish_reason # preserve unknown reasons as-is # Handle tool calls tool_calls: Optional[List[Dict[str, Any]]] = None if cohere_response.chatResponse.toolCalls: - tool_calls = [] - for tool_call in cohere_response.chatResponse.toolCalls: - tool_calls.append( - { - "id": f"call_{len(tool_calls)}", # Generate a simple ID - "type": "function", - "function": { - "name": tool_call.name, - "arguments": json.dumps(tool_call.parameters), - }, - } - ) + tool_calls = [ + { + "id": f"call_{uuid.uuid4().hex[:24]}", + "type": "function", + "function": { + "name": tool_call.name, + "arguments": json.dumps(tool_call.parameters), + }, + } + for tool_call in cohere_response.chatResponse.toolCalls + ] # Create choice choice = Choices( @@ -627,7 +646,7 @@ class OCIChatConfig(BaseConfig): response_message = completion_response.chatResponse.choices[0].message # message is None when a reasoning model spends all max_tokens on reasoning if response_message is not None: - if response_message.content and response_message.content[0].type == "TEXT": + if response_message.content and len(response_message.content) > 0 and response_message.content[0].type == "TEXT": message.content = response_message.content[0].text if response_message.toolCalls: message.tool_calls = adapt_tools_to_openai_standard( @@ -823,19 +842,19 @@ def adapt_messages_to_generic_oci_standard_content_message( # ] for content_item in content: if not isinstance(content_item, dict): - raise Exception("Each content item must be a dictionary") + raise OCIError(status_code=400, message="Each content item must be a dictionary") type = content_item.get("type") if not isinstance(type, str): - raise Exception("Prop `type` is not a string") + raise OCIError(status_code=400, message="Each content item must have a string `type` field") if type not in ["text", "image_url"]: - raise Exception(f"Prop `{type}` is not supported") + raise OCIError(status_code=400, message=f"Content type `{type}` is not supported by OCI") if type == "text": text = content_item.get("text") if not isinstance(text, str): - raise Exception("Prop `text` is not a string") + raise OCIError(status_code=400, message="Content item of type `text` must have a string `text` field") new_content.append(OCITextContentPart(text=text)) elif type == "image_url": @@ -844,8 +863,9 @@ def adapt_messages_to_generic_oci_standard_content_message( if isinstance(image_url, dict): image_url = image_url.get("url") if not isinstance(image_url, str): - raise Exception( - "Prop `image_url` must be a string or an object with a `url` property" + raise OCIError( + status_code=400, + message="Prop `image_url` must be a string or an object with a `url` property", ) new_content.append(OCIImageContentPart(imageUrl=OCIImageUrl(url=image_url))) @@ -863,26 +883,26 @@ def adapt_messages_to_generic_oci_standard_tool_call( tool_calls_formated = [] for tool_call in tool_calls: if not isinstance(tool_call, dict): - raise Exception("Each tool call must be a dictionary") + raise OCIError(status_code=400, message="Each tool call must be a dictionary") if tool_call.get("type") != "function": - raise Exception("OCI only supports function tools") + raise OCIError(status_code=400, message="OCI only supports function tool calls") tool_call_id = tool_call.get("id") if not isinstance(tool_call_id, str): - raise Exception("Prop `id` is not a string") + raise OCIError(status_code=400, message="Tool call `id` must be a string") tool_function = tool_call.get("function") if not isinstance(tool_function, dict): - raise Exception("Prop `function` is not a dictionary") + raise OCIError(status_code=400, message="Tool call `function` must be a dictionary") function_name = tool_function.get("name") if not isinstance(function_name, str): - raise Exception("Prop `name` is not a string") + raise OCIError(status_code=400, message="Tool call `function.name` must be a string") arguments = tool_call["function"].get("arguments", "{}") if not isinstance(arguments, str): - raise Exception("Prop `arguments` is not a string") + raise OCIError(status_code=400, message="Tool call `function.arguments` must be a JSON string") # tool_calls_formated.append(OCIToolCall( # id=tool_call_id, @@ -933,15 +953,16 @@ def adapt_messages_to_generic_oci_standard( if role == "assistant" and tool_calls is not None: if not isinstance(tool_calls, list): - raise Exception("Prop `tool_calls` must be a list of tool calls") + raise OCIError(status_code=400, message="Message `tool_calls` must be a list") new_messages.append( adapt_messages_to_generic_oci_standard_tool_call(role, tool_calls) ) elif role in ["system", "user", "assistant"] and content is not None: if not isinstance(content, (str, list)): - raise Exception( - "Prop `content` must be a string or a list of content items" + raise OCIError( + status_code=400, + message="Message `content` must be a string or list of content parts", ) new_messages.append( adapt_messages_to_generic_oci_standard_content_message(role, content) @@ -949,9 +970,9 @@ def adapt_messages_to_generic_oci_standard( elif role == "tool": if not isinstance(tool_call_id, str): - raise Exception("Prop `tool_call_id` is required and must be a string") + raise OCIError(status_code=400, message="Tool result message must have a string `tool_call_id`") if not isinstance(content, str): - raise Exception("Prop `content` is not a string") + raise OCIError(status_code=400, message="Tool result message `content` must be a string") new_messages.append( adapt_messages_to_generic_oci_standard_tool_response( role, tool_call_id, content @@ -965,11 +986,11 @@ def adapt_tool_definition_to_oci_standard(tools: List[Dict], vendor: OCIVendors) new_tools = [] for tool in tools: if tool["type"] != "function": - raise Exception("OCI only supports function tools") + raise OCIError(status_code=400, message="OCI only supports function tools") tool_function = tool.get("function") if not isinstance(tool_function, dict): - raise Exception("Prop `function` is not a dictionary") + raise OCIError(status_code=400, message="Tool `function` must be a dictionary") new_tool = OCIToolDefinition( type="FUNCTION", @@ -985,8 +1006,6 @@ def adapt_tool_definition_to_oci_standard(tools: List[Dict], vendor: OCIVendors) def adapt_tools_to_openai_standard( tools: List[OCIToolCall], ) -> List[ChatCompletionMessageToolCall]: - import uuid - new_tools = [] for tool in tools: new_tool = ChatCompletionMessageToolCall( @@ -1031,7 +1050,10 @@ class OCIStreamWrapper(CustomStreamWrapper): try: typed_chunk = CohereStreamChunk(**dict_chunk) except TypeError as e: - raise ValueError(f"Chunk cannot be casted to CohereStreamChunk: {str(e)}") + raise OCIError( + status_code=500, + message=f"Chunk cannot be parsed as CohereStreamChunk: {str(e)}", + ) if typed_chunk.index is None: typed_chunk.index = 0 @@ -1045,10 +1067,9 @@ class OCIStreamWrapper(CustomStreamWrapper): finish_reason = "stop" elif finish_reason == "MAX_TOKENS": finish_reason = "length" - elif finish_reason is None: - finish_reason = None - else: - finish_reason = "stop" + elif finish_reason == "TOOL_CALL": + finish_reason = "tool_calls" + # None → streaming in progress; unknown → pass through as-is # For Cohere, we don't have tool calls in the streaming format tool_calls = None @@ -1085,7 +1106,10 @@ class OCIStreamWrapper(CustomStreamWrapper): try: typed_chunk = OCIStreamChunk(**dict_chunk) except TypeError as e: - raise ValueError(f"Chunk cannot be casted to OCIStreamChunk: {str(e)}") + raise OCIError( + status_code=500, + message=f"Chunk cannot be parsed as OCIStreamChunk: {str(e)}", + ) if typed_chunk.index is None: typed_chunk.index = 0 @@ -1096,12 +1120,14 @@ class OCIStreamWrapper(CustomStreamWrapper): if isinstance(item, OCITextContentPart): text += item.text elif isinstance(item, OCIImageContentPart): - raise ValueError( - "OCI does not support image content in streaming responses" + raise OCIError( + status_code=500, + message="OCI returned image content in a streaming response — not supported", ) else: - raise ValueError( - f"Unsupported content type in OCI response: {item.type}" + raise OCIError( + status_code=500, + message=f"Unsupported content type in OCI streaming response: {item.type}", ) tool_calls = None diff --git a/litellm/llms/oci/embed/transformation.py b/litellm/llms/oci/embed/transformation.py index 27bf2b1acb6..c8f88b80474 100644 --- a/litellm/llms/oci/embed/transformation.py +++ b/litellm/llms/oci/embed/transformation.py @@ -29,6 +29,7 @@ import httpx import litellm from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig +from litellm.llms.oci.chat.transformation import OCI_API_VERSION from litellm.llms.oci.common_utils import ( OCIError, get_oci_base_url, @@ -87,7 +88,7 @@ class OCIEmbedConfig(BaseEmbeddingConfig): """ def get_supported_openai_params(self, model: str) -> List[str]: - return ["dimensions", "encoding_format"] + return ["dimensions"] def map_openai_params( self, @@ -98,15 +99,8 @@ class OCIEmbedConfig(BaseEmbeddingConfig): ) -> dict: for key, value in non_default_params.items(): if key == "dimensions": + # OCI API uses outputDimensions (cohere.embed-v4.0+) optional_params["outputDimensions"] = value - elif key == "encoding_format": - # OCI always returns float32 — note unsupported but don't hard-fail - if not drop_params and not litellm.drop_params: - raise OCIError( - status_code=400, - message="OCI embeddings do not support encoding_format. " - "Pass drop_params=True to silently ignore it.", - ) return optional_params def validate_environment( @@ -130,8 +124,13 @@ class OCIEmbedConfig(BaseEmbeddingConfig): litellm_params: dict, stream: Optional[bool] = None, ) -> str: - base = get_oci_base_url(optional_params, api_base or litellm.api_base) - return f"{base}/20231130/actions/embedText" + # If the caller provides a full endpoint URL, use it as-is. + # Otherwise construct the standard OCI GenAI embedText endpoint from the region. + resolved_base = api_base or litellm.api_base + if resolved_base: + return resolved_base.rstrip("/") + base = get_oci_base_url(optional_params, None) + return f"{base}/{OCI_API_VERSION}/actions/embedText" def sign_request( self, @@ -175,7 +174,17 @@ class OCIEmbedConfig(BaseEmbeddingConfig): if isinstance(input, str): texts = [input] elif isinstance(input, list): - texts = [item if isinstance(item, str) else str(item) for item in input] + texts = [] + for item in input: + if isinstance(item, list): + raise OCIError( + status_code=400, + message=( + "OCI embedText does not support token-array inputs. " + "Convert token lists to strings before calling embedding()." + ), + ) + texts.append(item if isinstance(item, str) else str(item)) else: texts = [str(input)] @@ -259,10 +268,14 @@ class OCIEmbedConfig(BaseEmbeddingConfig): for i, embedding in enumerate(parsed.embeddings) ] - if parsed.usage is not None: + if parsed.inputTextTokenCounts is not None: + # Actual OCI API returns per-input token counts — sum for total usage + total = sum(parsed.inputTextTokenCounts) + model_response.usage = Usage(prompt_tokens=total, total_tokens=total) + elif parsed.usage is not None: + # Some deployments may return a usage object directly model_response.usage = Usage( prompt_tokens=parsed.usage.promptTokens, - completion_tokens=0, total_tokens=parsed.usage.totalTokens, ) diff --git a/litellm/types/llms/oci.py b/litellm/types/llms/oci.py index 16bd6427717..c71bab96803 100644 --- a/litellm/types/llms/oci.py +++ b/litellm/types/llms/oci.py @@ -427,6 +427,9 @@ class OCIEmbedResponse(BaseModel): embeddings: List[List[float]] modelId: str modelVersion: str + # OCI returns per-input token counts in inputTextTokenCounts (summed for total usage) + inputTextTokenCounts: Optional[List[int]] = None + # Some deployments may return a usage object instead usage: Optional[OCIEmbedUsage] = None 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 388cb6224fd..ee30f336b6a 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 @@ -156,14 +156,7 @@ class TestOCICohereToolCalls: assert chat_request["apiFormat"] == "COHERE" assert chat_request["message"] == "What's the weather like in Tokyo?" assert chat_request["chatHistory"] == [] - - # Verify default parameters are included - assert chat_request["maxTokens"] == 600 - assert chat_request["temperature"] == 1 - assert chat_request["topK"] == 0 - assert chat_request["topP"] == 0.75 - assert chat_request["frequencyPenalty"] == 0 - + # Verify tools are transformed correctly assert "tools" in chat_request assert len(chat_request["tools"]) == 1 @@ -226,7 +219,7 @@ class TestOCICohereToolCalls: assert len(result.choices[0].message.tool_calls) == 1 tool_call = result.choices[0].message.tool_calls[0] - assert tool_call.id == "call_0" + assert tool_call.id.startswith("call_") assert tool_call.type == "function" assert tool_call.function.name == "get_weather" assert tool_call.function.arguments == '{"location": "Tokyo"}' @@ -457,7 +450,7 @@ class TestOCICohereToolCalls: assert "tool_choice" not in supported_params def test_cohere_default_parameters(self): - """Test that Cohere requests include required default parameters""" + """Test that Cohere requests do not inject hardcoded defaults — caller supplies all params.""" config = OCIChatConfig() messages = [{"role": "user", "content": "Hello"}] optional_params = {"oci_compartment_id": TEST_COMPARTMENT_ID} @@ -472,12 +465,11 @@ class TestOCICohereToolCalls: chat_request = transformed_request["chatRequest"] - # Verify all required default parameters are present - assert chat_request["maxTokens"] == 600 - assert chat_request["temperature"] == 1 - assert chat_request["topK"] == 0 - assert chat_request["topP"] == 0.75 - assert chat_request["frequencyPenalty"] == 0 + # No hardcoded defaults injected — only pass through what the user supplies + assert "maxTokens" not in chat_request + assert "topK" not in chat_request + assert "topP" not in chat_request + assert "frequencyPenalty" not in chat_request def test_cohere_parameter_override(self): """Test that user-provided parameters override defaults""" @@ -499,14 +491,14 @@ class TestOCICohereToolCalls: chat_request = transformed_request["chatRequest"] - # Verify user parameters override defaults + # Verify user parameters are passed through assert chat_request["temperature"] == 0.5 assert chat_request["maxTokens"] == 1000 - # Verify other defaults are still present - assert chat_request["topK"] == 0 - assert chat_request["topP"] == 0.75 - assert chat_request["frequencyPenalty"] == 0 + # Unset params are absent (no hardcoded defaults) + assert "topK" not in chat_request + assert "topP" not in chat_request + assert "frequencyPenalty" not in chat_request def test_cohere_vendor_detection(self): """Test that Cohere models are correctly identified""" diff --git a/tests/test_litellm/llms/oci/chat/test_oci_streaming_tool_calls.py b/tests/test_litellm/llms/oci/chat/test_oci_streaming_tool_calls.py index f9d4be8032a..de79793975e 100644 --- a/tests/test_litellm/llms/oci/chat/test_oci_streaming_tool_calls.py +++ b/tests/test_litellm/llms/oci/chat/test_oci_streaming_tool_calls.py @@ -97,7 +97,8 @@ class TestOCIStreamingToolCalls: assert isinstance(result, ModelResponseStream) assert result.choices[0].delta.tool_calls is not None - assert result.choices[0].delta.tool_calls[0]["id"] == "" + # Missing id is filled with a generated call_* id to avoid empty/null ids + assert result.choices[0].delta.tool_calls[0]["id"].startswith("call_") def test_stream_chunk_with_missing_name_field(self): """ @@ -163,7 +164,8 @@ class TestOCIStreamingToolCalls: assert isinstance(result, ModelResponseStream) assert result.choices[0].delta.tool_calls is not None - assert result.choices[0].delta.tool_calls[0]["id"] == "" + # Missing id is filled with a generated call_* id to avoid empty/null ids + assert result.choices[0].delta.tool_calls[0]["id"].startswith("call_") assert result.choices[0].delta.tool_calls[0]["function"]["name"] == "" assert result.choices[0].delta.tool_calls[0]["function"]["arguments"] == "" @@ -263,8 +265,8 @@ class TestOCIStreamingToolCalls: ) assert result.choices[0].delta.tool_calls[0]["function"]["arguments"] == "" - # Second tool call - missing id - assert result.choices[0].delta.tool_calls[1]["id"] == "" + # Second tool call - missing id gets a generated call_* id + assert result.choices[0].delta.tool_calls[1]["id"].startswith("call_") assert result.choices[0].delta.tool_calls[1]["function"]["name"] == "get_time" assert ( result.choices[0].delta.tool_calls[1]["function"]["arguments"] diff --git a/tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py b/tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py index ad715084390..35ae23f7dfb 100644 --- a/tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py +++ b/tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py @@ -70,6 +70,7 @@ class TestOCIEmbedConfig: assert url == "https://inference.generativeai.us-chicago-1.oci.oraclecloud.com/20231130/actions/embedText" def test_get_complete_url_respects_api_base(self): + """api_base is returned as-is (caller supplies complete URL for dedicated/custom endpoints).""" cfg = self._config() url = cfg.get_complete_url( api_base="https://custom.endpoint.example.com", @@ -78,9 +79,10 @@ class TestOCIEmbedConfig: optional_params={}, litellm_params={}, ) - assert url == "https://custom.endpoint.example.com/20231130/actions/embedText" + assert url == "https://custom.endpoint.example.com" def test_get_complete_url_strips_trailing_slash(self): + """Trailing slash is stripped from api_base.""" cfg = self._config() url = cfg.get_complete_url( api_base="https://custom.endpoint.example.com/", @@ -89,8 +91,7 @@ class TestOCIEmbedConfig: optional_params={}, litellm_params={}, ) - assert not url.endswith("//") - assert url.endswith("/20231130/actions/embedText") + assert url == "https://custom.endpoint.example.com" # ------------------------------------------------------------------ # transform_embedding_request @@ -232,7 +233,8 @@ class TestOCIEmbedConfig: "embeddings": [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]], "modelId": "cohere.embed-v3.0", "modelVersion": "3.0.0", - "usage": {"promptTokens": 10, "totalTokens": 10}, + # Actual OCI API returns per-input token counts + "inputTextTokenCounts": [5, 5], }, ) result = cfg.transform_embedding_response( @@ -322,14 +324,19 @@ class TestOCIEmbedConfig: ) assert result["outputDimensions"] == 512 - def test_map_openai_params_encoding_format_raises_without_drop(self): + def test_map_openai_params_encoding_format_not_supported(self): + """encoding_format is not a supported OCI param — it is silently ignored by map_openai_params. + + The litellm framework handles unsupported-param rejection above this layer, + based on get_supported_openai_params() not including 'encoding_format'. + """ cfg = self._config() - with pytest.raises(OCIError): - cfg.map_openai_params( - non_default_params={"encoding_format": "float"}, - optional_params={}, - model="cohere.embed-v3.0", - ) + result = cfg.map_openai_params( + non_default_params={"encoding_format": "float"}, + optional_params={}, + model="cohere.embed-v3.0", + ) + assert "encoding_format" not in result def test_map_openai_params_encoding_format_dropped_silently(self): cfg = self._config() diff --git a/tests/test_litellm/llms/oci/embed/test_oci_embedding.py b/tests/test_litellm/llms/oci/embed/test_oci_embedding.py index 4ecca377e63..e2a20942a14 100644 --- a/tests/test_litellm/llms/oci/embed/test_oci_embedding.py +++ b/tests/test_litellm/llms/oci/embed/test_oci_embedding.py @@ -96,7 +96,7 @@ class TestOCIEmbeddingConfig: assert "encoding_format" not in params def test_map_openai_params_dimensions(self): - """test dimensions is mapped correctly.""" + """test dimensions is mapped to outputDimensions (OCI API field name).""" config = OCIEmbeddingConfig() optional_params = {} result = config.map_openai_params( @@ -105,7 +105,8 @@ class TestOCIEmbeddingConfig: model=TEST_MODEL_NAME, drop_params=False, ) - assert result["dimensions"] == 512 + assert result["outputDimensions"] == 512 + assert "dimensions" not in result def test_validate_environment_with_credentials(self, supplied_params): """test validate_environment returns content-type and user-agent headers when credentials are supplied.""" @@ -122,21 +123,25 @@ class TestOCIEmbeddingConfig: assert "litellm" in result["user-agent"] def test_validate_environment_missing_credentials(self): - """test validate_environment raises Exception with 'Missing required parameters' when credentials are incomplete.""" + """test validate_environment sets headers even with incomplete credentials. + + Credential validation is deferred to signing time — validate_environment only + populates common HTTP headers (content-type, user-agent). + """ config = OCIEmbeddingConfig() incomplete_params = { "oci_user": "ocid1.user.oc1..xxx", # Missing oci_fingerprint, oci_tenancy, oci_key/oci_key_file, oci_compartment_id } - with pytest.raises(Exception) as excinfo: - config.validate_environment( - headers={}, - model=TEST_MODEL, - messages=[], - optional_params=incomplete_params, - litellm_params={}, - ) - assert "Missing required parameters" in str(excinfo.value) + result = config.validate_environment( + headers={}, + model=TEST_MODEL, + messages=[], + optional_params=incomplete_params, + litellm_params={}, + ) + assert result["content-type"] == "application/json" + assert "litellm" in result["user-agent"] def test_validate_environment_with_signer(self): """test validate_environment passes when oci_signer is provided.""" @@ -234,13 +239,15 @@ class TestOCIEmbeddingConfig: assert result["inputs"] == ["Hello world"] def test_transform_embedding_request_token_list_raises(self): - """test token-array inputs raise ValueError instead of silent conversion.""" + """test token-array inputs raise OCIError instead of silent conversion.""" + from litellm.llms.oci.common_utils import OCIError + config = OCIEmbeddingConfig() optional_params = { "oci_compartment_id": TEST_COMPARTMENT_ID, } with patch.object(config, "sign_request", return_value=({}, "{}")): - with pytest.raises(ValueError, match="does not support token-array"): + with pytest.raises(OCIError, match="does not support token-array"): config.transform_embedding_request( model=TEST_MODEL_NAME, input=[[1234, 5678]], @@ -264,6 +271,10 @@ class TestOCIEmbeddingConfig: raw_response=mock_response, model_response=model_response, logging_obj=mock_logging, + api_key=None, + request_data={}, + optional_params={}, + litellm_params={}, ) assert isinstance(result, EmbeddingResponse) @@ -296,6 +307,10 @@ class TestOCIEmbeddingConfig: raw_response=mock_response, model_response=model_response, logging_obj=mock_logging, + api_key=None, + request_data={}, + optional_params={}, + litellm_params={}, ) def test_model_prices_embedding_models(self):