feat(oci): add embeddings, fix streaming/reasoning, expand model catalog

- Add OCIEmbedConfig with full Cohere embed support (7 models, batch up to 96)
- Fix sync streaming: split SSE events on \n\n before JSON parsing
- Fix reasoning models (Gemini 2.5, xAI Grok): make completionTokens and message
  optional in OCIResponseChoice to handle max_tokens exhausted on reasoning
- Fix compartment_id resolution in chat transform to use resolve_oci_credentials
- Fix tool call id: make OCIToolCall.id optional, generate UUID fallback for
  providers (Google via OCI) that omit it
- Add OCI_KEY env var support for inline PEM keys
- Fix datetime.utcnow() deprecation in request signing
- Expand model catalog: 29 OCI models including Llama 4, Gemini 2.5, xAI Grok,
  Cohere Command A, and all Cohere embed variants
- Add 37 live integration tests: sync/async completions for Meta/Google/xAI/Cohere,
  sync/async embeddings, tool use across all vendors, streaming, env var auth
- Add 23 embed unit tests covering all transform and validation paths
This commit is contained in:
Federico Kamelhar 2026-04-05 00:16:19 -04:00
parent d9672ebc05
commit 0647074846
11 changed files with 2064 additions and 718 deletions

View file

@ -890,6 +890,18 @@ class BaseLLMHTTPHandler:
headers=headers,
)
# Some providers (e.g. OCI) require request signing after the body is built.
# The default BaseConfig.sign_request returns (headers, None) — a no-op for
# providers that don't need signing.
headers, signed_body = provider_config.sign_request(
headers=headers,
optional_params=optional_params,
request_data=data,
api_base=api_base,
api_key=api_key,
model=model,
)
## LOGGING
logging_obj.pre_call(
input=input,
@ -916,6 +928,7 @@ class BaseLLMHTTPHandler:
client=client,
optional_params=optional_params,
litellm_params=litellm_params,
signed_body=signed_body,
)
if client is None or not isinstance(client, HTTPHandler):
@ -929,7 +942,7 @@ class BaseLLMHTTPHandler:
response = sync_httpx_client.post(
url=api_base,
headers=headers,
data=json.dumps(data),
data=signed_body if signed_body is not None else json.dumps(data),
timeout=timeout,
)
except Exception as e:
@ -964,6 +977,7 @@ class BaseLLMHTTPHandler:
api_key: Optional[str] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
signed_body: Optional[bytes] = None,
) -> EmbeddingResponse:
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
@ -974,12 +988,20 @@ class BaseLLMHTTPHandler:
async_httpx_client = client
try:
response = await async_httpx_client.post(
url=api_base,
headers=headers,
json=request_data,
timeout=timeout,
)
if signed_body is not None:
response = await async_httpx_client.post(
url=api_base,
headers=headers,
data=signed_body,
timeout=timeout,
)
else:
response = await async_httpx_client.post(
url=api_base,
headers=headers,
json=request_data,
timeout=timeout,
)
except Exception as e:
raise self._handle_error(e=e, provider_config=provider_config)

View file

@ -1,20 +1,16 @@
import base64
import datetime
import hashlib
import json
from dataclasses import dataclass
from typing import (
TYPE_CHECKING,
Any,
AsyncIterator,
Dict,
Iterator,
List,
Optional,
Protocol,
Tuple,
Union,
)
from urllib.parse import urlparse
import httpx
@ -28,11 +24,18 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
version,
)
from litellm.llms.oci.common_utils import OCIError
from litellm.llms.oci.common_utils import (
OCIError,
OCIRequestWrapper, # re-exported for backwards compatibility
get_oci_base_url,
resolve_oci_credentials,
sign_oci_request,
validate_oci_environment,
)
from litellm.types.llms.oci import (
CohereChatRequest,
CohereMessage,
CohereChatResult,
CohereMessage,
CohereParameterDefinition,
CohereStreamChunk,
CohereTool,
@ -54,6 +57,7 @@ from litellm.types.llms.oci import (
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import (
Choices,
Delta,
LlmProviders,
ModelResponse,
@ -74,103 +78,6 @@ else:
LiteLLMLoggingObj = Any
class OCISignerProtocol(Protocol):
"""
Protocol for OCI request signers (e.g., oci.signer.Signer).
This protocol defines the interface expected for OCI SDK signer objects.
Compatible with the OCI Python SDK's Signer class.
See: https://docs.oracle.com/en-us/iaas/tools/python/latest/api/signing.html
"""
def do_request_sign(
self, request: Any, *, enforce_content_headers: bool = False
) -> None:
"""
Sign an HTTP request by adding authentication headers.
Args:
request: Request object with method, url, headers, body, and path_url attributes
enforce_content_headers: Whether to enforce content-type and content-length headers
"""
...
@dataclass
class OCIRequestWrapper:
"""
Wrapper for HTTP requests compatible with OCI signer interface.
This class wraps request data in a format compatible with OCI SDK signers,
which expect objects with method, url, headers, body, and path_url attributes.
"""
method: str
url: str
headers: dict
body: bytes
@property
def path_url(self) -> str:
"""Returns the path + query string for OCI signing."""
parsed_url = urlparse(self.url)
return parsed_url.path + ("?" + parsed_url.query if parsed_url.query else "")
def sha256_base64(data: bytes) -> str:
digest = hashlib.sha256(data).digest()
return base64.b64encode(digest).decode()
def build_signature_string(method, path, headers, signed_headers):
lines = []
for header in signed_headers:
if header == "(request-target)":
value = f"{method.lower()} {path}"
else:
value = headers[header]
lines.append(f"{header}: {value}")
return "\n".join(lines)
def load_private_key_from_str(key_str: str):
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(
key_str.encode("utf-8"),
password=None,
)
if not isinstance(key, rsa.RSAPrivateKey):
raise TypeError(
"The provided private key is not an RSA key, which is required for OCI signing."
)
return key
def load_private_key_from_file(file_path: str):
"""Loads a private key from a file path"""
try:
with open(file_path, "r", encoding="utf-8") as f:
key_str = f.read().strip()
except FileNotFoundError:
raise FileNotFoundError(f"Private key file not found: {file_path}")
except OSError as e:
raise OSError(f"Failed to read private key file '{file_path}': {e}") from e
if not key_str:
raise ValueError(f"Private key file is empty: {file_path}")
return load_private_key_from_str(key_str)
def get_vendor_from_model(model: str) -> OCIVendors:
"""
Extracts the vendor from the model name.
@ -187,8 +94,7 @@ def get_vendor_from_model(model: str) -> OCIVendors:
vendor = model.split(".")[0].lower()
if vendor == "cohere":
return OCIVendors.COHERE
else:
return OCIVendors.GENERIC
return OCIVendors.GENERIC
# 5 minute timeout (models may need to load)
@ -239,25 +145,21 @@ class OCIChatConfig(BaseConfig):
"response_format": "responseFormat",
}
# Cohere and Gemini use the same parameter mapping as GENERIC
self.openai_to_oci_cohere_param_map = (
self.openai_to_oci_generic_param_map.copy()
)
# Cohere uses the same parameter keys as GENERIC except tool_choice is unsupported.
# Build a *separate* frozen reference map so callers never mutate the canonical dict.
self._openai_to_oci_cohere_param_map = {
k: v
for k, v in self.openai_to_oci_generic_param_map.items()
if k not in ("tool_choice", "max_retries")
}
def get_supported_openai_params(self, model: str) -> List[str]:
supported_params = []
vendor = get_vendor_from_model(model)
if vendor == OCIVendors.COHERE:
open_ai_to_oci_param_map = self.openai_to_oci_cohere_param_map
open_ai_to_oci_param_map.pop("tool_choice")
open_ai_to_oci_param_map.pop("max_retries")
else:
open_ai_to_oci_param_map = self.openai_to_oci_generic_param_map
for key, value in open_ai_to_oci_param_map.items():
if value:
supported_params.append(key)
return supported_params
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
)
return [key for key, value in param_map.items() if value]
def map_openai_params(
self,
@ -269,7 +171,7 @@ class OCIChatConfig(BaseConfig):
adapted_params = {}
vendor = get_vendor_from_model(model)
if vendor == OCIVendors.COHERE:
open_ai_to_oci_param_map = self.openai_to_oci_cohere_param_map
open_ai_to_oci_param_map = self._openai_to_oci_cohere_param_map
else:
open_ai_to_oci_param_map = self.openai_to_oci_generic_param_map
@ -295,211 +197,6 @@ class OCIChatConfig(BaseConfig):
return adapted_params
def _sign_with_oci_signer(
self,
headers: dict,
optional_params: dict,
request_data: dict,
api_base: str,
) -> Tuple[dict, bytes]:
"""
Sign request using OCI SDK Signer object.
Args:
headers: Request headers to be signed
optional_params: Optional parameters including oci_signer
request_data: The request body dict to be sent in HTTP request
api_base: The complete URL for the HTTP request
Returns:
Tuple of (signed_headers, encoded_body)
Raises:
OCIError: If signing fails
ValueError: If HTTP method is unsupported
"""
oci_signer = optional_params.get("oci_signer")
body = json.dumps(request_data).encode("utf-8")
method = str(optional_params.get("method", "POST")).upper()
if method not in ["POST", "GET", "PUT", "DELETE", "PATCH"]:
raise ValueError(f"Unsupported HTTP method: {method}")
prepared_headers = headers.copy()
prepared_headers.setdefault("content-type", "application/json")
prepared_headers.setdefault("content-length", str(len(body)))
request_wrapper = OCIRequestWrapper(
method=method, url=api_base, headers=prepared_headers, body=body
)
if oci_signer is None:
raise ValueError(
"oci_signer cannot be None when calling _sign_with_oci_signer"
)
try:
oci_signer.do_request_sign(request_wrapper, enforce_content_headers=True)
except Exception as e:
raise OCIError(
status_code=500,
message=(
f"Failed to sign request with provided oci_signer: {str(e)}. "
"The signer must implement the OCI SDK Signer interface with a "
"do_request_sign(request, enforce_content_headers=True) method. "
"See: https://docs.oracle.com/en-us/iaas/tools/python/latest/api/signing.html"
),
) from e
headers.update(request_wrapper.headers)
return headers, body
def _sign_with_manual_credentials(
self,
headers: dict,
optional_params: dict,
request_data: dict,
api_base: str,
) -> Tuple[dict, None]:
"""
Sign request using manual OCI credentials.
Args:
headers: Request headers to be signed
optional_params: Optional parameters including OCI credentials
request_data: The request body dict to be sent in HTTP request
api_base: The complete URL for the HTTP request
Returns:
Tuple of (signed_headers, None)
Raises:
Exception: If required credentials are missing
ImportError: If cryptography package is not installed
"""
oci_region = optional_params.get("oci_region", "us-ashburn-1")
api_base = (
api_base
or litellm.api_base
or f"https://inference.generativeai.{oci_region}.oci.oraclecloud.com"
)
oci_user = optional_params.get("oci_user")
oci_fingerprint = optional_params.get("oci_fingerprint")
oci_tenancy = optional_params.get("oci_tenancy")
oci_key = optional_params.get("oci_key")
oci_key_file = optional_params.get("oci_key_file")
if (
not oci_user
or not oci_fingerprint
or not oci_tenancy
or not (oci_key or oci_key_file)
):
raise Exception(
"Missing required parameters: oci_user, oci_fingerprint, oci_tenancy, "
"and at least one of oci_key or oci_key_file."
)
method = str(optional_params.get("method", "POST")).upper()
body = json.dumps(request_data).encode("utf-8")
parsed = urlparse(api_base)
path = parsed.path or "/"
host = parsed.netloc
date = datetime.datetime.utcnow().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)
headers_to_sign = {
"date": date,
"host": host,
"content-type": content_type,
"content-length": content_length,
"x-content-sha256": x_content_sha256,
}
signed_headers = [
"date",
"(request-target)",
"host",
"content-length",
"content-type",
"x-content-sha256",
]
signing_string = build_signature_string(
method, path, headers_to_sign, signed_headers
)
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
# Handle oci_key - it should be a string (PEM content)
oci_key_content = None
if oci_key:
if isinstance(oci_key, str):
oci_key_content = oci_key
# Fix common issues with PEM content
# Replace escaped newlines with actual newlines
oci_key_content = oci_key_content.replace("\\n", "\n")
# Ensure proper line endings
if "\r\n" in oci_key_content:
oci_key_content = oci_key_content.replace("\r\n", "\n")
else:
raise OCIError(
status_code=400,
message=f"oci_key must be a string containing the PEM private key content. "
f"Got type: {type(oci_key).__name__}",
)
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
)
if private_key is None:
raise OCIError(
status_code=400,
message="Private key is required for OCI authentication. Please provide either oci_key or oci_key_file.",
)
signature = private_key.sign(
signing_string.encode("utf-8"),
padding.PKCS1v15(),
hashes.SHA256(),
)
signature_b64 = base64.b64encode(signature).decode()
key_id = f"{oci_tenancy}/{oci_user}/{oci_fingerprint}"
authorization = (
'Signature version="1",'
f'keyId="{key_id}",'
'algorithm="rsa-sha256",'
f'headers="{" ".join(signed_headers)}",'
f'signature="{signature_b64}"'
)
headers.update(
{
"authorization": authorization,
"date": date,
"host": host,
"content-type": content_type,
"content-length": content_length,
"x-content-sha256": x_content_sha256,
}
)
return headers, None
def sign_request(
self,
headers: dict,
@ -515,56 +212,23 @@ class OCIChatConfig(BaseConfig):
Sign the OCI request by adding authentication headers.
Supports two signing modes:
1. OCI SDK Signer: Use an oci_signer object to sign the request
2. Manual Signing: Use OCI credentials to manually sign the request
Args:
headers: Request headers to be signed
optional_params: Optional parameters including auth credentials or oci_signer
request_data: The request body dict to be sent in HTTP request
api_base: The complete URL for the HTTP request
api_key: Optional API key (not used for OCI)
model: Optional model name
stream: Optional streaming flag
fake_stream: Optional fake streaming flag
Returns:
Tuple of (signed_headers, encoded_body):
- If oci_signer is provided: Returns (headers, body) where body is the encoded JSON
- If manual credentials are provided: Returns (headers, None) as body is not returned
for the manual signing path
Raises:
OCIError: If signing fails with oci_signer
Exception: If required credentials are missing
ImportError: If cryptography package is not installed (manual signing only)
Example:
>>> from oci.signer import Signer
>>> signer = Signer(
... tenancy="ocid1.tenancy.oc1..",
... user="ocid1.user.oc1..",
... fingerprint="xx:xx:xx",
... private_key_file_location="~/.oci/key.pem"
... )
>>> headers, body = config.sign_request(
... headers={},
... optional_params={"oci_signer": signer},
... request_data={"message": "Hello"},
... api_base="https://inference.generativeai.us-ashburn-1.oci.oraclecloud.com/..."
... )
1. **OCI SDK Signer** — pass ``oci_signer`` (an ``oci.signer.Signer`` instance or any
object implementing :class:`~litellm.llms.oci.common_utils.OCISignerProtocol`).
2. **Manual RSA-SHA256** — pass ``oci_user``, ``oci_fingerprint``, ``oci_tenancy``, and
``oci_key`` (PEM string) or ``oci_key_file`` (path). All of these can also be
supplied via ``OCI_USER``, ``OCI_FINGERPRINT``, ``OCI_TENANCY``, and
``OCI_KEY_FILE`` environment variables.
"""
oci_signer = optional_params.get("oci_signer")
# If a signer is provided, use it for request signing
if oci_signer is not None:
return self._sign_with_oci_signer(
headers, optional_params, request_data, api_base
)
# Standard manual credential signing
return self._sign_with_manual_credentials(
headers, optional_params, request_data, api_base
return sign_oci_request(
headers=headers,
optional_params=optional_params,
request_data=request_data,
api_base=api_base,
api_key=api_key,
model=model,
stream=stream,
fake_stream=fake_stream,
)
def validate_environment(
@ -577,80 +241,29 @@ class OCIChatConfig(BaseConfig):
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
"""
Validate the OCI environment and credentials.
Supports two authentication modes:
1. OCI SDK Signer: Pass an oci_signer object (e.g., oci.signer.Signer)
2. Manual Credentials: Pass oci_user, oci_fingerprint, oci_tenancy, and oci_key/oci_key_file
Args:
headers: Request headers to populate
model: Model name
messages: List of chat messages
optional_params: Optional parameters including authentication credentials
litellm_params: LiteLLM parameters
api_key: Optional API key (not used for OCI)
api_base: Optional API base URL
Returns:
Updated headers dict
Raises:
Exception: If required parameters are missing or invalid
"""
oci_signer = optional_params.get("oci_signer")
oci_region = optional_params.get("oci_region", "us-ashburn-1")
# Determine api_base
api_base = (
api_base
or litellm.api_base
or f"https://inference.generativeai.{oci_region}.oci.oraclecloud.com"
)
if not api_base:
raise Exception(
"Either `api_base` must be provided or `litellm.api_base` must be set. "
"Alternatively, you can set the `oci_region` optional parameter to use the default OCI region."
)
# Validate credentials only if signer is not provided
if oci_signer is None:
oci_user = optional_params.get("oci_user")
oci_fingerprint = optional_params.get("oci_fingerprint")
oci_tenancy = optional_params.get("oci_tenancy")
oci_key = optional_params.get("oci_key")
oci_key_file = optional_params.get("oci_key_file")
oci_compartment_id = optional_params.get("oci_compartment_id")
if (
not oci_user
or not oci_fingerprint
or not oci_tenancy
or not (oci_key or oci_key_file)
or not oci_compartment_id
):
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. "
"Alternatively, provide an oci_signer object from the OCI SDK."
)
# Common header setup
headers.update(
{
"content-type": "application/json",
"user-agent": f"litellm/{version}",
}
)
if not messages:
raise Exception(
"kwarg `messages` must be an array of messages that follow the openai chat standard"
)
return headers
# Validate credentials early so the caller gets a clear error immediately
# rather than a cryptic signing failure at request time.
# Credentials may come from optional_params or OCI_* env vars.
if optional_params.get("oci_signer") is None:
creds = resolve_oci_credentials(optional_params)
missing = [
k
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")):
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."
)
return validate_oci_environment(headers, optional_params, api_key)
def get_complete_url(
self,
@ -661,15 +274,13 @@ class OCIChatConfig(BaseConfig):
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
oci_region = optional_params.get("oci_region", "us-ashburn-1")
return f"https://inference.generativeai.{oci_region}.oci.oraclecloud.com/20231130/actions/chat"
base = get_oci_base_url(optional_params, api_base or litellm.api_base)
return f"{base}/20231130/actions/chat"
def _get_optional_params(self, vendor: OCIVendors, optional_params: dict) -> Dict:
selected_params = {}
selected_params: Dict = {}
if vendor == OCIVendors.COHERE:
open_ai_to_oci_param_map = self.openai_to_oci_cohere_param_map
# remove tool_choice from the map
open_ai_to_oci_param_map.pop("tool_choice")
open_ai_to_oci_param_map = self._openai_to_oci_cohere_param_map
# Add default values for Cohere API
selected_params = {
"maxTokens": 600,
@ -853,9 +464,16 @@ class OCIChatConfig(BaseConfig):
litellm_params: dict,
headers: dict,
) -> dict:
oci_compartment_id = optional_params.get("oci_compartment_id", None)
creds = resolve_oci_credentials(optional_params)
oci_compartment_id = creds["oci_compartment_id"]
if not oci_compartment_id:
raise Exception("kwarg `oci_compartment_id` is required for OCI requests")
raise OCIError(
status_code=400,
message=(
"oci_compartment_id is required for OCI chat requests. "
"Pass it as optional_params or set the OCI_COMPARTMENT_ID env var."
),
)
vendor = get_vendor_from_model(model)
@ -968,8 +586,6 @@ class OCIChatConfig(BaseConfig):
)
# Create choice
from litellm.types.utils import Choices
choice = Choices(
index=0,
message={
@ -983,8 +599,6 @@ class OCIChatConfig(BaseConfig):
# Extract usage info
usage_info = cohere_response.chatResponse.usage
from litellm.types.utils import Usage
model_response.usage = Usage( # type: ignore[attr-defined]
prompt_tokens=usage_info.promptTokens, # type: ignore[union-attr]
completion_tokens=usage_info.completionTokens, # type: ignore[union-attr]
@ -1017,17 +631,20 @@ class OCIChatConfig(BaseConfig):
message = model_response.choices[0].message # type: ignore
response_message = completion_response.chatResponse.choices[0].message
if response_message.content 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(
response_message.toolCalls
)
# 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":
message.content = response_message.content[0].text
if response_message.toolCalls:
message.tool_calls = adapt_tools_to_openai_standard(
response_message.toolCalls
)
oci_usage = completion_response.chatResponse.usage
usage = Usage(
prompt_tokens=completion_response.chatResponse.usage.promptTokens,
completion_tokens=completion_response.chatResponse.usage.completionTokens,
total_tokens=completion_response.chatResponse.usage.totalTokens,
prompt_tokens=oci_usage.promptTokens,
completion_tokens=oci_usage.completionTokens or 0,
total_tokens=oci_usage.totalTokens,
)
model_response.usage = usage # type: ignore
@ -1111,10 +728,16 @@ class OCIChatConfig(BaseConfig):
if response.status_code != 200:
raise OCIError(status_code=response.status_code, message=response.text)
completion_stream = response.iter_text()
def split_chunks(stream: Iterator[str]) -> Iterator[str]:
"""SSE events are separated by \\n\\n — yield one data line at a time."""
for item in stream:
for chunk in item.split("\n\n"):
stripped = chunk.strip()
if stripped:
yield stripped
streaming_response = OCIStreamWrapper(
completion_stream=completion_stream,
completion_stream=split_chunks(response.iter_text()),
model=model,
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
@ -1139,7 +762,7 @@ class OCIChatConfig(BaseConfig):
del data["stream"]
if client is None or isinstance(client, HTTPHandler):
client = get_async_httpx_client(llm_provider=LlmProviders.BYTEZ, params={})
client = get_async_httpx_client(llm_provider=LlmProviders.OCI, params={})
try:
response = await client.post(
@ -1368,10 +991,12 @@ 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(
id=tool.id,
id=tool.id or f"call_{uuid.uuid4().hex[:24]}",
type="function",
function={
"name": tool.name,

View file

@ -1,4 +1,11 @@
from typing import Optional
import base64
import datetime
import hashlib
import json
import os
from dataclasses import dataclass
from typing import Any, Dict, Optional, Protocol, Tuple
from urllib.parse import urlparse
import httpx
@ -17,3 +24,361 @@ class OCIError(BaseLLMException):
message=message,
headers=headers,
)
# ---------------------------------------------------------------------------
# OCI signing protocol and helpers
# ---------------------------------------------------------------------------
class OCISignerProtocol(Protocol):
"""
Protocol for OCI request signers (e.g., oci.signer.Signer).
Compatible with the OCI Python SDK's Signer class.
See: https://docs.oracle.com/en-us/iaas/tools/python/latest/api/signing.html
"""
def do_request_sign(
self, request: Any, *, enforce_content_headers: bool = False
) -> None: ...
@dataclass
class OCIRequestWrapper:
"""
Wrapper for HTTP requests compatible with OCI signer interface.
Wraps request data in the format expected by OCI SDK signers, which require
objects with method, url, headers, body, and path_url attributes.
"""
method: str
url: str
headers: dict
body: bytes
@property
def path_url(self) -> str:
"""Returns the path + query string for OCI signing."""
parsed = urlparse(self.url)
return parsed.path + ("?" + parsed.query if parsed.query else "")
def sha256_base64(data: bytes) -> str:
digest = hashlib.sha256(data).digest()
return base64.b64encode(digest).decode()
def build_signature_string(
method: str, path: str, headers: dict, signed_headers: list
) -> str:
lines = []
for header in signed_headers:
if header == "(request-target)":
value = f"{method.lower()} {path}"
else:
value = headers[header]
lines.append(f"{header}: {value}")
return "\n".join(lines)
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(
key_str.encode("utf-8"),
password=None,
)
if not isinstance(key, rsa.RSAPrivateKey):
raise TypeError(
"The provided private key is not an RSA key, which is required for OCI signing."
)
return key
def load_private_key_from_file(file_path: str) -> Any:
"""Loads a private key from a file path."""
try:
with open(file_path, "r", encoding="utf-8") as f:
key_str = f.read().strip()
except FileNotFoundError:
raise FileNotFoundError(f"Private key file not found: {file_path}")
except OSError as e:
raise OSError(f"Failed to read private key file '{file_path}': {e}") from e
if not key_str:
raise ValueError(f"Private key file is empty: {file_path}")
return load_private_key_from_str(key_str)
# ---------------------------------------------------------------------------
# Env-var credential resolution
# ---------------------------------------------------------------------------
_OCI_REGION_ENV = "OCI_REGION"
_OCI_USER_ENV = "OCI_USER"
_OCI_FINGERPRINT_ENV = "OCI_FINGERPRINT"
_OCI_TENANCY_ENV = "OCI_TENANCY"
_OCI_KEY_FILE_ENV = "OCI_KEY_FILE"
_OCI_KEY_ENV = "OCI_KEY"
_OCI_COMPARTMENT_ID_ENV = "OCI_COMPARTMENT_ID"
def resolve_oci_credentials(optional_params: dict) -> dict:
"""
Merge OCI credentials from optional_params (explicit, always wins) and
environment variables (fallback).
Returns a dict with resolved values for:
oci_region, oci_user, oci_fingerprint, oci_tenancy,
oci_key, oci_key_file, oci_compartment_id
"""
return {
"oci_region": optional_params.get("oci_region")
or os.environ.get(_OCI_REGION_ENV)
or "us-ashburn-1",
"oci_user": optional_params.get("oci_user") or os.environ.get(_OCI_USER_ENV),
"oci_fingerprint": optional_params.get("oci_fingerprint")
or os.environ.get(_OCI_FINGERPRINT_ENV),
"oci_tenancy": optional_params.get("oci_tenancy")
or os.environ.get(_OCI_TENANCY_ENV),
"oci_key": optional_params.get("oci_key") or os.environ.get(_OCI_KEY_ENV),
"oci_key_file": optional_params.get("oci_key_file")
or os.environ.get(_OCI_KEY_FILE_ENV),
"oci_compartment_id": optional_params.get("oci_compartment_id")
or os.environ.get(_OCI_COMPARTMENT_ID_ENV),
}
def get_oci_base_url(optional_params: dict, api_base: Optional[str] = None) -> str:
"""Return the OCI inference base URL, respecting any explicit api_base override."""
if api_base:
return api_base.rstrip("/")
creds = resolve_oci_credentials(optional_params)
region = creds["oci_region"]
return f"https://inference.generativeai.{region}.oci.oraclecloud.com"
# ---------------------------------------------------------------------------
# Signing implementations (shared by chat, embed, and rerank configs)
# ---------------------------------------------------------------------------
def sign_with_oci_signer(
headers: dict,
optional_params: dict,
request_data: dict,
api_base: str,
) -> Tuple[dict, bytes]:
"""Sign a request using an OCI SDK Signer object passed in optional_params."""
oci_signer = optional_params.get("oci_signer")
body = json.dumps(request_data).encode("utf-8")
method = str(optional_params.get("method", "POST")).upper()
if method not in {"POST", "GET", "PUT", "DELETE", "PATCH"}:
raise ValueError(f"Unsupported HTTP method: {method}")
prepared_headers = {**headers}
prepared_headers.setdefault("content-type", "application/json")
prepared_headers.setdefault("content-length", str(len(body)))
request_wrapper = OCIRequestWrapper(
method=method, url=api_base, headers=prepared_headers, body=body
)
if oci_signer is None:
raise ValueError("oci_signer cannot be None when calling sign_with_oci_signer")
try:
oci_signer.do_request_sign(request_wrapper, enforce_content_headers=True)
except Exception as e:
raise OCIError(
status_code=500,
message=(
f"Failed to sign request with provided oci_signer: {str(e)}. "
"The signer must implement the OCI SDK Signer interface with a "
"do_request_sign(request, enforce_content_headers=True) method. "
"See: https://docs.oracle.com/en-us/iaas/tools/python/latest/api/signing.html"
),
) from e
headers.update(request_wrapper.headers)
return headers, body
def sign_with_manual_credentials(
headers: dict,
optional_params: dict,
request_data: dict,
api_base: str,
) -> Tuple[dict, None]:
"""Sign a request using manually provided OCI credentials (user/fingerprint/tenancy/key)."""
creds = resolve_oci_credentials(optional_params)
oci_user = creds["oci_user"]
oci_fingerprint = creds["oci_fingerprint"]
oci_tenancy = creds["oci_tenancy"]
oci_key = creds["oci_key"]
oci_key_file = creds["oci_key_file"]
if (
not oci_user
or not oci_fingerprint
or not oci_tenancy
or not (oci_key or oci_key_file)
):
raise OCIError(
status_code=401,
message=(
"Missing required OCI credentials: oci_user, oci_fingerprint, oci_tenancy, "
"and at least one of oci_key or oci_key_file. "
"These can also be supplied via environment variables: "
f"{_OCI_USER_ENV}, {_OCI_FINGERPRINT_ENV}, {_OCI_TENANCY_ENV}, {_OCI_KEY_ENV} (or {_OCI_KEY_FILE_ENV}). "
"Alternatively, provide an oci_signer object from the OCI SDK."
),
)
method = str(optional_params.get("method", "POST")).upper()
body = json.dumps(request_data).encode("utf-8")
parsed = urlparse(api_base)
path = parsed.path or "/"
host = parsed.netloc
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)
headers_to_sign: Dict[str, str] = {
"date": date,
"host": host,
"content-type": content_type,
"content-length": content_length,
"x-content-sha256": x_content_sha256,
}
signed_header_names = [
"date",
"(request-target)",
"host",
"content-length",
"content-type",
"x-content-sha256",
]
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
# Resolve the private key — prefer inline PEM content over file path
oci_key_content: Optional[str] = None
if oci_key:
if not isinstance(oci_key, str):
raise OCIError(
status_code=400,
message=(
f"oci_key must be a string containing the PEM private key content. "
f"Got type: {type(oci_key).__name__}"
),
)
oci_key_content = oci_key.replace("\\n", "\n").replace("\r\n", "\n")
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
)
if private_key is None:
raise OCIError(
status_code=400,
message="Private key is required for OCI authentication. Provide either oci_key or oci_key_file.",
)
signature = private_key.sign(
signing_string.encode("utf-8"),
padding.PKCS1v15(),
hashes.SHA256(),
)
signature_b64 = base64.b64encode(signature).decode()
key_id = f"{oci_tenancy}/{oci_user}/{oci_fingerprint}"
authorization = (
'Signature version="1",'
f'keyId="{key_id}",'
'algorithm="rsa-sha256",'
f'headers="{" ".join(signed_header_names)}",'
f'signature="{signature_b64}"'
)
headers.update(
{
"authorization": authorization,
"date": date,
"host": host,
"content-type": content_type,
"content-length": content_length,
"x-content-sha256": x_content_sha256,
}
)
return headers, None
def sign_oci_request(
headers: dict,
optional_params: dict,
request_data: dict,
api_base: str,
api_key: Optional[str] = None,
model: Optional[str] = None,
stream: Optional[bool] = None,
fake_stream: Optional[bool] = None,
) -> Tuple[dict, Optional[bytes]]:
"""
Route to the appropriate OCI signing method based on what credentials are present.
If ``oci_signer`` is in optional_params, use the OCI SDK signer object.
Otherwise use manual RSA-SHA256 signing with explicit credentials (which can
also be supplied via OCI_* environment variables).
Returns:
Tuple of (signed_headers, body_bytes_or_None)
"""
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)
def validate_oci_environment(
headers: dict,
optional_params: dict,
api_key: Optional[str] = None,
) -> dict:
"""
Populate common OCI request headers (content-type, user-agent).
Full credential validation is deferred to signing time so that credentials
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}")
return headers

View file

@ -1,8 +1,14 @@
"""
OCI Generative AI Embedding Configuration
OCI Generative AI — Embedding transformation.
Supports embedding models available on Oracle Cloud Infrastructure Generative AI service.
Uses the same authentication mechanisms as OCI chat (manual signing or OCI SDK Signer).
Endpoint: POST /20231130/actions/embedText
Supported models: cohere.embed-english-v3.0, cohere.embed-multilingual-v3.0,
cohere.embed-v4.0, and all other Cohere embed variants available on OCI
(including dedicated endpoints).
Authentication follows the same RSA-SHA256 / OCI SDK signer pattern as chat.
The base handler (base_llm_http_handler.embedding) calls sign_request after
building the body, so signing happens automatically.
Supported models:
- cohere.embed-english-v3.0
@ -10,25 +16,44 @@ Supported models:
- cohere.embed-multilingual-v3.0
- cohere.embed-multilingual-light-v3.0
- cohere.embed-english-image-v3.0
- cohere.embed-english-light-image-v3.0
- cohere.embed-multilingual-light-image-v3.0
- cohere.embed-multilingual-image-v3.0
- cohere.embed-v4.0
Reference: https://docs.oracle.com/en-us/iaas/api/#/en/generative-ai-inference/latest/EmbedTextResult/EmbedText
"""
from typing import Any, Dict, List, Optional, Union
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
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 OCIChatConfig
from litellm.llms.oci.common_utils import OCIError
from litellm.llms.oci.common_utils import (
OCIError,
get_oci_base_url,
resolve_oci_credentials,
sign_oci_request,
validate_oci_environment,
)
from litellm.types.llms.oci import (
OCIEmbedRequest,
OCIEmbedResponse,
OCIServingMode,
)
from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
from litellm.types.utils import EmbeddingResponse, Usage
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
# OCI sends up to 96 texts per embedText request (Cohere limit).
OCI_EMBED_BATCH_LIMIT = 96
# Input type mapping from OpenAI conventions to OCI/Cohere conventions
_INPUT_TYPE_MAP = {
"search_document": "SEARCH_DOCUMENT",
@ -38,65 +63,50 @@ _INPUT_TYPE_MAP = {
}
class OCIEmbeddingConfig(BaseEmbeddingConfig):
class OCIEmbedConfig(BaseEmbeddingConfig):
"""
Configuration for OCI Generative AI Embedding API.
Transformation config for OCI Generative AI embeddings.
The OCI embedding endpoint uses the Cohere embed models hosted on OCI.
Authentication is handled via OCI request signing (manual credentials or OCI SDK Signer).
Supports both text and (on cohere.embed-v4.0) multimodal inputs.
Usage:
```python
import litellm
Authentication — same two modes as chat:
- **OCI SDK signer**: pass ``oci_signer`` in optional_params.
- **Manual RSA-SHA256**: pass ``oci_user``, ``oci_fingerprint``, ``oci_tenancy``,
and ``oci_key`` or ``oci_key_file``, or set the corresponding ``OCI_*`` env vars.
response = litellm.embedding(
model="oci/cohere.embed-english-v3.0",
input=["Hello world", "Goodbye world"],
oci_compartment_id="ocid1.compartment.oc1..xxx",
oci_region="us-ashburn-1",
oci_user="ocid1.user.oc1..xxx",
oci_fingerprint="xx:xx:xx:xx",
oci_tenancy="ocid1.tenancy.oc1..xxx",
oci_key_file="~/.oci/key.pem",
)
```
Required call-time params (via optional_params or env vars):
- ``oci_compartment_id`` / ``OCI_COMPARTMENT_ID``
- ``oci_region`` / ``OCI_REGION`` (default: ``us-ashburn-1``)
Optional call-time params:
- ``oci_serving_mode``: ``"ON_DEMAND"`` (default) or ``"DEDICATED"``
- ``oci_endpoint_id``: endpoint OCID for dedicated serving mode
- ``input_type``: ``SEARCH_DOCUMENT``, ``SEARCH_QUERY``, ``CLASSIFICATION``, ``CLUSTERING``
- ``truncate``: ``NONE``, ``START``, or ``END`` (default ``END``)
- ``dimensions``: output embedding dimensions (cohere.embed-v4.0+)
"""
def __init__(self) -> None:
# We reuse OCIChatConfig for signing logic
self._chat_config = OCIChatConfig()
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
if api_base:
return api_base
oci_region = optional_params.get("oci_region", "us-ashburn-1")
return f"https://inference.generativeai.{oci_region}.oci.oraclecloud.com/20231130/actions/embedText"
def get_supported_openai_params(self, model: str) -> list:
return [
"dimensions",
]
def get_supported_openai_params(self, model: str) -> List[str]:
return ["dimensions", "encoding_format"]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
drop_params: bool = False,
) -> dict:
# Note: OCI Cohere embed does not support custom dimensions natively,
# but we pass it through in case future models support it
if "dimensions" in non_default_params:
optional_params["dimensions"] = non_default_params["dimensions"]
for key, value in non_default_params.items():
if key == "dimensions":
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(
@ -109,49 +119,19 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
"""
Validate OCI credentials for embedding requests.
Supports both OCI SDK Signer and manual credential signing.
"""
oci_signer = optional_params.get("oci_signer")
oci_region = optional_params.get("oci_region", "us-ashburn-1")
return validate_oci_environment(headers, optional_params, api_key)
api_base = (
api_base
or f"https://inference.generativeai.{oci_region}.oci.oraclecloud.com"
)
if oci_signer is None:
oci_user = optional_params.get("oci_user")
oci_fingerprint = optional_params.get("oci_fingerprint")
oci_tenancy = optional_params.get("oci_tenancy")
oci_key = optional_params.get("oci_key")
oci_key_file = optional_params.get("oci_key_file")
oci_compartment_id = optional_params.get("oci_compartment_id")
if (
not oci_user
or not oci_fingerprint
or not oci_tenancy
or not (oci_key or oci_key_file)
or not oci_compartment_id
):
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. "
"Alternatively, provide an oci_signer object from the OCI SDK."
)
from litellm.llms.custom_httpx.http_handler import version
headers.update(
{
"content-type": "application/json",
"user-agent": f"litellm/{version}",
}
)
return headers
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
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"
def sign_request(
self,
@ -163,17 +143,14 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
model: Optional[str] = None,
stream: Optional[bool] = None,
fake_stream: Optional[bool] = None,
):
"""Delegate to OCIChatConfig's signing logic."""
return self._chat_config.sign_request(
) -> Tuple[dict, Optional[bytes]]:
return sign_oci_request(
headers=headers,
optional_params=optional_params,
request_data=request_data,
api_base=api_base,
api_key=api_key,
model=model,
stream=stream,
fake_stream=fake_stream,
)
def transform_embedding_request(
@ -182,91 +159,62 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
input: AllEmbeddingInputValues,
optional_params: dict,
headers: dict,
api_base: Optional[str] = None,
) -> dict:
"""
Transform the embedding request to OCI format.
OCI embedText API expects:
{
"compartmentId": "...",
"servingMode": {"servingType": "ON_DEMAND", "modelId": "..."},
"inputs": ["text1", "text2"],
"truncate": "END",
"inputType": "SEARCH_DOCUMENT"
}
"""
oci_compartment_id = optional_params.get("oci_compartment_id")
if not oci_compartment_id:
raise Exception(
"kwarg `oci_compartment_id` is required for OCI embedding requests"
creds = resolve_oci_credentials(optional_params)
compartment_id = creds["oci_compartment_id"]
if not compartment_id:
raise OCIError(
status_code=400,
message=(
"oci_compartment_id is required for OCI embedding requests. "
"Pass it as optional_params or set the OCI_COMPARTMENT_ID env var."
),
)
# Build serving mode
oci_serving_mode = optional_params.get("oci_serving_mode", "ON_DEMAND")
if oci_serving_mode == "DEDICATED":
oci_endpoint_id = optional_params.get("oci_endpoint_id", model)
serving_mode = {
"servingType": "DEDICATED",
"endpointId": oci_endpoint_id,
}
else:
serving_mode = {
"servingType": "ON_DEMAND",
"modelId": model,
}
# Normalize input to list of strings
# Normalise input to a flat list of strings
if isinstance(input, str):
inputs = [input]
texts = [input]
elif isinstance(input, list):
inputs = []
for item in input:
if isinstance(item, str):
inputs.append(item)
elif isinstance(item, list):
raise ValueError(
"OCI embedding does not support token-array inputs. "
"Please convert token lists to strings before calling embedding()."
)
else:
inputs.append(str(item))
texts = [item if isinstance(item, str) else str(item) for item in input]
else:
inputs = [str(input)]
texts = [str(input)]
# Build request data — OCI embedText API expects inputs, truncate,
# and inputType at the top level alongside compartmentId and servingMode
request_data: Dict[str, Any] = {
"compartmentId": oci_compartment_id,
"servingMode": serving_mode,
"inputs": inputs,
"truncate": optional_params.get("truncate", "END"),
}
if len(texts) > OCI_EMBED_BATCH_LIMIT:
raise OCIError(
status_code=400,
message=(
f"OCI embedText accepts at most {OCI_EMBED_BATCH_LIMIT} inputs per request "
f"(got {len(texts)}). Batch your requests."
),
)
# Map input_type if provided
serving_mode_type = optional_params.get("oci_serving_mode", "ON_DEMAND").upper()
if serving_mode_type not in {"ON_DEMAND", "DEDICATED"}:
raise OCIError(
status_code=400,
message="oci_serving_mode must be 'ON_DEMAND' or 'DEDICATED'.",
)
if serving_mode_type == "DEDICATED":
endpoint_id = optional_params.get("oci_endpoint_id", model)
serving_mode = OCIServingMode(servingType="DEDICATED", endpointId=endpoint_id)
else:
serving_mode = OCIServingMode(servingType="ON_DEMAND", modelId=model)
# Map input_type from OpenAI convention to OCI/Cohere convention
input_type = optional_params.get("input_type")
if input_type:
mapped_type = _INPUT_TYPE_MAP.get(input_type.lower(), input_type.upper())
request_data["inputType"] = mapped_type
input_type = _INPUT_TYPE_MAP.get(input_type.lower(), input_type.upper())
# Sign the request using the same URL the HTTP handler will POST to
signing_url = self.get_complete_url(
api_base=api_base,
api_key=None,
model=model,
optional_params=optional_params,
litellm_params={},
request = OCIEmbedRequest(
compartmentId=compartment_id,
servingMode=serving_mode,
inputs=texts,
inputType=input_type,
truncate=optional_params.get("truncate", "END"),
outputDimensions=optional_params.get("outputDimensions"),
)
signed_headers, body = self.sign_request(
headers=headers,
optional_params=optional_params,
request_data=request_data,
api_base=signing_url,
)
headers.update(signed_headers)
return request_data
return request.model_dump(exclude_none=True)
def transform_embedding_response(
self,
@ -274,63 +222,49 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
raw_response: httpx.Response,
model_response: EmbeddingResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
request_data: dict = {},
optional_params: dict = {},
litellm_params: dict = {},
api_key: Optional[str],
request_data: dict,
optional_params: dict,
litellm_params: dict,
) -> EmbeddingResponse:
"""
Transform OCI embedding response to standard EmbeddingResponse format.
OCI response format:
{
"embeddings": [[0.1, 0.2, ...], [0.3, 0.4, ...]],
"modelId": "cohere.embed-english-v3.0",
"modelVersion": "3.0",
"inputTextTokenCounts": [5, 4]
}
"""
if raw_response.status_code != 200:
raise OCIError(
message=raw_response.text,
status_code=raw_response.status_code,
message=raw_response.text,
)
try:
raw_response_json = raw_response.json()
except Exception:
json_response = raw_response.json()
except Exception as e:
raise OCIError(
message=raw_response.text,
status_code=raw_response.status_code,
message=f"Failed to parse OCI embed response as JSON: {e}",
)
embeddings = raw_response_json.get("embeddings", [])
model_id = raw_response_json.get("modelId", model)
# Build response data in OpenAI format
embedding_data = []
for idx, embedding in enumerate(embeddings):
embedding_data.append(
{
"object": "embedding",
"index": idx,
"embedding": embedding,
}
try:
parsed = OCIEmbedResponse(**json_response)
except Exception as e:
raise OCIError(
status_code=500,
message=f"OCI embed response does not match expected schema: {e}",
)
model_response.model = model_id
model_response.data = embedding_data
model_response.object = "list"
model_response.model = parsed.modelId
model_response.data = [
{
"object": "embedding",
"index": i,
"embedding": embedding,
}
for i, embedding in enumerate(parsed.embeddings)
]
# Calculate token usage
input_token_counts = raw_response_json.get("inputTextTokenCounts", [])
total_tokens = sum(input_token_counts) if input_token_counts else 0
usage = Usage(
prompt_tokens=total_tokens,
total_tokens=total_tokens,
)
model_response.usage = usage
if parsed.usage is not None:
model_response.usage = Usage(
prompt_tokens=parsed.usage.promptTokens,
completion_tokens=0,
total_tokens=parsed.usage.totalTokens,
)
return model_response
@ -340,8 +274,8 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
status_code: int,
headers: Union[dict, httpx.Headers],
) -> BaseLLMException:
return OCIError(
message=error_message,
status_code=status_code,
headers=headers if isinstance(headers, httpx.Headers) else None,
)
return OCIError(status_code=status_code, message=error_message)
# Alias for backwards compatibility with any code that imports OCIEmbeddingConfig
OCIEmbeddingConfig = OCIEmbedConfig

View file

@ -5116,6 +5116,24 @@ def embedding( # noqa: PLR0915
client=client,
aembedding=aembedding,
)
elif custom_llm_provider == "oci":
if headers is None:
headers = {}
response = base_llm_http_handler.embedding(
model=model,
input=input,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
api_key=api_key,
logging_obj=logging,
timeout=timeout,
model_response=EmbeddingResponse(),
optional_params=optional_params,
client=client,
aembedding=aembedding,
litellm_params=litellm_params_dict,
headers=headers,
)
elif custom_llm_provider == "cohere" or custom_llm_provider == "cohere_chat":
cohere_key = (
api_key

View file

@ -57,7 +57,7 @@ OCIContentPartUnion = Union[OCITextContentPart, OCIImageContentPart]
class OCIToolCall(BaseModel):
"""Represents a tool call made by the model."""
id: str
id: Optional[str] = None # absent in some provider responses (e.g. Google via OCI)
type: Literal["FUNCTION"] = "FUNCTION"
name: str
arguments: str # Arguments should be a JSON-serialized string
@ -141,7 +141,9 @@ class OCIResponseUsage(BaseModel):
"""Token usage in the OCI response."""
promptTokens: int
completionTokens: int
# completionTokens may be absent for reasoning models when all the output
# budget is consumed by reasoning tokens before any visible content is produced.
completionTokens: Optional[int] = None
totalTokens: int
completionTokensDetails: Optional[OCICompletionTokenDetails] = None
promptTokensDetails: Optional[OCIPromptTokensDetails] = None
@ -151,7 +153,9 @@ class OCIResponseChoice(BaseModel):
"""A completion choice in the OCI response."""
index: int
message: OCIMessage
# message is absent when a reasoning model exhausts max_tokens in the
# reasoning phase without producing any visible content.
message: Optional[OCIMessage] = None
finishReason: Optional[str] = None
logprobs: Optional[Dict[str, Any]] = None
@ -394,3 +398,35 @@ class CohereChatResult(BaseModel):
modelId: str
modelVersion: str
chatResponse: CohereChatResponse
# ---------------------------------------------------------------------------
# OCI Embed types
# ---------------------------------------------------------------------------
class OCIEmbedRequest(BaseModel):
"""Request body for POST /20231130/actions/embedText."""
compartmentId: str
servingMode: OCIServingMode
inputs: List[str]
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
class OCIEmbedUsage(BaseModel):
promptTokens: int
totalTokens: int
class OCIEmbedResponse(BaseModel):
"""Response body from POST /20231130/actions/embedText."""
embeddings: List[List[float]]
modelId: str
modelVersion: str
usage: Optional[OCIEmbedUsage] = None

View file

@ -8383,6 +8383,10 @@ class ProviderConfigManager:
return litellm.InfinityEmbeddingConfig()
elif litellm.LlmProviders.SAMBANOVA == provider:
return litellm.SambaNovaEmbeddingConfig()
elif litellm.LlmProviders.OCI == provider:
from litellm.llms.oci.embed.transformation import OCIEmbedConfig
return OCIEmbedConfig()
elif (
litellm.LlmProviders.COHERE == provider
or litellm.LlmProviders.COHERE_CHAT == provider

View file

@ -25641,6 +25641,32 @@
"supports_vision": true,
"supports_web_search": true
},
"oci/meta.llama-3.1-8b-instruct": {
"input_cost_per_token": 7.2e-07,
"litellm_provider": "oci",
"max_input_tokens": 128000,
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
"output_cost_per_token": 7.2e-07,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_streaming": true
},
"oci/meta.llama-3.1-70b-instruct": {
"input_cost_per_token": 7.2e-07,
"litellm_provider": "oci",
"max_input_tokens": 128000,
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
"output_cost_per_token": 7.2e-07,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_streaming": true
},
"oci/meta.llama-3.1-405b-instruct": {
"input_cost_per_token": 1.068e-05,
"litellm_provider": "oci",
@ -25651,7 +25677,8 @@
"output_cost_per_token": 1.068e-05,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing",
"supports_function_calling": true,
"supports_response_schema": false
"supports_response_schema": false,
"supports_streaming": true
},
"oci/meta.llama-3.2-90b-vision-instruct": {
"input_cost_per_token": 2e-06,
@ -25664,6 +25691,7 @@
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_streaming": true,
"supports_vision": true
},
"oci/meta.llama-3.3-70b-instruct": {
@ -25676,7 +25704,8 @@
"output_cost_per_token": 7.2e-07,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing",
"supports_function_calling": true,
"supports_response_schema": false
"supports_response_schema": false,
"supports_streaming": true
},
"oci/meta.llama-4-maverick-17b-128e-instruct-fp8": {
"input_cost_per_token": 7.2e-07,
@ -25688,7 +25717,8 @@
"output_cost_per_token": 7.2e-07,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing",
"supports_function_calling": true,
"supports_response_schema": false
"supports_response_schema": false,
"supports_streaming": true
},
"oci/meta.llama-4-scout-17b-16e-instruct": {
"input_cost_per_token": 7.2e-07,
@ -25700,7 +25730,8 @@
"output_cost_per_token": 7.2e-07,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing",
"supports_function_calling": true,
"supports_response_schema": false
"supports_response_schema": false,
"supports_streaming": true
},
"oci/xai.grok-3": {
"input_cost_per_token": 3e-06,
@ -25712,7 +25743,8 @@
"output_cost_per_token": 1.5e-05,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing",
"supports_function_calling": true,
"supports_response_schema": false
"supports_response_schema": false,
"supports_streaming": true
},
"oci/xai.grok-3-fast": {
"input_cost_per_token": 5e-06,
@ -25724,7 +25756,8 @@
"output_cost_per_token": 2.5e-05,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing",
"supports_function_calling": true,
"supports_response_schema": false
"supports_response_schema": false,
"supports_streaming": true
},
"oci/xai.grok-3-mini": {
"input_cost_per_token": 3e-07,
@ -25736,7 +25769,8 @@
"output_cost_per_token": 5e-07,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing",
"supports_function_calling": true,
"supports_response_schema": false
"supports_response_schema": false,
"supports_streaming": true
},
"oci/xai.grok-3-mini-fast": {
"input_cost_per_token": 6e-07,
@ -25748,7 +25782,8 @@
"output_cost_per_token": 4e-06,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing",
"supports_function_calling": true,
"supports_response_schema": false
"supports_response_schema": false,
"supports_streaming": true
},
"oci/xai.grok-4": {
"input_cost_per_token": 3e-06,
@ -25760,7 +25795,8 @@
"output_cost_per_token": 1.5e-05,
"source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing",
"supports_function_calling": true,
"supports_response_schema": false
"supports_response_schema": false,
"supports_streaming": true
},
"oci/cohere.command-latest": {
"input_cost_per_token": 1.56e-06,
@ -25772,7 +25808,8 @@
"output_cost_per_token": 1.56e-06,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_function_calling": true,
"supports_response_schema": false
"supports_response_schema": false,
"supports_streaming": true
},
"oci/cohere.command-a-03-2025": {
"input_cost_per_token": 1.56e-06,
@ -25784,7 +25821,8 @@
"output_cost_per_token": 1.56e-06,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_function_calling": true,
"supports_response_schema": false
"supports_response_schema": false,
"supports_streaming": true
},
"oci/cohere.command-plus-latest": {
"input_cost_per_token": 1.56e-06,
@ -25796,7 +25834,189 @@
"output_cost_per_token": 1.56e-06,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_function_calling": true,
"supports_response_schema": false
"supports_response_schema": false,
"supports_streaming": true
},
"oci/google.gemini-2.5-flash": {
"input_cost_per_token": 1.875e-07,
"litellm_provider": "oci",
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_streaming": true,
"supports_vision": true
},
"oci/google.gemini-2.5-pro": {
"input_cost_per_token": 1.875e-07,
"litellm_provider": "oci",
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_streaming": true,
"supports_vision": true
},
"oci/google.gemini-2.5-flash-lite": {
"input_cost_per_token": 1.875e-07,
"litellm_provider": "oci",
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_streaming": true,
"supports_vision": true
},
"oci/openai.gpt-oss-120b": {
"input_cost_per_token": 0.0,
"litellm_provider": "oci",
"max_input_tokens": 131072,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 0.0,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_streaming": true
},
"oci/openai.gpt-oss-20b": {
"input_cost_per_token": 0.0,
"litellm_provider": "oci",
"max_input_tokens": 131072,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 0.0,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_streaming": true
},
"oci/cohere.command-a-vision": {
"input_cost_per_token": 1.56e-06,
"litellm_provider": "oci",
"max_input_tokens": 256000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 1.56e-06,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_streaming": true,
"supports_vision": true
},
"oci/cohere.command-a-reasoning": {
"input_cost_per_token": 1.56e-06,
"litellm_provider": "oci",
"max_input_tokens": 256000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 1.56e-06,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_function_calling": false,
"supports_response_schema": false,
"supports_streaming": true
},
"oci/meta.llama-4-maverick-17b-128e-instruct-fp8": {
"input_cost_per_token": 0.0,
"litellm_provider": "oci",
"max_input_tokens": 1048576,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 0.0,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_streaming": true,
"supports_vision": true
},
"oci/meta.llama-4-scout-17b-16e-instruct": {
"input_cost_per_token": 0.0,
"litellm_provider": "oci",
"max_input_tokens": 10485760,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 0.0,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_function_calling": true,
"supports_response_schema": false,
"supports_streaming": true
},
"oci/cohere.embed-english-v3.0": {
"input_cost_per_token": 1e-07,
"litellm_provider": "oci",
"max_input_tokens": 512,
"mode": "embedding",
"output_vector_size": 1024,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/"
},
"oci/cohere.embed-multilingual-v3.0": {
"input_cost_per_token": 1e-07,
"litellm_provider": "oci",
"max_input_tokens": 512,
"mode": "embedding",
"output_vector_size": 1024,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/"
},
"oci/cohere.embed-english-light-v3.0": {
"input_cost_per_token": 1e-07,
"litellm_provider": "oci",
"max_input_tokens": 512,
"mode": "embedding",
"output_vector_size": 384,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/"
},
"oci/cohere.embed-multilingual-light-v3.0": {
"input_cost_per_token": 1e-07,
"litellm_provider": "oci",
"max_input_tokens": 512,
"mode": "embedding",
"output_vector_size": 384,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/"
},
"oci/cohere.embed-english-image-v3.0": {
"input_cost_per_token": 1e-07,
"litellm_provider": "oci",
"max_input_tokens": 512,
"mode": "embedding",
"output_vector_size": 1024,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_vision": true
},
"oci/cohere.embed-multilingual-image-v3.0": {
"input_cost_per_token": 1e-07,
"litellm_provider": "oci",
"max_input_tokens": 512,
"mode": "embedding",
"output_vector_size": 1024,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_vision": true
},
"oci/cohere.embed-v4.0": {
"input_cost_per_token": 1e-07,
"litellm_provider": "oci",
"max_input_tokens": 128000,
"mode": "embedding",
"output_vector_size": 1536,
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
"supports_vision": true
},
"oci/cohere.command-a-reasoning-08-2025": {
"input_cost_per_token": 1.56e-06,

View file

@ -0,0 +1,380 @@
"""
Unit tests for OCI Generative AI embedding transformation.
These tests exercise the transformation layer only — no real OCI calls are made.
"""
import json
import os
import sys
from typing import Any
from unittest.mock import MagicMock
import httpx
import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.oci.common_utils import OCIError
from litellm.llms.oci.embed.transformation import OCI_EMBED_BATCH_LIMIT, OCIEmbedConfig
from litellm.types.utils import EmbeddingResponse, Usage
# ---------------------------------------------------------------------------
# Test fixtures
# ---------------------------------------------------------------------------
COMPARTMENT_ID = "ocid1.compartment.oc1..test"
BASE_PARAMS = {
"oci_region": "us-ashburn-1",
"oci_user": "ocid1.user.oc1..test",
"oci_fingerprint": "aa:bb:cc:dd",
"oci_tenancy": "ocid1.tenancy.oc1..test",
"oci_compartment_id": COMPARTMENT_ID,
"oci_key": "-----BEGIN RSA PRIVATE KEY-----\nfakekey\n-----END RSA PRIVATE KEY-----",
}
class TestOCIEmbedConfig:
def _config(self) -> OCIEmbedConfig:
return OCIEmbedConfig()
# ------------------------------------------------------------------
# validate_environment
# ------------------------------------------------------------------
def test_validate_environment_sets_headers(self):
cfg = self._config()
headers = cfg.validate_environment(
headers={},
model="oci/cohere.embed-v3.0",
messages=[],
optional_params=BASE_PARAMS,
litellm_params={},
)
assert headers["content-type"] == "application/json"
assert "litellm/" in headers["user-agent"]
# ------------------------------------------------------------------
# get_complete_url
# ------------------------------------------------------------------
def test_get_complete_url_default_region(self):
cfg = self._config()
url = cfg.get_complete_url(
api_base=None,
api_key=None,
model="cohere.embed-v3.0",
optional_params={"oci_region": "us-chicago-1"},
litellm_params={},
)
assert url == "https://inference.generativeai.us-chicago-1.oci.oraclecloud.com/20231130/actions/embedText"
def test_get_complete_url_respects_api_base(self):
cfg = self._config()
url = cfg.get_complete_url(
api_base="https://custom.endpoint.example.com",
api_key=None,
model="cohere.embed-v3.0",
optional_params={},
litellm_params={},
)
assert url == "https://custom.endpoint.example.com/20231130/actions/embedText"
def test_get_complete_url_strips_trailing_slash(self):
cfg = self._config()
url = cfg.get_complete_url(
api_base="https://custom.endpoint.example.com/",
api_key=None,
model="cohere.embed-v3.0",
optional_params={},
litellm_params={},
)
assert not url.endswith("//")
assert url.endswith("/20231130/actions/embedText")
# ------------------------------------------------------------------
# transform_embedding_request
# ------------------------------------------------------------------
def test_transform_request_single_string(self):
cfg = self._config()
result = cfg.transform_embedding_request(
model="cohere.embed-v3.0",
input="hello world",
optional_params={"oci_compartment_id": COMPARTMENT_ID},
headers={},
)
assert result["compartmentId"] == COMPARTMENT_ID
assert result["servingMode"]["servingType"] == "ON_DEMAND"
assert result["servingMode"]["modelId"] == "cohere.embed-v3.0"
assert result["inputs"] == ["hello world"]
def test_transform_request_list_of_texts(self):
cfg = self._config()
texts = ["hello", "world"]
result = cfg.transform_embedding_request(
model="cohere.embed-v3.0",
input=texts,
optional_params={"oci_compartment_id": COMPARTMENT_ID},
headers={},
)
assert result["inputs"] == texts
def test_transform_request_with_input_type(self):
cfg = self._config()
result = cfg.transform_embedding_request(
model="cohere.embed-v3.0",
input=["query"],
optional_params={
"oci_compartment_id": COMPARTMENT_ID,
"input_type": "SEARCH_QUERY",
},
headers={},
)
assert result["inputType"] == "SEARCH_QUERY"
def test_transform_request_with_output_dimensions(self):
cfg = self._config()
result = cfg.transform_embedding_request(
model="cohere.embed-v4.0",
input=["text"],
optional_params={
"oci_compartment_id": COMPARTMENT_ID,
"outputDimensions": 512,
},
headers={},
)
assert result["outputDimensions"] == 512
def test_transform_request_dedicated_serving_mode(self):
cfg = self._config()
result = cfg.transform_embedding_request(
model="cohere.embed-v3.0",
input=["text"],
optional_params={
"oci_compartment_id": COMPARTMENT_ID,
"oci_serving_mode": "DEDICATED",
"oci_endpoint_id": "ocid1.genaiendpoint.oc1..test",
},
headers={},
)
assert result["servingMode"]["servingType"] == "DEDICATED"
assert result["servingMode"]["endpointId"] == "ocid1.genaiendpoint.oc1..test"
assert "modelId" not in result["servingMode"]
def test_transform_request_missing_compartment_id_raises(self):
cfg = self._config()
with pytest.raises(OCIError) as exc_info:
cfg.transform_embedding_request(
model="cohere.embed-v3.0",
input=["text"],
optional_params={},
headers={},
)
assert exc_info.value.status_code == 400
assert "oci_compartment_id" in str(exc_info.value)
def test_transform_request_batch_limit_exceeded_raises(self):
cfg = self._config()
texts = ["text"] * (OCI_EMBED_BATCH_LIMIT + 1)
with pytest.raises(OCIError) as exc_info:
cfg.transform_embedding_request(
model="cohere.embed-v3.0",
input=texts,
optional_params={"oci_compartment_id": COMPARTMENT_ID},
headers={},
)
assert exc_info.value.status_code == 400
assert str(OCI_EMBED_BATCH_LIMIT) in str(exc_info.value)
def test_transform_request_invalid_serving_mode_raises(self):
cfg = self._config()
with pytest.raises(OCIError) as exc_info:
cfg.transform_embedding_request(
model="cohere.embed-v3.0",
input=["text"],
optional_params={
"oci_compartment_id": COMPARTMENT_ID,
"oci_serving_mode": "INVALID",
},
headers={},
)
assert exc_info.value.status_code == 400
def test_transform_request_none_input_becomes_string(self):
"""Non-list, non-string inputs are coerced to str."""
cfg = self._config()
result = cfg.transform_embedding_request(
model="cohere.embed-v3.0",
input=42, # type: ignore
optional_params={"oci_compartment_id": COMPARTMENT_ID},
headers={},
)
assert result["inputs"] == ["42"]
# ------------------------------------------------------------------
# transform_embedding_response
# ------------------------------------------------------------------
def _mock_response(self, status_code: int, body: dict) -> httpx.Response:
return httpx.Response(
status_code=status_code,
content=json.dumps(body).encode(),
headers={"content-type": "application/json"},
)
def test_transform_response_success(self):
cfg = self._config()
model_response = EmbeddingResponse()
raw = self._mock_response(
200,
{
"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},
},
)
result = cfg.transform_embedding_response(
model="cohere.embed-v3.0",
raw_response=raw,
model_response=model_response,
logging_obj=MagicMock(),
api_key=None,
request_data={},
optional_params={},
litellm_params={},
)
assert len(result.data) == 2
assert result.data[0]["embedding"] == [0.1, 0.2, 0.3]
assert result.data[1]["index"] == 1
assert result.model == "cohere.embed-v3.0"
assert result.usage.prompt_tokens == 10
def test_transform_response_no_usage(self):
cfg = self._config()
model_response = EmbeddingResponse()
raw = self._mock_response(
200,
{
"embeddings": [[0.1]],
"modelId": "cohere.embed-v3.0",
"modelVersion": "3.0.0",
},
)
result = cfg.transform_embedding_response(
model="cohere.embed-v3.0",
raw_response=raw,
model_response=model_response,
logging_obj=MagicMock(),
api_key=None,
request_data={},
optional_params={},
litellm_params={},
)
assert len(result.data) == 1
def test_transform_response_http_error_raises(self):
cfg = self._config()
raw = self._mock_response(401, {"error": "Unauthorized"})
with pytest.raises(OCIError) as exc_info:
cfg.transform_embedding_response(
model="cohere.embed-v3.0",
raw_response=raw,
model_response=EmbeddingResponse(),
logging_obj=MagicMock(),
api_key=None,
request_data={},
optional_params={},
litellm_params={},
)
assert exc_info.value.status_code == 401
def test_transform_response_invalid_json_raises(self):
cfg = self._config()
raw = httpx.Response(
status_code=200,
content=b"not-json",
headers={"content-type": "text/plain"},
)
with pytest.raises(OCIError):
cfg.transform_embedding_response(
model="cohere.embed-v3.0",
raw_response=raw,
model_response=EmbeddingResponse(),
logging_obj=MagicMock(),
api_key=None,
request_data={},
optional_params={},
litellm_params={},
)
# ------------------------------------------------------------------
# map_openai_params
# ------------------------------------------------------------------
def test_map_openai_params_dimensions(self):
cfg = self._config()
result = cfg.map_openai_params(
non_default_params={"dimensions": 512},
optional_params={},
model="cohere.embed-v4.0",
)
assert result["outputDimensions"] == 512
def test_map_openai_params_encoding_format_raises_without_drop(self):
cfg = self._config()
with pytest.raises(OCIError):
cfg.map_openai_params(
non_default_params={"encoding_format": "float"},
optional_params={},
model="cohere.embed-v3.0",
)
def test_map_openai_params_encoding_format_dropped_silently(self):
cfg = self._config()
result = cfg.map_openai_params(
non_default_params={"encoding_format": "float"},
optional_params={},
model="cohere.embed-v3.0",
drop_params=True,
)
assert "encoding_format" not in result
# ------------------------------------------------------------------
# env var credential resolution
# ------------------------------------------------------------------
def test_env_var_compartment_id(self, monkeypatch):
monkeypatch.setenv("OCI_COMPARTMENT_ID", "ocid1.compartment.from.env")
cfg = self._config()
result = cfg.transform_embedding_request(
model="cohere.embed-v3.0",
input=["hello"],
optional_params={}, # no compartment_id in params
headers={},
)
assert result["compartmentId"] == "ocid1.compartment.from.env"
def test_explicit_param_overrides_env_var(self, monkeypatch):
monkeypatch.setenv("OCI_COMPARTMENT_ID", "ocid1.compartment.from.env")
cfg = self._config()
result = cfg.transform_embedding_request(
model="cohere.embed-v3.0",
input=["hello"],
optional_params={"oci_compartment_id": "ocid1.compartment.explicit"},
headers={},
)
assert result["compartmentId"] == "ocid1.compartment.explicit"
def test_env_var_region_used_in_url(self, monkeypatch):
monkeypatch.setenv("OCI_REGION", "eu-frankfurt-1")
cfg = self._config()
url = cfg.get_complete_url(
api_base=None,
api_key=None,
model="cohere.embed-v3.0",
optional_params={}, # no explicit region
litellm_params={},
)
assert "eu-frankfurt-1" in url

View file

@ -0,0 +1,742 @@
"""
OCI Generative AI — end-to-end integration tests.
These tests make REAL calls to OCI. They are skipped automatically when the
standard ~/.oci/config is absent or when OCI_TEST_COMPARTMENT_ID is not set.
Prerequisites
-------------
- ~/.oci/config with a valid [DEFAULT] profile
- Private key referenced by key_file in that profile
- Sufficient IAM policies to call the Generative AI inference service
Environment variables (all optional — fall back to ~/.oci/config values):
OCI_TEST_REGION OCI region (default: us-chicago-1)
OCI_TEST_COMPARTMENT_ID compartment OCID (default: tenancy root from config)
Run only these tests:
pytest tests/test_litellm/llms/oci/test_oci_integration.py -v
"""
import os
import sys
from typing import Generator
import pytest
sys.path.insert(0, os.path.abspath("../../../.."))
# ---------------------------------------------------------------------------
# Fixtures / helpers
# ---------------------------------------------------------------------------
OCI_CONFIG_FILE = os.path.expanduser("~/.oci/config")
_OCI_AVAILABLE = os.path.isfile(OCI_CONFIG_FILE)
pytestmark = pytest.mark.skipif(
not _OCI_AVAILABLE,
reason="~/.oci/config not found — skipping OCI integration tests",
)
@pytest.fixture(scope="module")
def oci_signer():
"""Return an oci.Signer built from ~/.oci/config [DEFAULT]."""
oci = pytest.importorskip("oci")
config = oci.config.from_file()
return oci.Signer(
tenancy=config["tenancy"],
user=config["user"],
fingerprint=config["fingerprint"],
private_key_file_location=config["key_file"],
)
@pytest.fixture(scope="module")
def oci_params(oci_signer) -> dict:
"""Common OCI call-time parameters shared by all tests."""
oci = pytest.importorskip("oci")
config = oci.config.from_file()
compartment_id = os.environ.get("OCI_TEST_COMPARTMENT_ID", config["tenancy"])
region = os.environ.get("OCI_TEST_REGION", "us-chicago-1")
return {
"oci_signer": oci_signer,
"oci_compartment_id": compartment_id,
"oci_region": region,
}
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _chat(model: str, message: str, params: dict, max_tokens: int = 64) -> str:
"""Run a single-turn completion and return the text content."""
import litellm
resp = litellm.completion(
model=f"oci/{model}",
messages=[{"role": "user", "content": message}],
max_tokens=max_tokens,
**params,
)
# Reasoning models may return None content when all budget is used by reasoning
return resp.choices[0].message.content or ""
# ---------------------------------------------------------------------------
# Chat tests — one per vendor family
# ---------------------------------------------------------------------------
class TestOCIChatMeta:
"""Meta Llama models (GENERIC apiFormat)."""
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,
)
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
)
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 "")
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
class TestOCIChatXAI:
"""xAI Grok models (GENERIC apiFormat)."""
MODEL = "xai.grok-3-mini"
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
class TestOCIChatCohere:
"""Cohere Command models (COHERE apiFormat)."""
MODEL = "cohere.command-latest"
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
# ---------------------------------------------------------------------------
# Embedding tests
# ---------------------------------------------------------------------------
class TestOCIEmbeddings:
def test_english_v3_basic(self, oci_params):
import litellm
resp = litellm.embedding(
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
def test_english_v3_batch(self, oci_params):
import litellm
texts = [
"The quick brown fox",
"jumps over the lazy dog",
"Paris is the capital of France",
]
resp = litellm.embedding(
model="oci/cohere.embed-english-v3.0",
input=texts,
input_type="SEARCH_DOCUMENT",
**oci_params,
)
assert len(resp.data) == 3
for i, item in enumerate(resp.data):
assert item["index"] == i
assert len(item["embedding"]) == 1024
def test_multilingual_v3(self, oci_params):
import litellm
resp = litellm.embedding(
model="oci/cohere.embed-multilingual-v3.0",
input=["Bonjour le monde", "Hola mundo"],
input_type="SEARCH_DOCUMENT",
**oci_params,
)
assert len(resp.data) == 2
assert len(resp.data[0]["embedding"]) == 1024
def test_search_query_input_type(self, oci_params):
import litellm
resp = litellm.embedding(
model="oci/cohere.embed-english-v3.0",
input=["What is the capital of France?"],
input_type="SEARCH_QUERY",
**oci_params,
)
assert len(resp.data[0]["embedding"]) == 1024
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",
input=[
"The cat sat on the mat",
"A feline rested on the rug",
"The stock market crashed today",
],
input_type="SEARCH_DOCUMENT",
**oci_params,
)
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))
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})"
)
def test_embed_v4(self, oci_params):
import litellm
resp = litellm.embedding(
model="oci/cohere.embed-v4.0",
input=["Hello world"],
input_type="SEARCH_DOCUMENT",
**oci_params,
)
assert len(resp.data) == 1
assert len(resp.data[0]["embedding"]) == 1536
def test_usage_tokens(self, oci_params):
import litellm
resp = litellm.embedding(
model="oci/cohere.embed-english-v3.0",
input=["short text", "another short text"],
input_type="SEARCH_DOCUMENT",
**oci_params,
)
assert resp.usage.prompt_tokens > 0
assert resp.usage.total_tokens == resp.usage.prompt_tokens
# ---------------------------------------------------------------------------
# Env-var credential path
# ---------------------------------------------------------------------------
class TestOCIEnvVarCredentials:
"""Verify that OCI_* env vars are picked up without explicit params."""
def test_completion_via_env_vars(self, monkeypatch):
"""Completion works when credentials are set through environment variables."""
oci = pytest.importorskip("oci")
config = oci.config.from_file()
key_path = os.path.expanduser(config["key_file"])
with open(key_path) as f:
key_pem = f.read()
monkeypatch.setenv("OCI_REGION", "us-chicago-1")
monkeypatch.setenv("OCI_USER", config["user"])
monkeypatch.setenv("OCI_FINGERPRINT", config["fingerprint"])
monkeypatch.setenv("OCI_TENANCY", config["tenancy"])
monkeypatch.setenv("OCI_KEY", key_pem)
monkeypatch.setenv("OCI_COMPARTMENT_ID", config["tenancy"])
import litellm
resp = litellm.completion(
model="oci/meta.llama-3.3-70b-instruct",
messages=[{"role": "user", "content": "Reply with only the word: pong"}],
max_tokens=10,
)
assert resp.choices[0].message.content is not None
def test_embedding_via_env_vars(self, monkeypatch):
oci = pytest.importorskip("oci")
config = oci.config.from_file()
key_path = os.path.expanduser(config["key_file"])
with open(key_path) as f:
key_pem = f.read()
monkeypatch.setenv("OCI_REGION", "us-chicago-1")
monkeypatch.setenv("OCI_USER", config["user"])
monkeypatch.setenv("OCI_FINGERPRINT", config["fingerprint"])
monkeypatch.setenv("OCI_TENANCY", config["tenancy"])
monkeypatch.setenv("OCI_KEY", key_pem)
monkeypatch.setenv("OCI_COMPARTMENT_ID", config["tenancy"])
import litellm
resp = litellm.embedding(
model="oci/cohere.embed-english-v3.0",
input=["hello"],
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)