mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
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:
parent
30ecbf1dad
commit
91be3a866d
7 changed files with 184 additions and 126 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue