refactor(oci): principal-level code quality pass

- Remove _extract_text_content duplication — single definition in cohere.py,
  imported where needed; instance method on OCIChatConfig eliminated
- Move cryptography imports to module level with _CRYPTOGRAPHY_AVAILABLE flag
  and _require_cryptography() guard; no more re-import on every signing call
- Move litellm version import to module level via litellm._version; remove
  inline import inside validate_oci_environment
- sign_with_manual_credentials now returns Tuple[dict, bytes] matching
  sign_with_oci_signer — asymmetry eliminated, Optional[bytes] guards removed
  throughout stream wrappers (signed_json_body: bytes = b"")
- Rename _openai_to_oci_cohere_param_map → openai_to_oci_cohere_param_map
  for consistency with openai_to_oci_generic_param_map
- Remove double-key bug in map_openai_params where responseFormat was stored
  under both OCI and OpenAI key names simultaneously
- Remove delegating shims (adapt_messages_to_cohere_standard,
  adapt_tool_definitions_to_cohere_standard, _handle_generic_stream_chunk)
  from OCIChatConfig/OCIStreamWrapper; tests now import directly from
  cohere.py and generic.py where symbols live
- Trim __all__ to 7 genuine public symbols; remove the 13-symbol list that
  existed only to support test imports
- Collapse per-model integration test classes into pytest.mark.parametrize;
  CHAT_MODELS list is the single source of truth for model-specific config
- Black + Ruff clean across all OCI files
This commit is contained in:
Federico Kamelhar 2026-04-05 13:16:39 -04:00
parent eb184a5691
commit 17d52c208c
11 changed files with 433 additions and 721 deletions

View file

@ -11,8 +11,6 @@ import json
import uuid
from typing import Any, Dict, List, Optional
import httpx
from litellm.llms.oci.common_utils import (
OCI_JSON_TO_PYTHON_TYPES,
OCIError,

View file

@ -7,9 +7,8 @@ parsing, and streaming chunk parsing for models served with
"""
import datetime
import json
import uuid
from typing import Any, Dict, List, Optional, Union
from typing import Dict, List, Optional, Union
import httpx
@ -70,7 +69,9 @@ def adapt_messages_to_generic_oci_standard_content_message(
for content_item in content:
if not isinstance(content_item, dict):
raise OCIError(status_code=400, message="Each content item must be a dictionary")
raise OCIError(
status_code=400, message="Each content item must be a dictionary"
)
item_type = content_item.get("type")
if not isinstance(item_type, str):
@ -119,9 +120,13 @@ def adapt_messages_to_generic_oci_standard_tool_call(
tool_calls_formatted = []
for tool_call in tool_calls:
if not isinstance(tool_call, dict):
raise OCIError(status_code=400, message="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 OCIError(status_code=400, message="OCI only supports function tool calls")
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):
@ -129,7 +134,9 @@ def adapt_messages_to_generic_oci_standard_tool_call(
tool_function = tool_call.get("function")
if not isinstance(tool_function, dict):
raise OCIError(status_code=400, message="Tool call `function` must be 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):
@ -186,7 +193,9 @@ def adapt_messages_to_generic_oci_standard(
if role == "assistant" and tool_calls is not None:
if not isinstance(tool_calls, list):
raise OCIError(status_code=400, message="Message `tool_calls` must be a list")
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)
)
@ -240,7 +249,9 @@ def adapt_tool_definition_to_oci_standard(
tool_function = tool.get("function")
if not isinstance(tool_function, dict):
raise OCIError(status_code=400, message="Tool `function` must be a dictionary")
raise OCIError(
status_code=400, message="Tool `function` must be a dictionary"
)
raw_params = tool_function.get("parameters", {})
resolved_params = sanitize_oci_schema(
@ -308,7 +319,9 @@ def handle_generic_response(
):
message.content = response_message.content[0].text
if response_message.toolCalls:
message.tool_calls = adapt_tools_to_openai_standard(response_message.toolCalls)
message.tool_calls = adapt_tools_to_openai_standard(
response_message.toolCalls
)
oci_usage = completion_response.chatResponse.usage
model_response.usage = Usage( # type: ignore[attr-defined]
@ -377,7 +390,9 @@ def handle_generic_stream_chunk(dict_chunk: dict) -> ModelResponseStream:
delta=Delta(
content=text,
tool_calls=(
[tool.model_dump() for tool in tool_calls] if tool_calls else None
[tool.model_dump() for tool in tool_calls]
if tool_calls
else None
),
provider_specific_fields=None,
thinking_blocks=None,

View file

@ -35,6 +35,7 @@ from litellm.llms.custom_httpx.http_handler import (
version,
)
from litellm.llms.oci.chat.cohere import (
_extract_text_content,
adapt_messages_to_cohere_standard,
adapt_tool_definitions_to_cohere_standard,
handle_cohere_response,
@ -43,10 +44,8 @@ from litellm.llms.oci.chat.cohere import (
from litellm.llms.oci.chat.generic import (
adapt_messages_to_generic_oci_standard,
adapt_tool_definition_to_oci_standard,
adapt_tools_to_openai_standard,
handle_generic_response,
handle_generic_stream_chunk,
open_ai_to_generic_oci_role_map,
)
from litellm.llms.oci.common_utils import (
OCIError,
@ -143,7 +142,7 @@ class OCIChatConfig(BaseConfig):
# - tool_choice is unsupported
# - stop sequences key is "stopSequences" not "stop"
# - n (numGenerations) is GENERIC-only
self._openai_to_oci_cohere_param_map = {
self.openai_to_oci_cohere_param_map = {
k: ("stopSequences" if k == "stop" else v)
for k, v in self.openai_to_oci_generic_param_map.items()
if k not in ("tool_choice", "max_retries", "n")
@ -151,7 +150,7 @@ class OCIChatConfig(BaseConfig):
def get_supported_openai_params(self, model: str) -> List[str]:
param_map = (
self._openai_to_oci_cohere_param_map
self.openai_to_oci_cohere_param_map
if get_vendor_from_model(model) == OCIVendors.COHERE
else self.openai_to_oci_generic_param_map
)
@ -167,7 +166,7 @@ class OCIChatConfig(BaseConfig):
adapted_params = {}
vendor = get_vendor_from_model(model)
param_map = (
self._openai_to_oci_cohere_param_map
self.openai_to_oci_cohere_param_map
if vendor == OCIVendors.COHERE
else self.openai_to_oci_generic_param_map
)
@ -185,8 +184,6 @@ class OCIChatConfig(BaseConfig):
adapted_params[key] = value
continue
adapted_params[alias] = value
if alias == "responseFormat":
adapted_params["response_format"] = value
return adapted_params
@ -200,7 +197,7 @@ class OCIChatConfig(BaseConfig):
model: Optional[str] = None,
stream: Optional[bool] = None,
fake_stream: Optional[bool] = None,
) -> Tuple[dict, Optional[bytes]]:
) -> Tuple[dict, bytes]:
return sign_oci_request(
headers=headers,
optional_params=optional_params,
@ -231,7 +228,12 @@ class OCIChatConfig(BaseConfig):
creds = resolve_oci_credentials(optional_params)
missing = [
k
for k in ("oci_user", "oci_fingerprint", "oci_tenancy", "oci_compartment_id")
for k in (
"oci_user",
"oci_fingerprint",
"oci_tenancy",
"oci_compartment_id",
)
if not creds.get(k)
]
if missing or not (creds.get("oci_key") or creds.get("oci_key_file")):
@ -261,7 +263,7 @@ class OCIChatConfig(BaseConfig):
def _get_optional_params(self, vendor: OCIVendors, optional_params: dict) -> Dict:
param_map = (
self._openai_to_oci_cohere_param_map
self.openai_to_oci_cohere_param_map
if vendor == OCIVendors.COHERE
else self.openai_to_oci_generic_param_map
)
@ -272,7 +274,11 @@ class OCIChatConfig(BaseConfig):
selected_params[oci_key] = optional_params[openai_key] # type: ignore[index]
for oci_value in param_map.values():
if oci_value and oci_value in optional_params and oci_value not in selected_params:
if (
oci_value
and oci_value in optional_params
and oci_value not in selected_params
):
selected_params[oci_value] = optional_params[oci_value] # type: ignore[index]
if "tools" in selected_params:
@ -309,7 +315,9 @@ class OCIChatConfig(BaseConfig):
schema_payload: Optional[Any] = None
if "json_schema" in rf_payload:
raw_schema = rf_payload.pop("json_schema")
schema_payload = dict(raw_schema) if isinstance(raw_schema, dict) else raw_schema
schema_payload = (
dict(raw_schema) if isinstance(raw_schema, dict) else raw_schema
)
if schema_payload is not None:
rf_payload["jsonSchema"] = schema_payload
if vendor == OCIVendors.COHERE:
@ -320,35 +328,6 @@ class OCIChatConfig(BaseConfig):
return selected_params
# ------------------------------------------------------------------
# Delegating wrappers — keep instance-method call sites working
# while the real logic lives in cohere.py / generic.py.
# ------------------------------------------------------------------
def adapt_messages_to_cohere_standard(
self, messages: List[AllMessageValues]
): # type: ignore[return]
return adapt_messages_to_cohere_standard(messages)
def adapt_tool_definitions_to_cohere_standard(
self, tools: List[Dict[str, Any]]
): # type: ignore[return]
return adapt_tool_definitions_to_cohere_standard(tools)
def _extract_text_content(self, content: Any) -> str:
"""Return plain-text for message content (string or content-part list)."""
if content is None:
return ""
if isinstance(content, str):
return content
if isinstance(content, list):
return "".join(
item.get("text", "")
for item in content
if isinstance(item, dict) and item.get("type") == "text"
)
return str(content)
def transform_request(
self,
model: str,
@ -397,14 +376,14 @@ class OCIChatConfig(BaseConfig):
preamble_override = None
if system_messages:
preamble = "\n".join(
self._extract_text_content(m["content"]) for m in system_messages
_extract_text_content(m["content"]) for m in system_messages
)
if preamble:
preamble_override = preamble
chat_request = CohereChatRequest(
apiFormat="COHERE",
message=self._extract_text_content(user_messages[-1]["content"]),
message=_extract_text_content(user_messages[-1]["content"]),
chatHistory=adapt_messages_to_cohere_standard(messages),
preambleOverride=preamble_override,
**self._get_optional_params(OCIVendors.COHERE, optional_params),
@ -456,7 +435,9 @@ class OCIChatConfig(BaseConfig):
vendor = get_vendor_from_model(model)
if vendor == OCIVendors.COHERE:
model_response = handle_cohere_response(response_json, model, model_response)
model_response = handle_cohere_response(
response_json, model, model_response
)
else:
model_response = handle_generic_response(
response_json, model, model_response, raw_response
@ -477,7 +458,7 @@ class OCIChatConfig(BaseConfig):
messages: list,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
json_mode: Optional[bool] = None,
signed_json_body: Optional[bytes] = None,
signed_json_body: bytes = b"",
) -> "OCIStreamWrapper":
if "stream" in data:
del data["stream"]
@ -488,7 +469,7 @@ class OCIChatConfig(BaseConfig):
response = client.post(
api_base,
headers=headers,
data=signed_json_body if signed_json_body is not None else json.dumps(data),
data=signed_json_body or json.dumps(data),
stream=True,
logging_obj=logging_obj,
timeout=STREAMING_TIMEOUT,
@ -525,7 +506,7 @@ class OCIChatConfig(BaseConfig):
messages: list,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
json_mode: Optional[bool] = None,
signed_json_body: Optional[bytes] = None,
signed_json_body: bytes = b"",
) -> "OCIStreamWrapper":
if "stream" in data:
del data["stream"]
@ -536,7 +517,7 @@ class OCIChatConfig(BaseConfig):
response = await client.post(
api_base,
headers=headers,
data=signed_json_body if signed_json_body is not None else json.dumps(data),
data=signed_json_body or json.dumps(data),
stream=True,
logging_obj=logging_obj,
timeout=STREAMING_TIMEOUT,
@ -586,18 +567,7 @@ class OCIStreamWrapper(CustomStreamWrapper):
return handle_cohere_stream_chunk(dict_chunk)
return handle_generic_stream_chunk(dict_chunk)
# Delegating shims so any code that calls these as instance methods keeps working.
def _handle_cohere_stream_chunk(self, dict_chunk: dict) -> ModelResponseStream:
return handle_cohere_stream_chunk(dict_chunk)
def _handle_generic_stream_chunk(self, dict_chunk: dict) -> ModelResponseStream:
return handle_generic_stream_chunk(dict_chunk)
# ---------------------------------------------------------------------------
# Backward-compatibility re-exports
# Keep all symbols that test files import directly from this module.
# ---------------------------------------------------------------------------
__all__ = [
"OCIChatConfig",
"OCIStreamWrapper",
@ -606,15 +576,4 @@ __all__ = [
"STREAMING_TIMEOUT",
"get_vendor_from_model",
"version",
# generic helpers (imported in tests)
"open_ai_to_generic_oci_role_map",
"adapt_messages_to_generic_oci_standard",
"adapt_messages_to_generic_oci_standard_content_message",
"adapt_messages_to_generic_oci_standard_tool_call",
"adapt_messages_to_generic_oci_standard_tool_response",
"adapt_tool_definition_to_oci_standard",
"adapt_tools_to_openai_standard",
# cohere helpers
"adapt_messages_to_cohere_standard",
"adapt_tool_definitions_to_cohere_standard",
]

View file

@ -4,13 +4,34 @@ import hashlib
import json
import os
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Protocol, Tuple
from typing import Any, Dict, Optional, Protocol, Tuple
from urllib.parse import urlparse
import httpx
from litellm.llms.base_llm.chat.transformation import BaseLLMException
try:
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import padding, rsa
_CRYPTOGRAPHY_AVAILABLE = True
except ImportError:
_CRYPTOGRAPHY_AVAILABLE = False
try:
from litellm._version import version as _litellm_version
except ImportError:
_litellm_version = "0.0.0"
def _require_cryptography() -> None:
if not _CRYPTOGRAPHY_AVAILABLE:
raise ImportError(
"cryptography package is required for OCI authentication. "
"Please install it with: pip install cryptography"
)
class OCIError(BaseLLMException):
def __init__(
@ -84,20 +105,12 @@ def build_signature_string(
def load_private_key_from_str(key_str: str) -> Any:
try:
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
except ImportError as e:
raise ImportError(
"cryptography package is required for OCI authentication. "
"Please install it with: pip install cryptography"
) from e
key = serialization.load_pem_private_key(
_require_cryptography()
key = serialization.load_pem_private_key( # type: ignore[union-attr]
key_str.encode("utf-8"),
password=None,
)
if not isinstance(key, rsa.RSAPrivateKey):
if not isinstance(key, rsa.RSAPrivateKey): # type: ignore[union-attr]
raise TypeError(
"The provided private key is not an RSA key, which is required for OCI signing."
)
@ -220,7 +233,7 @@ def sign_with_manual_credentials(
optional_params: dict,
request_data: dict,
api_base: str,
) -> Tuple[dict, None]:
) -> Tuple[dict, bytes]:
"""Sign a request using manually provided OCI credentials (user/fingerprint/tenancy/key)."""
creds = resolve_oci_credentials(optional_params)
oci_user = creds["oci_user"]
@ -252,7 +265,9 @@ def sign_with_manual_credentials(
path = parsed.path or "/"
host = parsed.netloc
date = datetime.datetime.now(datetime.timezone.utc).strftime("%a, %d %b %Y %H:%M:%S GMT")
date = datetime.datetime.now(datetime.timezone.utc).strftime(
"%a, %d %b %Y %H:%M:%S GMT"
)
content_type = headers.get("content-type", "application/json")
content_length = str(len(body))
x_content_sha256 = sha256_base64(body)
@ -273,16 +288,11 @@ def sign_with_manual_credentials(
"content-type",
"x-content-sha256",
]
signing_string = build_signature_string(method, path, headers_to_sign, signed_header_names)
signing_string = build_signature_string(
method, path, headers_to_sign, signed_header_names
)
try:
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.asymmetric import padding
except ImportError as e:
raise ImportError(
"cryptography package is required for OCI authentication. "
"Please install it with: pip install cryptography"
) from e
_require_cryptography()
# Resolve the private key — prefer inline PEM content over file path
oci_key_content: Optional[str] = None
@ -300,9 +310,7 @@ def sign_with_manual_credentials(
private_key = (
load_private_key_from_str(oci_key_content)
if oci_key_content
else load_private_key_from_file(oci_key_file)
if oci_key_file
else None
else load_private_key_from_file(oci_key_file) if oci_key_file else None
)
if private_key is None:
@ -313,8 +321,8 @@ def sign_with_manual_credentials(
signature = private_key.sign(
signing_string.encode("utf-8"),
padding.PKCS1v15(),
hashes.SHA256(),
padding.PKCS1v15(), # type: ignore[union-attr]
hashes.SHA256(), # type: ignore[union-attr]
)
signature_b64 = base64.b64encode(signature).decode()
@ -337,7 +345,7 @@ def sign_with_manual_credentials(
"x-content-sha256": x_content_sha256,
}
)
return headers, None
return headers, body
def sign_oci_request(
@ -349,7 +357,7 @@ def sign_oci_request(
model: Optional[str] = None,
stream: Optional[bool] = None,
fake_stream: Optional[bool] = None,
) -> Tuple[dict, Optional[bytes]]:
) -> Tuple[dict, bytes]:
"""
Route to the appropriate OCI signing method based on what credentials are present.
@ -358,11 +366,13 @@ def sign_oci_request(
also be supplied via OCI_* environment variables).
Returns:
Tuple of (signed_headers, body_bytes_or_None)
Tuple of (signed_headers, signed_body_bytes)
"""
if optional_params.get("oci_signer") is not None:
return sign_with_oci_signer(headers, optional_params, request_data, api_base)
return sign_with_manual_credentials(headers, optional_params, request_data, api_base)
return sign_with_manual_credentials(
headers, optional_params, request_data, api_base
)
def validate_oci_environment(
@ -377,10 +387,8 @@ def validate_oci_environment(
supplied via environment variables are resolved at call time rather than
at construction time.
"""
from litellm.llms.custom_httpx.http_handler import version
headers.setdefault("content-type", "application/json")
headers.setdefault("user-agent", f"litellm/{version}")
headers.setdefault("user-agent", f"litellm/{_litellm_version}")
return headers
@ -447,7 +455,8 @@ def resolve_oci_schema_anyof(obj: Any) -> Any:
if isinstance(obj, dict):
if "anyOf" in obj and "type" not in obj:
non_null = [
t for t in obj["anyOf"]
t
for t in obj["anyOf"]
if not (isinstance(t, dict) and t.get("type") == "null")
]
if non_null:
@ -504,7 +513,9 @@ def sanitize_oci_schema(schema: Any) -> Any:
return sanitized
def enrich_cohere_param_description(description: str, param_schema: Dict[str, Any]) -> str:
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

View file

@ -206,7 +206,9 @@ class OCIEmbedConfig(BaseEmbeddingConfig):
if serving_mode_type == "DEDICATED":
endpoint_id = optional_params.get("oci_endpoint_id", model)
serving_mode = OCIServingMode(servingType="DEDICATED", endpointId=endpoint_id)
serving_mode = OCIServingMode(
servingType="DEDICATED", endpointId=endpoint_id
)
else:
serving_mode = OCIServingMode(servingType="ON_DEMAND", modelId=model)

View file

@ -421,9 +421,13 @@ class OCIEmbedRequest(BaseModel):
compartmentId: str
servingMode: OCIServingMode
inputs: List[str]
inputType: Optional[str] = None # SEARCH_DOCUMENT | SEARCH_QUERY | CLASSIFICATION | CLUSTERING | IMAGE
inputType: Optional[str] = (
None # SEARCH_DOCUMENT | SEARCH_QUERY | CLASSIFICATION | CLUSTERING | IMAGE
)
truncate: Optional[str] = "END" # NONE | START | END
outputDimensions: Optional[int] = None # cohere.embed-v4.0+; valid: 256, 512, 1024, 1536
outputDimensions: Optional[int] = (
None # cohere.embed-v4.0+; valid: 256, 512, 1024, 1536
)
class OCIEmbedUsage(BaseModel):
@ -442,5 +446,3 @@ class OCIEmbedResponse(BaseModel):
inputTextTokenCounts: Optional[List[int]] = None
# Some deployments may return a usage object instead
usage: Optional[OCIEmbedUsage] = None

View file

@ -18,9 +18,10 @@ Run only these tests:
pytest tests/llm_translation/test_oci_integration.py -v
"""
import math
import os
import sys
from typing import Generator
from typing import NamedTuple, Optional
import pytest
@ -67,244 +68,243 @@ def oci_params(oci_signer) -> dict:
# ---------------------------------------------------------------------------
# Helpers
# Model registry
#
# Each entry drives the runtime pivot inside OCI's own transformation layer —
# the tests themselves are format-agnostic. Per-model quirks are captured in
# the config fields below rather than in separate test classes.
# ---------------------------------------------------------------------------
def _chat(model: str, message: str, params: dict, max_tokens: int = 64) -> str:
"""Run a single-turn completion and return the text content."""
class _M(NamedTuple):
"""Per-model test configuration."""
model: str
max_tokens: int
# Reasoning models (Gemini 2.5, Grok mini) may return None content when the
# reasoning budget is exhausted before the answer token budget starts.
reasoning: bool = False
# tool_choice value to send; None means omit the parameter entirely.
tool_choice: Optional[str] = "auto"
# Whether to include the model in tool-use parametrize list.
supports_tool_use: bool = True
# All chat models under test.
CHAT_MODELS = [
pytest.param(_M("meta.llama-3.3-70b-instruct", 64), id="meta"),
pytest.param(_M("google.gemini-2.5-flash", 200, reasoning=True), id="google"),
pytest.param(_M("xai.grok-3-mini", 100, reasoning=True), id="xai"),
pytest.param(_M("cohere.command-latest", 64, tool_choice=None), id="cohere"),
]
# Subset of models that reliably support tool use in OCI.
# xAI Grok mini is omitted — OCI does not expose tool-use for it yet.
TOOL_USE_MODELS = [
pytest.param(_M("meta.llama-3.3-70b-instruct", 100), id="meta"),
pytest.param(_M("cohere.command-latest", 200, tool_choice=None), id="cohere"),
pytest.param(_M("google.gemini-2.5-flash", 200, reasoning=True), id="google"),
]
# Simple weather tool used by all tool-use tests.
_WEATHER_TOOL = {
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the current weather for a city.",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string", "description": "The city name."}},
"required": ["city"],
},
},
}
# ---------------------------------------------------------------------------
# Sync chat tests — model list drives the pivot, not separate test classes
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("m", CHAT_MODELS)
def test_basic_completion(m: _M, oci_params):
import litellm
resp = litellm.completion(
model=f"oci/{model}",
messages=[{"role": "user", "content": message}],
max_tokens=max_tokens,
**params,
model=f"oci/{m.model}",
messages=[{"role": "user", "content": "Reply with only the word: pong"}],
max_tokens=m.max_tokens,
**oci_params,
)
# Reasoning models may return None content when all budget is used by reasoning
return resp.choices[0].message.content or ""
assert resp.choices[0].finish_reason is not None
assert resp.usage.prompt_tokens > 0
if not m.reasoning:
assert resp.choices[0].message.content is not None
@pytest.mark.parametrize("m", CHAT_MODELS)
def test_usage_populated(m: _M, oci_params):
import litellm
resp = litellm.completion(
model=f"oci/{m.model}",
messages=[{"role": "user", "content": "What is 2+2?"}],
max_tokens=m.max_tokens,
**oci_params,
)
assert resp.usage.prompt_tokens > 0
assert resp.usage.total_tokens >= resp.usage.prompt_tokens
@pytest.mark.parametrize("m", CHAT_MODELS)
def test_system_message(m: _M, oci_params):
import litellm
resp = litellm.completion(
model=f"oci/{m.model}",
messages=[
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Say hello."},
],
max_tokens=m.max_tokens,
**oci_params,
)
assert resp.choices[0].finish_reason is not None
@pytest.mark.parametrize("m", CHAT_MODELS)
def test_streaming(m: _M, oci_params):
import litellm
chunks = list(
litellm.completion(
model=f"oci/{m.model}",
messages=[{"role": "user", "content": "Count to 3."}],
max_tokens=m.max_tokens,
stream=True,
**oci_params,
)
)
assert len(chunks) > 0
# Reasoning models may stream only reasoning tokens and return empty content.
if not m.reasoning:
content = "".join(c.choices[0].delta.content or "" for c in chunks if c.choices)
assert len(content) > 0
@pytest.mark.parametrize("m", CHAT_MODELS)
def test_multi_turn(m: _M, oci_params):
import litellm
resp = litellm.completion(
model=f"oci/{m.model}",
messages=[
{"role": "user", "content": "My name is Alice."},
{"role": "assistant", "content": "Nice to meet you, Alice!"},
{"role": "user", "content": "What is my name?"},
],
max_tokens=m.max_tokens,
**oci_params,
)
# Reasoning models may have None content; skip text assertion for them.
content = resp.choices[0].message.content or ""
if not m.reasoning:
assert "Alice" in content
# ---------------------------------------------------------------------------
# Chat tests — one per vendor family
# Async chat tests
# ---------------------------------------------------------------------------
class TestOCIChatMeta:
"""Meta Llama models (GENERIC apiFormat)."""
@pytest.mark.asyncio
@pytest.mark.parametrize("m", CHAT_MODELS)
async def test_async_completion(m: _M, oci_params):
import litellm
MODEL = "meta.llama-3.3-70b-instruct"
def test_basic_completion(self, oci_params):
import litellm
resp = litellm.completion(
model=f"oci/{self.MODEL}",
messages=[{"role": "user", "content": "Reply with only the word: pong"}],
max_tokens=10,
**oci_params,
)
assert resp.choices[0].message.content is not None
assert resp.choices[0].finish_reason is not None
assert resp.usage.prompt_tokens > 0
def test_usage_populated(self, oci_params):
import litellm
resp = litellm.completion(
model=f"oci/{self.MODEL}",
messages=[{"role": "user", "content": "Say hi."}],
max_tokens=20,
**oci_params,
)
assert resp.usage.prompt_tokens > 0
assert resp.usage.total_tokens >= resp.usage.prompt_tokens
def test_system_message(self, oci_params):
import litellm
resp = litellm.completion(
model=f"oci/{self.MODEL}",
messages=[
{"role": "system", "content": "You only reply in pirate speak."},
{"role": "user", "content": "Hello!"},
],
max_tokens=30,
**oci_params,
)
resp = await litellm.acompletion(
model=f"oci/{m.model}",
messages=[{"role": "user", "content": "Reply with only the word: pong"}],
max_tokens=m.max_tokens,
**oci_params,
)
assert resp.choices[0].finish_reason is not None
assert resp.usage.total_tokens > 0
if not m.reasoning:
assert resp.choices[0].message.content is not None
def test_streaming(self, oci_params):
import litellm
chunks = list(
litellm.completion(
model=f"oci/{self.MODEL}",
messages=[{"role": "user", "content": "Count to 3."}],
max_tokens=30,
stream=True,
**oci_params,
)
)
assert len(chunks) > 0
content = "".join(
c.choices[0].delta.content or "" for c in chunks if c.choices
)
@pytest.mark.asyncio
@pytest.mark.parametrize("m", CHAT_MODELS)
async def test_async_streaming(m: _M, oci_params):
import litellm
chunks = []
async for chunk in await litellm.acompletion(
model=f"oci/{m.model}",
messages=[{"role": "user", "content": "Count to 3."}],
max_tokens=m.max_tokens,
stream=True,
**oci_params,
):
chunks.append(chunk)
assert len(chunks) > 0
if not m.reasoning:
content = "".join(c.choices[0].delta.content or "" for c in chunks if c.choices)
assert len(content) > 0
def test_multi_turn(self, oci_params):
import litellm
resp = litellm.completion(
model=f"oci/{self.MODEL}",
messages=[
{"role": "user", "content": "My name is Alice."},
{"role": "assistant", "content": "Nice to meet you, Alice!"},
{"role": "user", "content": "What is my name?"},
],
max_tokens=30,
**oci_params,
)
assert "Alice" in (resp.choices[0].message.content or "")
# ---------------------------------------------------------------------------
# Tool-use tests
# ---------------------------------------------------------------------------
class TestOCIChatGoogle:
"""Google Gemini models (GENERIC apiFormat)."""
MODEL = "google.gemini-2.5-flash"
def test_basic_completion(self, oci_params):
import litellm
resp = litellm.completion(
model=f"oci/{self.MODEL}",
messages=[{"role": "user", "content": "Reply with only the word: pong"}],
max_tokens=200,
**oci_params,
)
# Gemini 2.5 Flash is a reasoning model — content may be empty if reasoning
# consumed the budget, but no exception should be raised.
assert resp.choices[0].finish_reason is not None
assert resp.usage.prompt_tokens > 0
def test_usage_populated(self, oci_params):
import litellm
resp = litellm.completion(
model=f"oci/{self.MODEL}",
messages=[{"role": "user", "content": "What is 2+2?"}],
max_tokens=200,
**oci_params,
)
assert resp.usage.total_tokens > 0
def test_streaming(self, oci_params):
import litellm
chunks = list(
litellm.completion(
model=f"oci/{self.MODEL}",
messages=[{"role": "user", "content": "Say the word hello."}],
max_tokens=200,
stream=True,
**oci_params,
)
)
assert len(chunks) > 0
def _assert_tool_call(resp, expected_tool: str = "get_weather"):
"""Assert the response contains the expected tool call (or a plain stop)."""
choice = resp.choices[0]
assert choice.finish_reason in ("tool_calls", "stop")
if choice.finish_reason == "tool_calls":
assert choice.message.tool_calls is not None
assert len(choice.message.tool_calls) > 0
assert choice.message.tool_calls[0].function.name == expected_tool
class TestOCIChatXAI:
"""xAI Grok models (GENERIC apiFormat)."""
@pytest.mark.parametrize("m", TOOL_USE_MODELS)
def test_tool_use(m: _M, oci_params):
import litellm
MODEL = "xai.grok-3-mini"
call_kwargs = dict(
model=f"oci/{m.model}",
messages=[{"role": "user", "content": "What is the weather in Paris?"}],
tools=[_WEATHER_TOOL],
max_tokens=m.max_tokens,
**oci_params,
)
if m.tool_choice is not None:
call_kwargs["tool_choice"] = m.tool_choice
def test_basic_completion(self, oci_params):
import litellm
resp = litellm.completion(
model=f"oci/{self.MODEL}",
messages=[{"role": "user", "content": "Reply with only the word: pong"}],
max_tokens=50,
**oci_params,
)
assert resp.choices[0].message.content is not None
assert resp.usage.total_tokens > 0
def test_streaming(self, oci_params):
import litellm
chunks = list(
litellm.completion(
model=f"oci/{self.MODEL}",
messages=[{"role": "user", "content": "Count to 3."}],
max_tokens=50,
stream=True,
**oci_params,
)
)
assert len(chunks) > 0
def test_usage_has_reasoning_tokens(self, oci_params):
"""Grok mini exposes reasoning token breakdown in usage."""
import litellm
resp = litellm.completion(
model=f"oci/{self.MODEL}",
messages=[{"role": "user", "content": "What is 5*7?"}],
max_tokens=100,
**oci_params,
)
# totalTokens >= completionTokens + promptTokens (reasoning may add overhead)
assert resp.usage.total_tokens >= resp.usage.prompt_tokens
resp = litellm.completion(**call_kwargs)
_assert_tool_call(resp)
class TestOCIChatCohere:
"""Cohere Command models (COHERE apiFormat)."""
@pytest.mark.asyncio
@pytest.mark.parametrize("m", TOOL_USE_MODELS)
async def test_async_tool_use(m: _M, oci_params):
import litellm
MODEL = "cohere.command-latest"
call_kwargs = dict(
model=f"oci/{m.model}",
messages=[{"role": "user", "content": "What is the weather in Berlin?"}],
tools=[_WEATHER_TOOL],
max_tokens=m.max_tokens,
**oci_params,
)
if m.tool_choice is not None:
call_kwargs["tool_choice"] = m.tool_choice
def test_basic_completion(self, oci_params):
import litellm
resp = litellm.completion(
model=f"oci/{self.MODEL}",
messages=[{"role": "user", "content": "Reply with only the word: pong"}],
max_tokens=20,
**oci_params,
)
assert resp.choices[0].message.content is not None
assert resp.usage.prompt_tokens > 0
def test_streaming(self, oci_params):
import litellm
chunks = list(
litellm.completion(
model=f"oci/{self.MODEL}",
messages=[{"role": "user", "content": "Count to 3."}],
max_tokens=30,
stream=True,
**oci_params,
)
)
assert len(chunks) > 0
content = "".join(
c.choices[0].delta.content or "" for c in chunks if c.choices
)
assert len(content) > 0
def test_system_message(self, oci_params):
import litellm
resp = litellm.completion(
model=f"oci/{self.MODEL}",
messages=[
{"role": "system", "content": "Always end your response with 'cheers'."},
{"role": "user", "content": "Say hello."},
],
max_tokens=40,
**oci_params,
)
assert resp.choices[0].message.content is not None
resp = await litellm.acompletion(**call_kwargs)
_assert_tool_call(resp)
# ---------------------------------------------------------------------------
@ -372,7 +372,6 @@ class TestOCIEmbeddings:
def test_semantic_similarity(self, oci_params):
"""Semantically similar texts should have higher cosine similarity."""
import litellm
import math
resp = litellm.embedding(
model="oci/cohere.embed-english-v3.0",
@ -387,19 +386,18 @@ class TestOCIEmbeddings:
def cosine(a, b):
dot = sum(x * y for x, y in zip(a, b))
mag_a = math.sqrt(sum(x ** 2 for x in a))
mag_b = math.sqrt(sum(x ** 2 for x in b))
mag_a = math.sqrt(sum(x**2 for x in a))
mag_b = math.sqrt(sum(x**2 for x in b))
return dot / (mag_a * mag_b)
cat1 = resp.data[0]["embedding"]
cat2 = resp.data[1]["embedding"]
stock = resp.data[2]["embedding"]
sim_cats = cosine(cat1, cat2)
sim_diff = cosine(cat1, stock)
assert sim_cats > sim_diff, (
f"Expected similar sentences to score higher ({sim_cats:.3f} vs {sim_diff:.3f})"
)
assert (
sim_cats > sim_diff
), f"Expected similar sentences to score higher ({sim_cats:.3f} vs {sim_diff:.3f})"
def test_embed_v4(self, oci_params):
import litellm
@ -426,6 +424,54 @@ class TestOCIEmbeddings:
assert resp.usage.total_tokens == resp.usage.prompt_tokens
# ---------------------------------------------------------------------------
# Async embedding tests
# ---------------------------------------------------------------------------
class TestOCIAsyncEmbeddings:
@pytest.mark.asyncio
async def test_async_embedding_basic(self, oci_params):
import litellm
resp = await litellm.aembedding(
model="oci/cohere.embed-english-v3.0",
input=["Hello world"],
input_type="SEARCH_DOCUMENT",
**oci_params,
)
assert len(resp.data) == 1
assert len(resp.data[0]["embedding"]) == 1024
assert resp.usage.prompt_tokens > 0
@pytest.mark.asyncio
async def test_async_embedding_batch(self, oci_params):
import litellm
texts = ["The quick brown fox", "jumps over the lazy dog"]
resp = await litellm.aembedding(
model="oci/cohere.embed-english-v3.0",
input=texts,
input_type="SEARCH_DOCUMENT",
**oci_params,
)
assert len(resp.data) == 2
assert all(len(item["embedding"]) == 1024 for item in resp.data)
@pytest.mark.asyncio
async def test_async_embedding_multilingual(self, oci_params):
import litellm
resp = await litellm.aembedding(
model="oci/cohere.embed-multilingual-v3.0",
input=["Bonjour le monde"],
input_type="SEARCH_DOCUMENT",
**oci_params,
)
assert len(resp.data[0]["embedding"]) == 1024
# ---------------------------------------------------------------------------
# Env-var credential path
# ---------------------------------------------------------------------------
@ -482,261 +528,3 @@ class TestOCIEnvVarCredentials:
input_type="SEARCH_DOCUMENT",
)
assert len(resp.data[0]["embedding"]) == 1024
# ---------------------------------------------------------------------------
# Async chat tests
# ---------------------------------------------------------------------------
class TestOCIAsyncChat:
"""Verify async completion and streaming work for each vendor family."""
@pytest.mark.asyncio
async def test_async_completion_meta(self, oci_params):
import litellm
resp = await litellm.acompletion(
model="oci/meta.llama-3.3-70b-instruct",
messages=[{"role": "user", "content": "Reply with only the word: pong"}],
max_tokens=10,
**oci_params,
)
assert resp.choices[0].message.content is not None
assert resp.usage.prompt_tokens > 0
@pytest.mark.asyncio
async def test_async_completion_google(self, oci_params):
import litellm
resp = await litellm.acompletion(
model="oci/google.gemini-2.5-flash",
messages=[{"role": "user", "content": "Reply with only the word: pong"}],
max_tokens=200,
**oci_params,
)
assert resp.choices[0].finish_reason is not None
assert resp.usage.total_tokens > 0
@pytest.mark.asyncio
async def test_async_completion_xai(self, oci_params):
import litellm
resp = await litellm.acompletion(
model="oci/xai.grok-3-mini",
messages=[{"role": "user", "content": "Reply with only the word: pong"}],
max_tokens=50,
**oci_params,
)
assert resp.choices[0].message.content is not None
@pytest.mark.asyncio
async def test_async_completion_cohere(self, oci_params):
import litellm
resp = await litellm.acompletion(
model="oci/cohere.command-latest",
messages=[{"role": "user", "content": "Reply with only the word: pong"}],
max_tokens=20,
**oci_params,
)
assert resp.choices[0].message.content is not None
@pytest.mark.asyncio
async def test_async_streaming_meta(self, oci_params):
import litellm
chunks = []
async for chunk in await litellm.acompletion(
model="oci/meta.llama-3.3-70b-instruct",
messages=[{"role": "user", "content": "Count to 3."}],
max_tokens=30,
stream=True,
**oci_params,
):
chunks.append(chunk)
assert len(chunks) > 0
content = "".join(
c.choices[0].delta.content or "" for c in chunks if c.choices
)
assert len(content) > 0
@pytest.mark.asyncio
async def test_async_streaming_google(self, oci_params):
import litellm
chunks = []
async for chunk in await litellm.acompletion(
model="oci/google.gemini-2.5-flash",
messages=[{"role": "user", "content": "Say the word hello."}],
max_tokens=200,
stream=True,
**oci_params,
):
chunks.append(chunk)
assert len(chunks) > 0
@pytest.mark.asyncio
async def test_async_streaming_xai(self, oci_params):
import litellm
chunks = []
async for chunk in await litellm.acompletion(
model="oci/xai.grok-3-mini",
messages=[{"role": "user", "content": "Count to 3."}],
max_tokens=50,
stream=True,
**oci_params,
):
chunks.append(chunk)
assert len(chunks) > 0
# ---------------------------------------------------------------------------
# Async embedding tests
# ---------------------------------------------------------------------------
class TestOCIAsyncEmbeddings:
@pytest.mark.asyncio
async def test_async_embedding_basic(self, oci_params):
import litellm
resp = await litellm.aembedding(
model="oci/cohere.embed-english-v3.0",
input=["Hello world"],
input_type="SEARCH_DOCUMENT",
**oci_params,
)
assert len(resp.data) == 1
assert len(resp.data[0]["embedding"]) == 1024
assert resp.usage.prompt_tokens > 0
@pytest.mark.asyncio
async def test_async_embedding_batch(self, oci_params):
import litellm
texts = ["The quick brown fox", "jumps over the lazy dog"]
resp = await litellm.aembedding(
model="oci/cohere.embed-english-v3.0",
input=texts,
input_type="SEARCH_DOCUMENT",
**oci_params,
)
assert len(resp.data) == 2
assert all(len(item["embedding"]) == 1024 for item in resp.data)
@pytest.mark.asyncio
async def test_async_embedding_multilingual(self, oci_params):
import litellm
resp = await litellm.aembedding(
model="oci/cohere.embed-multilingual-v3.0",
input=["Bonjour le monde"],
input_type="SEARCH_DOCUMENT",
**oci_params,
)
assert len(resp.data[0]["embedding"]) == 1024
# ---------------------------------------------------------------------------
# Tool use / function calling tests
# ---------------------------------------------------------------------------
class TestOCIToolUse:
"""Verify tool use (function calling) works for models that support it."""
# Simple weather tool definition
WEATHER_TOOL = {
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the current weather for a city.",
"parameters": {
"type": "object",
"properties": {
"city": {
"type": "string",
"description": "The city name.",
}
},
"required": ["city"],
},
},
}
def _assert_tool_call(self, resp, expected_tool: str = "get_weather"):
"""Assert the response contains a tool call."""
choice = resp.choices[0]
assert choice.finish_reason in ("tool_calls", "stop")
if choice.finish_reason == "tool_calls":
assert choice.message.tool_calls is not None
assert len(choice.message.tool_calls) > 0
assert choice.message.tool_calls[0].function.name == expected_tool
# stop finish_reason can happen if model answers without calling the tool —
# acceptable behaviour, not a bug.
def test_tool_use_meta(self, oci_params):
import litellm
resp = litellm.completion(
model="oci/meta.llama-3.3-70b-instruct",
messages=[
{"role": "user", "content": "What is the weather in Paris?"}
],
tools=[self.WEATHER_TOOL],
tool_choice="auto",
max_tokens=100,
**oci_params,
)
self._assert_tool_call(resp)
def test_tool_use_cohere(self, oci_params):
import litellm
resp = litellm.completion(
model="oci/cohere.command-latest",
messages=[
{"role": "user", "content": "What is the weather in Tokyo?"}
],
tools=[self.WEATHER_TOOL],
max_tokens=200,
**oci_params,
)
self._assert_tool_call(resp)
def test_tool_use_google(self, oci_params):
import litellm
resp = litellm.completion(
model="oci/google.gemini-2.5-flash",
messages=[
{"role": "user", "content": "What is the weather in London?"}
],
tools=[self.WEATHER_TOOL],
tool_choice="auto",
max_tokens=200,
**oci_params,
)
self._assert_tool_call(resp)
@pytest.mark.asyncio
async def test_async_tool_use_meta(self, oci_params):
import litellm
resp = await litellm.acompletion(
model="oci/meta.llama-3.3-70b-instruct",
messages=[
{"role": "user", "content": "What is the weather in Berlin?"}
],
tools=[self.WEATHER_TOOL],
tool_choice="auto",
max_tokens=100,
**oci_params,
)
self._assert_tool_call(resp)

View file

@ -789,9 +789,11 @@ class TestOCISplitChunks:
async def _run_async_split(self, raw_chunks):
"""Invoke the async split_chunks logic directly."""
results = []
async def _gen():
for c in raw_chunks:
yield c
async for item in _gen():
for chunk in item.split("\n\n"):
stripped = chunk.strip()
@ -877,9 +879,8 @@ class TestOCIProviderEmbeddingConfig:
"""
import inspect
from litellm.utils import ProviderConfigManager
source = inspect.getsource(
ProviderConfigManager.get_provider_embedding_config
)
source = inspect.getsource(ProviderConfigManager.get_provider_embedding_config)
oci_count = source.count("LlmProviders.OCI")
assert oci_count == 1, (
f"Expected exactly 1 OCI branch in get_provider_embedding_config, found {oci_count}. "
@ -930,10 +931,16 @@ class TestOCICohereParamMapping:
model="cohere.command-latest",
drop_params=False,
)
for injected in ("maxTokens", "temperature", "topK", "topP", "frequencyPenalty"):
assert injected not in result, (
f"'{injected}' should not be injected when user did not provide it"
)
for injected in (
"maxTokens",
"temperature",
"topK",
"topP",
"frequencyPenalty",
):
assert (
injected not in result
), f"'{injected}' should not be injected when user did not provide it"
def test_cohere_explicit_params_still_passed(self):
"""User-provided Cohere params must still be forwarded correctly."""
@ -993,9 +1000,9 @@ class TestOCIStreamingSignedBody:
signed_json_body=signed_bytes,
)
assert posted_data["data"] == signed_bytes, (
"Streaming must use signed_json_body, not re-serialize data"
)
assert (
posted_data["data"] == signed_bytes
), "Streaming must use signed_json_body, not re-serialize data"
def test_get_custom_stream_wrapper_fallback_without_signed_body(self, monkeypatch):
"""When signed_json_body is None, fall back to json.dumps(data)."""
@ -1032,6 +1039,6 @@ class TestOCIStreamingSignedBody:
signed_json_body=None,
)
assert posted_data["data"] == json.dumps(payload), (
"Without signed_json_body, must fall back to json.dumps(data)"
)
assert posted_data["data"] == json.dumps(
payload
), "Without signed_json_body, must fall back to json.dumps(data)"

