fix(oci): comprehensive code quality pass — bugs, tests, schema accuracy

- Fix Cohere tool call IDs (was always call_0; now UUID per call)
- Fix TOOL_CALL finish reason mapping in both sync and streaming paths
- Fix Cohere stop parameter mapping (stop → stopSequences)
- Remove hardcoded Cohere defaults (maxTokens/topK/topP/frequencyPenalty)
- Fix content[0] safety guard against empty content arrays
- Fix streaming signed body used consistently (not re-serialized)
- Raise OCIError (not bare Exception/ValueError) throughout
- Centralize OCI_API_VERSION constant; import uuid at module level
- Fix embed get_complete_url to strip trailing slashes from api_base
- Fix OCIEmbedResponse schema: add inputTextTokenCounts (actual OCI field)
- Fix embed usage computed from inputTextTokenCounts (sum of per-input counts)
- Fix Cohere toolCallId included in tool result messages
- Add OCIToolCall.id as Optional (absent in Google/xAI streaming chunks)
- Update tests to reflect correct behavior (no hardcoded defaults, UUID ids,
  deferred credential validation, OCIError vs ValueError, real response schema)
This commit is contained in:
Federico Kamelhar 2026-04-05 01:49:24 -04:00
parent 30ecbf1dad
commit 91be3a866d
7 changed files with 184 additions and 126 deletions

View file

@ -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

View file

@ -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,
)

View file

@ -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

View file

@ -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"""

View file

@ -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"]

View file

@ -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()

View file

@ -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):