From 064707484651a9981029f29454d9b50d34854672 Mon Sep 17 00:00:00 2001 From: Federico Kamelhar Date: Sun, 5 Apr 2026 00:16:19 -0400 Subject: [PATCH] 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 --- litellm/llms/custom_httpx/llm_http_handler.py | 36 +- litellm/llms/oci/chat/transformation.py | 569 +++----------- litellm/llms/oci/common_utils.py | 367 ++++++++- litellm/llms/oci/embed/transformation.py | 380 ++++----- litellm/main.py | 18 + litellm/types/llms/oci.py | 42 +- litellm/utils.py | 4 + model_prices_and_context_window.json | 244 +++++- .../embed/test_oci_embed_transformation.py | 380 +++++++++ .../test_litellm/llms/oci/rerank/__init__.py | 0 .../llms/oci/test_oci_integration.py | 742 ++++++++++++++++++ 11 files changed, 2064 insertions(+), 718 deletions(-) create mode 100644 tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py create mode 100644 tests/test_litellm/llms/oci/rerank/__init__.py create mode 100644 tests/test_litellm/llms/oci/test_oci_integration.py diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 2ff63cc2d7f..27791f034bb 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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) diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index 62104e921a4..ce151dc78b8 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -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, diff --git a/litellm/llms/oci/common_utils.py b/litellm/llms/oci/common_utils.py index 661a6c89e4b..2c95f212576 100644 --- a/litellm/llms/oci/common_utils.py +++ b/litellm/llms/oci/common_utils.py @@ -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 diff --git a/litellm/llms/oci/embed/transformation.py b/litellm/llms/oci/embed/transformation.py index 1dcd8c5213c..27bf2b1acb6 100644 --- a/litellm/llms/oci/embed/transformation.py +++ b/litellm/llms/oci/embed/transformation.py @@ -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 diff --git a/litellm/main.py b/litellm/main.py index b5364f8ba17..ab2820fe057 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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 diff --git a/litellm/types/llms/oci.py b/litellm/types/llms/oci.py index e041810158a..16bd6427717 100644 --- a/litellm/types/llms/oci.py +++ b/litellm/types/llms/oci.py @@ -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 + + diff --git a/litellm/utils.py b/litellm/utils.py index 54cea313b0c..db61f454a41 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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 diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 94f0f1e78d3..c165413bb37 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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, diff --git a/tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py b/tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py new file mode 100644 index 00000000000..ad715084390 --- /dev/null +++ b/tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py @@ -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 diff --git a/tests/test_litellm/llms/oci/rerank/__init__.py b/tests/test_litellm/llms/oci/rerank/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/oci/test_oci_integration.py b/tests/test_litellm/llms/oci/test_oci_integration.py new file mode 100644 index 00000000000..5787f82614e --- /dev/null +++ b/tests/test_litellm/llms/oci/test_oci_integration.py @@ -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)