View file

@ -5,10 +5,14 @@ import json
from unittest.mock import patch, MagicMock
from litellm import ModelResponse
from litellm.llms.oci.chat.cohere import (
adapt_messages_to_cohere_standard,
adapt_tool_definitions_to_cohere_standard,
)
from litellm.llms.oci.chat.transformation import (
OCIChatConfig,
get_vendor_from_model,
OCIStreamWrapper,
get_vendor_from_model,
)
from litellm.types.llms.oci import OCIVendors
@ -75,7 +79,7 @@ class TestOCICohereToolCalls:
]
# Transform tools
cohere_tools = config.adapt_tool_definitions_to_cohere_standard(openai_tools)
cohere_tools = adapt_tool_definitions_to_cohere_standard(openai_tools)
# Verify transformation
assert len(cohere_tools) == 2
@ -95,7 +99,10 @@ class TestOCICohereToolCalls:
# Check unit parameter
unit_param = weather_tool.parameterDefinitions["unit"]
assert unit_param.description == "Temperature unit (celsius or fahrenheit). Allowed values: ['celsius', 'fahrenheit']"
assert (
unit_param.description
== "Temperature unit (celsius or fahrenheit). Allowed values: ['celsius', 'fahrenheit']"
)
assert unit_param.type == "str"
assert unit_param.isRequired == False
@ -156,7 +163,7 @@ class TestOCICohereToolCalls:
assert chat_request["apiFormat"] == "COHERE"
assert chat_request["message"] == "What's the weather like in Tokyo?"
assert chat_request["chatHistory"] == []
# Verify tools are transformed correctly
assert "tools" in chat_request
assert len(chat_request["tools"]) == 1
@ -317,7 +324,7 @@ class TestOCICohereToolCalls:
},
]
chat_history = config.adapt_messages_to_cohere_standard(messages)
chat_history = adapt_messages_to_cohere_standard(messages)
# First message is the user message
assert chat_history[0].role == "USER"
@ -358,7 +365,7 @@ class TestOCICohereToolCalls:
},
]
chat_history = config.adapt_messages_to_cohere_standard(messages)
chat_history = adapt_messages_to_cohere_standard(messages)
# Verify chat history structure (excludes last message)
assert len(chat_history) == 2
@ -524,7 +531,7 @@ class TestOCICohereToolCalls:
]
# The function should handle missing function key gracefully
cohere_tools = config.adapt_tool_definitions_to_cohere_standard(invalid_tools)
cohere_tools = adapt_tool_definitions_to_cohere_standard(invalid_tools)
# Should create a tool with empty name and description
assert len(cohere_tools) == 1
@ -679,7 +686,7 @@ class TestOCICoherePreambleOverride:
{"role": "user", "content": "Second question"},
]
chat_history = config.adapt_messages_to_cohere_standard(messages)
chat_history = adapt_messages_to_cohere_standard(messages)
# Should contain user and assistant only, no system
# Note: adapt_messages_to_cohere_standard excludes the last message
@ -706,7 +713,7 @@ class TestOCICohereStreaming:
stream_wrapper = self._create_stream_wrapper()
# chunk_creator is the public dispatch entry point
assert hasattr(stream_wrapper, 'chunk_creator')
assert hasattr(stream_wrapper, "chunk_creator")
assert callable(stream_wrapper.chunk_creator)
def test_cohere_streaming_chunk_parsing(self):

View file

@ -11,13 +11,10 @@ Error: ValidationError: 1 validation error for OCIStreamChunk message.toolCalls.
import os
import sys
import pytest
from unittest.mock import MagicMock
# Adds the parent directory to the system path
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.oci.chat.transformation import OCIStreamWrapper
from litellm.llms.oci.chat.generic import handle_generic_stream_chunk
from litellm.types.utils import ModelResponseStream
@ -26,12 +23,9 @@ class TestOCIStreamingToolCalls:
def test_stream_chunk_with_missing_arguments_field(self):
"""
Test that streaming chunks with tool calls missing 'arguments' field are handled.
OCI API can return tool calls in early chunks without the 'arguments' field,
which should be filled with an empty string to satisfy Pydantic validation.
"""
# Mock streaming chunk with tool call missing 'arguments' field
chunk_data = {
"index": 0,
"finishReason": None,
@ -43,21 +37,13 @@ class TestOCIStreamingToolCalls:
"type": "FUNCTION",
"id": "call_abc123",
"name": "get_weather",
# Note: 'arguments' field is missing
# 'arguments' field is missing
}
],
},
}
wrapper = OCIStreamWrapper(
completion_stream=iter([]),
model="meta.llama-3.1-405b-instruct",
custom_llm_provider="oci",
logging_obj=MagicMock(),
)
# This should not raise a ValidationError
result = wrapper._handle_generic_stream_chunk(chunk_data)
result = handle_generic_stream_chunk(chunk_data)
assert isinstance(result, ModelResponseStream)
assert len(result.choices) == 1
@ -66,9 +52,7 @@ class TestOCIStreamingToolCalls:
assert result.choices[0].delta.tool_calls[0]["function"]["arguments"] == ""
def test_stream_chunk_with_missing_id_field(self):
"""
Test that streaming chunks with tool calls missing 'id' field are handled.
"""
"""Missing 'id' gets a generated call_* id."""
chunk_data = {
"index": 0,
"finishReason": None,
@ -80,30 +64,20 @@ class TestOCIStreamingToolCalls:
"type": "FUNCTION",
"name": "get_weather",
"arguments": '{"location": "San Francisco"}',
# Note: 'id' field is missing
# 'id' field is missing
}
],
},
}
wrapper = OCIStreamWrapper(
completion_stream=iter([]),
model="meta.llama-3.1-405b-instruct",
custom_llm_provider="oci",
logging_obj=MagicMock(),
)
result = wrapper._handle_generic_stream_chunk(chunk_data)
result = handle_generic_stream_chunk(chunk_data)
assert isinstance(result, ModelResponseStream)
assert result.choices[0].delta.tool_calls is not None
# 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):
"""
Test that streaming chunks with tool calls missing 'name' field are handled.
"""
"""Missing 'name' defaults to empty string."""
chunk_data = {
"index": 0,
"finishReason": None,
@ -115,29 +89,20 @@ class TestOCIStreamingToolCalls:
"type": "FUNCTION",
"id": "call_abc123",
"arguments": '{"location": "San Francisco"}',
# Note: 'name' field is missing
# 'name' field is missing
}
],
},
}
wrapper = OCIStreamWrapper(
completion_stream=iter([]),
model="meta.llama-3.1-405b-instruct",
custom_llm_provider="oci",
logging_obj=MagicMock(),
)
result = wrapper._handle_generic_stream_chunk(chunk_data)
result = handle_generic_stream_chunk(chunk_data)
assert isinstance(result, ModelResponseStream)
assert result.choices[0].delta.tool_calls is not None
assert result.choices[0].delta.tool_calls[0]["function"]["name"] == ""
def test_stream_chunk_with_all_missing_fields(self):
"""
Test that streaming chunks with tool calls missing all optional fields are handled.
"""
"""All optional fields missing — all default gracefully."""
chunk_data = {
"index": 0,
"finishReason": None,
@ -147,32 +112,22 @@ class TestOCIStreamingToolCalls:
"toolCalls": [
{
"type": "FUNCTION"
# All fields missing: id, name, arguments
# id, name, arguments all missing
}
],
},
}
wrapper = OCIStreamWrapper(
completion_stream=iter([]),
model="meta.llama-3.1-405b-instruct",
custom_llm_provider="oci",
logging_obj=MagicMock(),
)
result = wrapper._handle_generic_stream_chunk(chunk_data)
result = handle_generic_stream_chunk(chunk_data)
assert isinstance(result, ModelResponseStream)
assert result.choices[0].delta.tool_calls is not None
# 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"] == ""
def test_stream_chunk_with_complete_tool_call(self):
"""
Test that streaming chunks with complete tool calls still work correctly.
"""
"""Fully-populated tool call passes through unchanged."""
chunk_data = {
"index": 0,
"finishReason": None,
@ -190,14 +145,7 @@ class TestOCIStreamingToolCalls:
},
}
wrapper = OCIStreamWrapper(
completion_stream=iter([]),
model="meta.llama-3.1-405b-instruct",
custom_llm_provider="oci",
logging_obj=MagicMock(),
)
result = wrapper._handle_generic_stream_chunk(chunk_data)
result = handle_generic_stream_chunk(chunk_data)
assert isinstance(result, ModelResponseStream)
assert result.choices[0].delta.tool_calls is not None
@ -212,9 +160,7 @@ class TestOCIStreamingToolCalls:
)
def test_stream_chunk_with_multiple_tool_calls_missing_fields(self):
"""
Test that streaming chunks with multiple tool calls, some with missing fields, are handled.
"""
"""Multiple tool calls with a mix of complete and incomplete entries."""
chunk_data = {
"index": 0,
"finishReason": None,
@ -222,50 +168,34 @@ class TestOCIStreamingToolCalls:
"role": "ASSISTANT",
"content": None,
"toolCalls": [
{
"type": "FUNCTION",
"id": "call_1",
"name": "get_weather",
# Missing arguments
},
{"type": "FUNCTION", "id": "call_1", "name": "get_weather"},
{
"type": "FUNCTION",
"name": "get_time",
"arguments": '{"timezone": "UTC"}',
# Missing id
},
{
"type": "FUNCTION",
"id": "call_3",
"name": "calculate",
"arguments": '{"expression": "2+2"}',
# Complete
},
],
},
}
wrapper = OCIStreamWrapper(
completion_stream=iter([]),
model="meta.llama-3.1-405b-instruct",
custom_llm_provider="oci",
logging_obj=MagicMock(),
)
result = wrapper._handle_generic_stream_chunk(chunk_data)
result = handle_generic_stream_chunk(chunk_data)
assert isinstance(result, ModelResponseStream)
assert result.choices[0].delta.tool_calls is not None
assert len(result.choices[0].delta.tool_calls) == 3
# First tool call - missing arguments
assert result.choices[0].delta.tool_calls[0]["id"] == "call_1"
assert (
result.choices[0].delta.tool_calls[0]["function"]["name"] == "get_weather"
)
assert result.choices[0].delta.tool_calls[0]["function"]["arguments"] == ""
# 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 (
@ -273,7 +203,6 @@ class TestOCIStreamingToolCalls:
== '{"timezone": "UTC"}'
)
# Third tool call - complete
assert result.choices[0].delta.tool_calls[2]["id"] == "call_3"
assert result.choices[0].delta.tool_calls[2]["function"]["name"] == "calculate"
assert (
@ -282,9 +211,7 @@ class TestOCIStreamingToolCalls:
)
def test_stream_chunk_without_tool_calls(self):
"""
Test that streaming chunks without tool calls continue to work as before.
"""
"""Plain text chunks (no tool calls) pass through correctly."""
chunk_data = {
"index": 0,
"finishReason": None,
@ -294,14 +221,7 @@ class TestOCIStreamingToolCalls:
},
}
wrapper = OCIStreamWrapper(
completion_stream=iter([]),
model="meta.llama-3.1-405b-instruct",
custom_llm_provider="oci",
logging_obj=MagicMock(),
)
result = wrapper._handle_generic_stream_chunk(chunk_data)
result = handle_generic_stream_chunk(chunk_data)
assert isinstance(result, ModelResponseStream)
assert result.choices[0].delta.content == "Hello, how can I help you?"

View file

@ -67,7 +67,10 @@ class TestOCIEmbedConfig:
optional_params={"oci_region": "us-chicago-1"},
litellm_params={},
)
assert url == "https://inference.generativeai.us-chicago-1.oci.oraclecloud.com/20231130/actions/embedText"
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)."""