From 17d52c208c1ebfa3f31ca9bef53128d5e9fdb778 Mon Sep 17 00:00:00 2001 From: Federico Kamelhar Date: Sun, 5 Apr 2026 13:16:39 -0400 Subject: [PATCH] refactor(oci): principal-level code quality pass MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Remove _extract_text_content duplication — single definition in cohere.py, imported where needed; instance method on OCIChatConfig eliminated - Move cryptography imports to module level with _CRYPTOGRAPHY_AVAILABLE flag and _require_cryptography() guard; no more re-import on every signing call - Move litellm version import to module level via litellm._version; remove inline import inside validate_oci_environment - sign_with_manual_credentials now returns Tuple[dict, bytes] matching sign_with_oci_signer — asymmetry eliminated, Optional[bytes] guards removed throughout stream wrappers (signed_json_body: bytes = b"") - Rename _openai_to_oci_cohere_param_map → openai_to_oci_cohere_param_map for consistency with openai_to_oci_generic_param_map - Remove double-key bug in map_openai_params where responseFormat was stored under both OCI and OpenAI key names simultaneously - Remove delegating shims (adapt_messages_to_cohere_standard, adapt_tool_definitions_to_cohere_standard, _handle_generic_stream_chunk) from OCIChatConfig/OCIStreamWrapper; tests now import directly from cohere.py and generic.py where symbols live - Trim __all__ to 7 genuine public symbols; remove the 13-symbol list that existed only to support test imports - Collapse per-model integration test classes into pytest.mark.parametrize; CHAT_MODELS list is the single source of truth for model-specific config - Black + Ruff clean across all OCI files --- litellm/llms/oci/chat/cohere.py | 2 - litellm/llms/oci/chat/generic.py | 35 +- litellm/llms/oci/chat/transformation.py | 99 +-- litellm/llms/oci/common_utils.py | 85 +- litellm/llms/oci/embed/transformation.py | 4 +- litellm/types/llms/oci.py | 10 +- tests/llm_translation/test_oci_integration.py | 738 +++++++----------- .../oci/chat/test_oci_chat_transformation.py | 33 +- .../oci/chat/test_oci_cohere_tool_calls.py | 25 +- .../oci/chat/test_oci_streaming_tool_calls.py | 118 +-- .../embed/test_oci_embed_transformation.py | 5 +- 11 files changed, 433 insertions(+), 721 deletions(-) diff --git a/litellm/llms/oci/chat/cohere.py b/litellm/llms/oci/chat/cohere.py index 50ea8320f3b..8bc97ddf9e2 100644 --- a/litellm/llms/oci/chat/cohere.py +++ b/litellm/llms/oci/chat/cohere.py @@ -11,8 +11,6 @@ import json import uuid from typing import Any, Dict, List, Optional -import httpx - from litellm.llms.oci.common_utils import ( OCI_JSON_TO_PYTHON_TYPES, OCIError, diff --git a/litellm/llms/oci/chat/generic.py b/litellm/llms/oci/chat/generic.py index eec837fb66d..07811d59327 100644 --- a/litellm/llms/oci/chat/generic.py +++ b/litellm/llms/oci/chat/generic.py @@ -7,9 +7,8 @@ parsing, and streaming chunk parsing for models served with """ import datetime -import json import uuid -from typing import Any, Dict, List, Optional, Union +from typing import Dict, List, Optional, Union import httpx @@ -70,7 +69,9 @@ def adapt_messages_to_generic_oci_standard_content_message( for content_item in content: if not isinstance(content_item, dict): - raise OCIError(status_code=400, message="Each content item must be a dictionary") + raise OCIError( + status_code=400, message="Each content item must be a dictionary" + ) item_type = content_item.get("type") if not isinstance(item_type, str): @@ -119,9 +120,13 @@ def adapt_messages_to_generic_oci_standard_tool_call( tool_calls_formatted = [] for tool_call in tool_calls: if not isinstance(tool_call, dict): - raise OCIError(status_code=400, message="Each tool call must be a dictionary") + raise OCIError( + status_code=400, message="Each tool call must be a dictionary" + ) if tool_call.get("type") != "function": - raise OCIError(status_code=400, message="OCI only supports function tool calls") + raise OCIError( + status_code=400, message="OCI only supports function tool calls" + ) tool_call_id = tool_call.get("id") if not isinstance(tool_call_id, str): @@ -129,7 +134,9 @@ def adapt_messages_to_generic_oci_standard_tool_call( tool_function = tool_call.get("function") if not isinstance(tool_function, dict): - raise OCIError(status_code=400, message="Tool call `function` must be a dictionary") + raise OCIError( + status_code=400, message="Tool call `function` must be a dictionary" + ) function_name = tool_function.get("name") if not isinstance(function_name, str): @@ -186,7 +193,9 @@ def adapt_messages_to_generic_oci_standard( if role == "assistant" and tool_calls is not None: if not isinstance(tool_calls, list): - raise OCIError(status_code=400, message="Message `tool_calls` must be a list") + raise OCIError( + status_code=400, message="Message `tool_calls` must be a list" + ) new_messages.append( adapt_messages_to_generic_oci_standard_tool_call(role, tool_calls) ) @@ -240,7 +249,9 @@ def adapt_tool_definition_to_oci_standard( tool_function = tool.get("function") if not isinstance(tool_function, dict): - raise OCIError(status_code=400, message="Tool `function` must be a dictionary") + raise OCIError( + status_code=400, message="Tool `function` must be a dictionary" + ) raw_params = tool_function.get("parameters", {}) resolved_params = sanitize_oci_schema( @@ -308,7 +319,9 @@ def handle_generic_response( ): message.content = response_message.content[0].text if response_message.toolCalls: - message.tool_calls = adapt_tools_to_openai_standard(response_message.toolCalls) + message.tool_calls = adapt_tools_to_openai_standard( + response_message.toolCalls + ) oci_usage = completion_response.chatResponse.usage model_response.usage = Usage( # type: ignore[attr-defined] @@ -377,7 +390,9 @@ def handle_generic_stream_chunk(dict_chunk: dict) -> ModelResponseStream: delta=Delta( content=text, tool_calls=( - [tool.model_dump() for tool in tool_calls] if tool_calls else None + [tool.model_dump() for tool in tool_calls] + if tool_calls + else None ), provider_specific_fields=None, thinking_blocks=None, diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index 5197fba32d9..9ddda9ed699 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -35,6 +35,7 @@ from litellm.llms.custom_httpx.http_handler import ( version, ) from litellm.llms.oci.chat.cohere import ( + _extract_text_content, adapt_messages_to_cohere_standard, adapt_tool_definitions_to_cohere_standard, handle_cohere_response, @@ -43,10 +44,8 @@ from litellm.llms.oci.chat.cohere import ( from litellm.llms.oci.chat.generic import ( adapt_messages_to_generic_oci_standard, adapt_tool_definition_to_oci_standard, - adapt_tools_to_openai_standard, handle_generic_response, handle_generic_stream_chunk, - open_ai_to_generic_oci_role_map, ) from litellm.llms.oci.common_utils import ( OCIError, @@ -143,7 +142,7 @@ class OCIChatConfig(BaseConfig): # - tool_choice is unsupported # - stop sequences key is "stopSequences" not "stop" # - n (numGenerations) is GENERIC-only - self._openai_to_oci_cohere_param_map = { + self.openai_to_oci_cohere_param_map = { k: ("stopSequences" if k == "stop" else v) for k, v in self.openai_to_oci_generic_param_map.items() if k not in ("tool_choice", "max_retries", "n") @@ -151,7 +150,7 @@ class OCIChatConfig(BaseConfig): def get_supported_openai_params(self, model: str) -> List[str]: param_map = ( - self._openai_to_oci_cohere_param_map + self.openai_to_oci_cohere_param_map if get_vendor_from_model(model) == OCIVendors.COHERE else self.openai_to_oci_generic_param_map ) @@ -167,7 +166,7 @@ class OCIChatConfig(BaseConfig): adapted_params = {} vendor = get_vendor_from_model(model) param_map = ( - self._openai_to_oci_cohere_param_map + self.openai_to_oci_cohere_param_map if vendor == OCIVendors.COHERE else self.openai_to_oci_generic_param_map ) @@ -185,8 +184,6 @@ class OCIChatConfig(BaseConfig): adapted_params[key] = value continue adapted_params[alias] = value - if alias == "responseFormat": - adapted_params["response_format"] = value return adapted_params @@ -200,7 +197,7 @@ class OCIChatConfig(BaseConfig): model: Optional[str] = None, stream: Optional[bool] = None, fake_stream: Optional[bool] = None, - ) -> Tuple[dict, Optional[bytes]]: + ) -> Tuple[dict, bytes]: return sign_oci_request( headers=headers, optional_params=optional_params, @@ -231,7 +228,12 @@ class OCIChatConfig(BaseConfig): creds = resolve_oci_credentials(optional_params) missing = [ k - for k in ("oci_user", "oci_fingerprint", "oci_tenancy", "oci_compartment_id") + for k in ( + "oci_user", + "oci_fingerprint", + "oci_tenancy", + "oci_compartment_id", + ) if not creds.get(k) ] if missing or not (creds.get("oci_key") or creds.get("oci_key_file")): @@ -261,7 +263,7 @@ class OCIChatConfig(BaseConfig): def _get_optional_params(self, vendor: OCIVendors, optional_params: dict) -> Dict: param_map = ( - self._openai_to_oci_cohere_param_map + self.openai_to_oci_cohere_param_map if vendor == OCIVendors.COHERE else self.openai_to_oci_generic_param_map ) @@ -272,7 +274,11 @@ class OCIChatConfig(BaseConfig): selected_params[oci_key] = optional_params[openai_key] # type: ignore[index] for oci_value in param_map.values(): - if oci_value and oci_value in optional_params and oci_value not in selected_params: + if ( + oci_value + and oci_value in optional_params + and oci_value not in selected_params + ): selected_params[oci_value] = optional_params[oci_value] # type: ignore[index] if "tools" in selected_params: @@ -309,7 +315,9 @@ class OCIChatConfig(BaseConfig): schema_payload: Optional[Any] = None if "json_schema" in rf_payload: raw_schema = rf_payload.pop("json_schema") - schema_payload = dict(raw_schema) if isinstance(raw_schema, dict) else raw_schema + schema_payload = ( + dict(raw_schema) if isinstance(raw_schema, dict) else raw_schema + ) if schema_payload is not None: rf_payload["jsonSchema"] = schema_payload if vendor == OCIVendors.COHERE: @@ -320,35 +328,6 @@ class OCIChatConfig(BaseConfig): return selected_params - # ------------------------------------------------------------------ - # Delegating wrappers — keep instance-method call sites working - # while the real logic lives in cohere.py / generic.py. - # ------------------------------------------------------------------ - - def adapt_messages_to_cohere_standard( - self, messages: List[AllMessageValues] - ): # type: ignore[return] - return adapt_messages_to_cohere_standard(messages) - - def adapt_tool_definitions_to_cohere_standard( - self, tools: List[Dict[str, Any]] - ): # type: ignore[return] - return adapt_tool_definitions_to_cohere_standard(tools) - - def _extract_text_content(self, content: Any) -> str: - """Return plain-text for message content (string or content-part list).""" - if content is None: - return "" - if isinstance(content, str): - return content - if isinstance(content, list): - return "".join( - item.get("text", "") - for item in content - if isinstance(item, dict) and item.get("type") == "text" - ) - return str(content) - def transform_request( self, model: str, @@ -397,14 +376,14 @@ class OCIChatConfig(BaseConfig): preamble_override = None if system_messages: preamble = "\n".join( - self._extract_text_content(m["content"]) for m in system_messages + _extract_text_content(m["content"]) for m in system_messages ) if preamble: preamble_override = preamble chat_request = CohereChatRequest( apiFormat="COHERE", - message=self._extract_text_content(user_messages[-1]["content"]), + message=_extract_text_content(user_messages[-1]["content"]), chatHistory=adapt_messages_to_cohere_standard(messages), preambleOverride=preamble_override, **self._get_optional_params(OCIVendors.COHERE, optional_params), @@ -456,7 +435,9 @@ class OCIChatConfig(BaseConfig): vendor = get_vendor_from_model(model) if vendor == OCIVendors.COHERE: - model_response = handle_cohere_response(response_json, model, model_response) + model_response = handle_cohere_response( + response_json, model, model_response + ) else: model_response = handle_generic_response( response_json, model, model_response, raw_response @@ -477,7 +458,7 @@ class OCIChatConfig(BaseConfig): messages: list, client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + signed_json_body: bytes = b"", ) -> "OCIStreamWrapper": if "stream" in data: del data["stream"] @@ -488,7 +469,7 @@ class OCIChatConfig(BaseConfig): response = client.post( api_base, headers=headers, - data=signed_json_body if signed_json_body is not None else json.dumps(data), + data=signed_json_body or json.dumps(data), stream=True, logging_obj=logging_obj, timeout=STREAMING_TIMEOUT, @@ -525,7 +506,7 @@ class OCIChatConfig(BaseConfig): messages: list, client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, json_mode: Optional[bool] = None, - signed_json_body: Optional[bytes] = None, + signed_json_body: bytes = b"", ) -> "OCIStreamWrapper": if "stream" in data: del data["stream"] @@ -536,7 +517,7 @@ class OCIChatConfig(BaseConfig): response = await client.post( api_base, headers=headers, - data=signed_json_body if signed_json_body is not None else json.dumps(data), + data=signed_json_body or json.dumps(data), stream=True, logging_obj=logging_obj, timeout=STREAMING_TIMEOUT, @@ -586,18 +567,7 @@ class OCIStreamWrapper(CustomStreamWrapper): return handle_cohere_stream_chunk(dict_chunk) return handle_generic_stream_chunk(dict_chunk) - # Delegating shims so any code that calls these as instance methods keeps working. - def _handle_cohere_stream_chunk(self, dict_chunk: dict) -> ModelResponseStream: - return handle_cohere_stream_chunk(dict_chunk) - def _handle_generic_stream_chunk(self, dict_chunk: dict) -> ModelResponseStream: - return handle_generic_stream_chunk(dict_chunk) - - -# --------------------------------------------------------------------------- -# Backward-compatibility re-exports -# Keep all symbols that test files import directly from this module. -# --------------------------------------------------------------------------- __all__ = [ "OCIChatConfig", "OCIStreamWrapper", @@ -606,15 +576,4 @@ __all__ = [ "STREAMING_TIMEOUT", "get_vendor_from_model", "version", - # generic helpers (imported in tests) - "open_ai_to_generic_oci_role_map", - "adapt_messages_to_generic_oci_standard", - "adapt_messages_to_generic_oci_standard_content_message", - "adapt_messages_to_generic_oci_standard_tool_call", - "adapt_messages_to_generic_oci_standard_tool_response", - "adapt_tool_definition_to_oci_standard", - "adapt_tools_to_openai_standard", - # cohere helpers - "adapt_messages_to_cohere_standard", - "adapt_tool_definitions_to_cohere_standard", ] diff --git a/litellm/llms/oci/common_utils.py b/litellm/llms/oci/common_utils.py index 931fb344ddc..f6ab6de3ff4 100644 --- a/litellm/llms/oci/common_utils.py +++ b/litellm/llms/oci/common_utils.py @@ -4,13 +4,34 @@ import hashlib import json import os from dataclasses import dataclass -from typing import Any, Dict, List, Optional, Protocol, Tuple +from typing import Any, Dict, Optional, Protocol, Tuple from urllib.parse import urlparse import httpx from litellm.llms.base_llm.chat.transformation import BaseLLMException +try: + from cryptography.hazmat.primitives import hashes, serialization + from cryptography.hazmat.primitives.asymmetric import padding, rsa + + _CRYPTOGRAPHY_AVAILABLE = True +except ImportError: + _CRYPTOGRAPHY_AVAILABLE = False + +try: + from litellm._version import version as _litellm_version +except ImportError: + _litellm_version = "0.0.0" + + +def _require_cryptography() -> None: + if not _CRYPTOGRAPHY_AVAILABLE: + raise ImportError( + "cryptography package is required for OCI authentication. " + "Please install it with: pip install cryptography" + ) + class OCIError(BaseLLMException): def __init__( @@ -84,20 +105,12 @@ def build_signature_string( def load_private_key_from_str(key_str: str) -> Any: - try: - from cryptography.hazmat.primitives import serialization - from cryptography.hazmat.primitives.asymmetric import rsa - except ImportError as e: - raise ImportError( - "cryptography package is required for OCI authentication. " - "Please install it with: pip install cryptography" - ) from e - - key = serialization.load_pem_private_key( + _require_cryptography() + key = serialization.load_pem_private_key( # type: ignore[union-attr] key_str.encode("utf-8"), password=None, ) - if not isinstance(key, rsa.RSAPrivateKey): + if not isinstance(key, rsa.RSAPrivateKey): # type: ignore[union-attr] raise TypeError( "The provided private key is not an RSA key, which is required for OCI signing." ) @@ -220,7 +233,7 @@ def sign_with_manual_credentials( optional_params: dict, request_data: dict, api_base: str, -) -> Tuple[dict, None]: +) -> Tuple[dict, bytes]: """Sign a request using manually provided OCI credentials (user/fingerprint/tenancy/key).""" creds = resolve_oci_credentials(optional_params) oci_user = creds["oci_user"] @@ -252,7 +265,9 @@ def sign_with_manual_credentials( path = parsed.path or "/" host = parsed.netloc - date = datetime.datetime.now(datetime.timezone.utc).strftime("%a, %d %b %Y %H:%M:%S GMT") + date = datetime.datetime.now(datetime.timezone.utc).strftime( + "%a, %d %b %Y %H:%M:%S GMT" + ) content_type = headers.get("content-type", "application/json") content_length = str(len(body)) x_content_sha256 = sha256_base64(body) @@ -273,16 +288,11 @@ def sign_with_manual_credentials( "content-type", "x-content-sha256", ] - signing_string = build_signature_string(method, path, headers_to_sign, signed_header_names) + signing_string = build_signature_string( + method, path, headers_to_sign, signed_header_names + ) - try: - from cryptography.hazmat.primitives import hashes - from cryptography.hazmat.primitives.asymmetric import padding - except ImportError as e: - raise ImportError( - "cryptography package is required for OCI authentication. " - "Please install it with: pip install cryptography" - ) from e + _require_cryptography() # Resolve the private key — prefer inline PEM content over file path oci_key_content: Optional[str] = None @@ -300,9 +310,7 @@ def sign_with_manual_credentials( private_key = ( load_private_key_from_str(oci_key_content) if oci_key_content - else load_private_key_from_file(oci_key_file) - if oci_key_file - else None + else load_private_key_from_file(oci_key_file) if oci_key_file else None ) if private_key is None: @@ -313,8 +321,8 @@ def sign_with_manual_credentials( signature = private_key.sign( signing_string.encode("utf-8"), - padding.PKCS1v15(), - hashes.SHA256(), + padding.PKCS1v15(), # type: ignore[union-attr] + hashes.SHA256(), # type: ignore[union-attr] ) signature_b64 = base64.b64encode(signature).decode() @@ -337,7 +345,7 @@ def sign_with_manual_credentials( "x-content-sha256": x_content_sha256, } ) - return headers, None + return headers, body def sign_oci_request( @@ -349,7 +357,7 @@ def sign_oci_request( model: Optional[str] = None, stream: Optional[bool] = None, fake_stream: Optional[bool] = None, -) -> Tuple[dict, Optional[bytes]]: +) -> Tuple[dict, bytes]: """ Route to the appropriate OCI signing method based on what credentials are present. @@ -358,11 +366,13 @@ def sign_oci_request( also be supplied via OCI_* environment variables). Returns: - Tuple of (signed_headers, body_bytes_or_None) + Tuple of (signed_headers, signed_body_bytes) """ if optional_params.get("oci_signer") is not None: return sign_with_oci_signer(headers, optional_params, request_data, api_base) - return sign_with_manual_credentials(headers, optional_params, request_data, api_base) + return sign_with_manual_credentials( + headers, optional_params, request_data, api_base + ) def validate_oci_environment( @@ -377,10 +387,8 @@ def validate_oci_environment( supplied via environment variables are resolved at call time rather than at construction time. """ - from litellm.llms.custom_httpx.http_handler import version - headers.setdefault("content-type", "application/json") - headers.setdefault("user-agent", f"litellm/{version}") + headers.setdefault("user-agent", f"litellm/{_litellm_version}") return headers @@ -447,7 +455,8 @@ def resolve_oci_schema_anyof(obj: Any) -> Any: if isinstance(obj, dict): if "anyOf" in obj and "type" not in obj: non_null = [ - t for t in obj["anyOf"] + t + for t in obj["anyOf"] if not (isinstance(t, dict) and t.get("type") == "null") ] if non_null: @@ -504,7 +513,9 @@ def sanitize_oci_schema(schema: Any) -> Any: return sanitized -def enrich_cohere_param_description(description: str, param_schema: Dict[str, Any]) -> str: +def enrich_cohere_param_description( + description: str, param_schema: Dict[str, Any] +) -> str: """Embed schema constraints into a Cohere parameter description. ``CohereParameterDefinition`` only has ``type``, ``description``, and diff --git a/litellm/llms/oci/embed/transformation.py b/litellm/llms/oci/embed/transformation.py index c8f88b80474..85bcabdb621 100644 --- a/litellm/llms/oci/embed/transformation.py +++ b/litellm/llms/oci/embed/transformation.py @@ -206,7 +206,9 @@ class OCIEmbedConfig(BaseEmbeddingConfig): if serving_mode_type == "DEDICATED": endpoint_id = optional_params.get("oci_endpoint_id", model) - serving_mode = OCIServingMode(servingType="DEDICATED", endpointId=endpoint_id) + serving_mode = OCIServingMode( + servingType="DEDICATED", endpointId=endpoint_id + ) else: serving_mode = OCIServingMode(servingType="ON_DEMAND", modelId=model) diff --git a/litellm/types/llms/oci.py b/litellm/types/llms/oci.py index 1817cbb489e..57512ec7c0c 100644 --- a/litellm/types/llms/oci.py +++ b/litellm/types/llms/oci.py @@ -421,9 +421,13 @@ class OCIEmbedRequest(BaseModel): compartmentId: str servingMode: OCIServingMode inputs: List[str] - inputType: Optional[str] = None # SEARCH_DOCUMENT | SEARCH_QUERY | CLASSIFICATION | CLUSTERING | IMAGE + inputType: Optional[str] = ( + None # SEARCH_DOCUMENT | SEARCH_QUERY | CLASSIFICATION | CLUSTERING | IMAGE + ) truncate: Optional[str] = "END" # NONE | START | END - outputDimensions: Optional[int] = None # cohere.embed-v4.0+; valid: 256, 512, 1024, 1536 + outputDimensions: Optional[int] = ( + None # cohere.embed-v4.0+; valid: 256, 512, 1024, 1536 + ) class OCIEmbedUsage(BaseModel): @@ -442,5 +446,3 @@ class OCIEmbedResponse(BaseModel): inputTextTokenCounts: Optional[List[int]] = None # Some deployments may return a usage object instead usage: Optional[OCIEmbedUsage] = None - - diff --git a/tests/llm_translation/test_oci_integration.py b/tests/llm_translation/test_oci_integration.py index bc3db7dc2e8..36510f53747 100644 --- a/tests/llm_translation/test_oci_integration.py +++ b/tests/llm_translation/test_oci_integration.py @@ -18,9 +18,10 @@ Run only these tests: pytest tests/llm_translation/test_oci_integration.py -v """ +import math import os import sys -from typing import Generator +from typing import NamedTuple, Optional import pytest @@ -67,244 +68,243 @@ def oci_params(oci_signer) -> dict: # --------------------------------------------------------------------------- -# Helpers +# Model registry +# +# Each entry drives the runtime pivot inside OCI's own transformation layer — +# the tests themselves are format-agnostic. Per-model quirks are captured in +# the config fields below rather than in separate test classes. # --------------------------------------------------------------------------- -def _chat(model: str, message: str, params: dict, max_tokens: int = 64) -> str: - """Run a single-turn completion and return the text content.""" + +class _M(NamedTuple): + """Per-model test configuration.""" + + model: str + max_tokens: int + # Reasoning models (Gemini 2.5, Grok mini) may return None content when the + # reasoning budget is exhausted before the answer token budget starts. + reasoning: bool = False + # tool_choice value to send; None means omit the parameter entirely. + tool_choice: Optional[str] = "auto" + # Whether to include the model in tool-use parametrize list. + supports_tool_use: bool = True + + +# All chat models under test. +CHAT_MODELS = [ + pytest.param(_M("meta.llama-3.3-70b-instruct", 64), id="meta"), + pytest.param(_M("google.gemini-2.5-flash", 200, reasoning=True), id="google"), + pytest.param(_M("xai.grok-3-mini", 100, reasoning=True), id="xai"), + pytest.param(_M("cohere.command-latest", 64, tool_choice=None), id="cohere"), +] + +# Subset of models that reliably support tool use in OCI. +# xAI Grok mini is omitted — OCI does not expose tool-use for it yet. +TOOL_USE_MODELS = [ + pytest.param(_M("meta.llama-3.3-70b-instruct", 100), id="meta"), + pytest.param(_M("cohere.command-latest", 200, tool_choice=None), id="cohere"), + pytest.param(_M("google.gemini-2.5-flash", 200, reasoning=True), id="google"), +] + +# Simple weather tool used by all tool-use tests. +_WEATHER_TOOL = { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get the current weather for a city.", + "parameters": { + "type": "object", + "properties": {"city": {"type": "string", "description": "The city name."}}, + "required": ["city"], + }, + }, +} + + +# --------------------------------------------------------------------------- +# Sync chat tests — model list drives the pivot, not separate test classes +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("m", CHAT_MODELS) +def test_basic_completion(m: _M, oci_params): import litellm resp = litellm.completion( - model=f"oci/{model}", - messages=[{"role": "user", "content": message}], - max_tokens=max_tokens, - **params, + model=f"oci/{m.model}", + messages=[{"role": "user", "content": "Reply with only the word: pong"}], + max_tokens=m.max_tokens, + **oci_params, ) - # Reasoning models may return None content when all budget is used by reasoning - return resp.choices[0].message.content or "" + assert resp.choices[0].finish_reason is not None + assert resp.usage.prompt_tokens > 0 + if not m.reasoning: + assert resp.choices[0].message.content is not None + + +@pytest.mark.parametrize("m", CHAT_MODELS) +def test_usage_populated(m: _M, oci_params): + import litellm + + resp = litellm.completion( + model=f"oci/{m.model}", + messages=[{"role": "user", "content": "What is 2+2?"}], + max_tokens=m.max_tokens, + **oci_params, + ) + assert resp.usage.prompt_tokens > 0 + assert resp.usage.total_tokens >= resp.usage.prompt_tokens + + +@pytest.mark.parametrize("m", CHAT_MODELS) +def test_system_message(m: _M, oci_params): + import litellm + + resp = litellm.completion( + model=f"oci/{m.model}", + messages=[ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "Say hello."}, + ], + max_tokens=m.max_tokens, + **oci_params, + ) + assert resp.choices[0].finish_reason is not None + + +@pytest.mark.parametrize("m", CHAT_MODELS) +def test_streaming(m: _M, oci_params): + import litellm + + chunks = list( + litellm.completion( + model=f"oci/{m.model}", + messages=[{"role": "user", "content": "Count to 3."}], + max_tokens=m.max_tokens, + stream=True, + **oci_params, + ) + ) + assert len(chunks) > 0 + # Reasoning models may stream only reasoning tokens and return empty content. + if not m.reasoning: + content = "".join(c.choices[0].delta.content or "" for c in chunks if c.choices) + assert len(content) > 0 + + +@pytest.mark.parametrize("m", CHAT_MODELS) +def test_multi_turn(m: _M, oci_params): + import litellm + + resp = litellm.completion( + model=f"oci/{m.model}", + messages=[ + {"role": "user", "content": "My name is Alice."}, + {"role": "assistant", "content": "Nice to meet you, Alice!"}, + {"role": "user", "content": "What is my name?"}, + ], + max_tokens=m.max_tokens, + **oci_params, + ) + # Reasoning models may have None content; skip text assertion for them. + content = resp.choices[0].message.content or "" + if not m.reasoning: + assert "Alice" in content # --------------------------------------------------------------------------- -# Chat tests — one per vendor family +# Async chat tests # --------------------------------------------------------------------------- -class TestOCIChatMeta: - """Meta Llama models (GENERIC apiFormat).""" +@pytest.mark.asyncio +@pytest.mark.parametrize("m", CHAT_MODELS) +async def test_async_completion(m: _M, oci_params): + import litellm - MODEL = "meta.llama-3.3-70b-instruct" - - def test_basic_completion(self, oci_params): - import litellm - - resp = litellm.completion( - model=f"oci/{self.MODEL}", - messages=[{"role": "user", "content": "Reply with only the word: pong"}], - max_tokens=10, - **oci_params, - ) - assert resp.choices[0].message.content is not None - assert resp.choices[0].finish_reason is not None - assert resp.usage.prompt_tokens > 0 - - def test_usage_populated(self, oci_params): - import litellm - - resp = litellm.completion( - model=f"oci/{self.MODEL}", - messages=[{"role": "user", "content": "Say hi."}], - max_tokens=20, - **oci_params, - ) - assert resp.usage.prompt_tokens > 0 - assert resp.usage.total_tokens >= resp.usage.prompt_tokens - - def test_system_message(self, oci_params): - import litellm - - resp = litellm.completion( - model=f"oci/{self.MODEL}", - messages=[ - {"role": "system", "content": "You only reply in pirate speak."}, - {"role": "user", "content": "Hello!"}, - ], - max_tokens=30, - **oci_params, - ) + resp = await litellm.acompletion( + model=f"oci/{m.model}", + messages=[{"role": "user", "content": "Reply with only the word: pong"}], + max_tokens=m.max_tokens, + **oci_params, + ) + assert resp.choices[0].finish_reason is not None + assert resp.usage.total_tokens > 0 + if not m.reasoning: assert resp.choices[0].message.content is not None - def test_streaming(self, oci_params): - import litellm - chunks = list( - litellm.completion( - model=f"oci/{self.MODEL}", - messages=[{"role": "user", "content": "Count to 3."}], - max_tokens=30, - stream=True, - **oci_params, - ) - ) - assert len(chunks) > 0 - content = "".join( - c.choices[0].delta.content or "" for c in chunks if c.choices - ) +@pytest.mark.asyncio +@pytest.mark.parametrize("m", CHAT_MODELS) +async def test_async_streaming(m: _M, oci_params): + import litellm + + chunks = [] + async for chunk in await litellm.acompletion( + model=f"oci/{m.model}", + messages=[{"role": "user", "content": "Count to 3."}], + max_tokens=m.max_tokens, + stream=True, + **oci_params, + ): + chunks.append(chunk) + + assert len(chunks) > 0 + if not m.reasoning: + content = "".join(c.choices[0].delta.content or "" for c in chunks if c.choices) assert len(content) > 0 - def test_multi_turn(self, oci_params): - import litellm - resp = litellm.completion( - model=f"oci/{self.MODEL}", - messages=[ - {"role": "user", "content": "My name is Alice."}, - {"role": "assistant", "content": "Nice to meet you, Alice!"}, - {"role": "user", "content": "What is my name?"}, - ], - max_tokens=30, - **oci_params, - ) - assert "Alice" in (resp.choices[0].message.content or "") +# --------------------------------------------------------------------------- +# Tool-use tests +# --------------------------------------------------------------------------- -class TestOCIChatGoogle: - """Google Gemini models (GENERIC apiFormat).""" - - MODEL = "google.gemini-2.5-flash" - - def test_basic_completion(self, oci_params): - import litellm - - resp = litellm.completion( - model=f"oci/{self.MODEL}", - messages=[{"role": "user", "content": "Reply with only the word: pong"}], - max_tokens=200, - **oci_params, - ) - # Gemini 2.5 Flash is a reasoning model — content may be empty if reasoning - # consumed the budget, but no exception should be raised. - assert resp.choices[0].finish_reason is not None - assert resp.usage.prompt_tokens > 0 - - def test_usage_populated(self, oci_params): - import litellm - - resp = litellm.completion( - model=f"oci/{self.MODEL}", - messages=[{"role": "user", "content": "What is 2+2?"}], - max_tokens=200, - **oci_params, - ) - assert resp.usage.total_tokens > 0 - - def test_streaming(self, oci_params): - import litellm - - chunks = list( - litellm.completion( - model=f"oci/{self.MODEL}", - messages=[{"role": "user", "content": "Say the word hello."}], - max_tokens=200, - stream=True, - **oci_params, - ) - ) - assert len(chunks) > 0 +def _assert_tool_call(resp, expected_tool: str = "get_weather"): + """Assert the response contains the expected tool call (or a plain stop).""" + choice = resp.choices[0] + assert choice.finish_reason in ("tool_calls", "stop") + if choice.finish_reason == "tool_calls": + assert choice.message.tool_calls is not None + assert len(choice.message.tool_calls) > 0 + assert choice.message.tool_calls[0].function.name == expected_tool -class TestOCIChatXAI: - """xAI Grok models (GENERIC apiFormat).""" +@pytest.mark.parametrize("m", TOOL_USE_MODELS) +def test_tool_use(m: _M, oci_params): + import litellm - MODEL = "xai.grok-3-mini" + call_kwargs = dict( + model=f"oci/{m.model}", + messages=[{"role": "user", "content": "What is the weather in Paris?"}], + tools=[_WEATHER_TOOL], + max_tokens=m.max_tokens, + **oci_params, + ) + if m.tool_choice is not None: + call_kwargs["tool_choice"] = m.tool_choice - def test_basic_completion(self, oci_params): - import litellm - - resp = litellm.completion( - model=f"oci/{self.MODEL}", - messages=[{"role": "user", "content": "Reply with only the word: pong"}], - max_tokens=50, - **oci_params, - ) - assert resp.choices[0].message.content is not None - assert resp.usage.total_tokens > 0 - - def test_streaming(self, oci_params): - import litellm - - chunks = list( - litellm.completion( - model=f"oci/{self.MODEL}", - messages=[{"role": "user", "content": "Count to 3."}], - max_tokens=50, - stream=True, - **oci_params, - ) - ) - assert len(chunks) > 0 - - def test_usage_has_reasoning_tokens(self, oci_params): - """Grok mini exposes reasoning token breakdown in usage.""" - import litellm - - resp = litellm.completion( - model=f"oci/{self.MODEL}", - messages=[{"role": "user", "content": "What is 5*7?"}], - max_tokens=100, - **oci_params, - ) - # totalTokens >= completionTokens + promptTokens (reasoning may add overhead) - assert resp.usage.total_tokens >= resp.usage.prompt_tokens + resp = litellm.completion(**call_kwargs) + _assert_tool_call(resp) -class TestOCIChatCohere: - """Cohere Command models (COHERE apiFormat).""" +@pytest.mark.asyncio +@pytest.mark.parametrize("m", TOOL_USE_MODELS) +async def test_async_tool_use(m: _M, oci_params): + import litellm - MODEL = "cohere.command-latest" + call_kwargs = dict( + model=f"oci/{m.model}", + messages=[{"role": "user", "content": "What is the weather in Berlin?"}], + tools=[_WEATHER_TOOL], + max_tokens=m.max_tokens, + **oci_params, + ) + if m.tool_choice is not None: + call_kwargs["tool_choice"] = m.tool_choice - def test_basic_completion(self, oci_params): - import litellm - - resp = litellm.completion( - model=f"oci/{self.MODEL}", - messages=[{"role": "user", "content": "Reply with only the word: pong"}], - max_tokens=20, - **oci_params, - ) - assert resp.choices[0].message.content is not None - assert resp.usage.prompt_tokens > 0 - - def test_streaming(self, oci_params): - import litellm - - chunks = list( - litellm.completion( - model=f"oci/{self.MODEL}", - messages=[{"role": "user", "content": "Count to 3."}], - max_tokens=30, - stream=True, - **oci_params, - ) - ) - assert len(chunks) > 0 - content = "".join( - c.choices[0].delta.content or "" for c in chunks if c.choices - ) - assert len(content) > 0 - - def test_system_message(self, oci_params): - import litellm - - resp = litellm.completion( - model=f"oci/{self.MODEL}", - messages=[ - {"role": "system", "content": "Always end your response with 'cheers'."}, - {"role": "user", "content": "Say hello."}, - ], - max_tokens=40, - **oci_params, - ) - assert resp.choices[0].message.content is not None + resp = await litellm.acompletion(**call_kwargs) + _assert_tool_call(resp) # --------------------------------------------------------------------------- @@ -372,7 +372,6 @@ class TestOCIEmbeddings: def test_semantic_similarity(self, oci_params): """Semantically similar texts should have higher cosine similarity.""" import litellm - import math resp = litellm.embedding( model="oci/cohere.embed-english-v3.0", @@ -387,19 +386,18 @@ class TestOCIEmbeddings: def cosine(a, b): dot = sum(x * y for x, y in zip(a, b)) - mag_a = math.sqrt(sum(x ** 2 for x in a)) - mag_b = math.sqrt(sum(x ** 2 for x in b)) + mag_a = math.sqrt(sum(x**2 for x in a)) + mag_b = math.sqrt(sum(x**2 for x in b)) return dot / (mag_a * mag_b) cat1 = resp.data[0]["embedding"] cat2 = resp.data[1]["embedding"] stock = resp.data[2]["embedding"] - sim_cats = cosine(cat1, cat2) sim_diff = cosine(cat1, stock) - assert sim_cats > sim_diff, ( - f"Expected similar sentences to score higher ({sim_cats:.3f} vs {sim_diff:.3f})" - ) + assert ( + sim_cats > sim_diff + ), f"Expected similar sentences to score higher ({sim_cats:.3f} vs {sim_diff:.3f})" def test_embed_v4(self, oci_params): import litellm @@ -426,6 +424,54 @@ class TestOCIEmbeddings: assert resp.usage.total_tokens == resp.usage.prompt_tokens +# --------------------------------------------------------------------------- +# Async embedding tests +# --------------------------------------------------------------------------- + + +class TestOCIAsyncEmbeddings: + + @pytest.mark.asyncio + async def test_async_embedding_basic(self, oci_params): + import litellm + + resp = await litellm.aembedding( + model="oci/cohere.embed-english-v3.0", + input=["Hello world"], + input_type="SEARCH_DOCUMENT", + **oci_params, + ) + assert len(resp.data) == 1 + assert len(resp.data[0]["embedding"]) == 1024 + assert resp.usage.prompt_tokens > 0 + + @pytest.mark.asyncio + async def test_async_embedding_batch(self, oci_params): + import litellm + + texts = ["The quick brown fox", "jumps over the lazy dog"] + resp = await litellm.aembedding( + model="oci/cohere.embed-english-v3.0", + input=texts, + input_type="SEARCH_DOCUMENT", + **oci_params, + ) + assert len(resp.data) == 2 + assert all(len(item["embedding"]) == 1024 for item in resp.data) + + @pytest.mark.asyncio + async def test_async_embedding_multilingual(self, oci_params): + import litellm + + resp = await litellm.aembedding( + model="oci/cohere.embed-multilingual-v3.0", + input=["Bonjour le monde"], + input_type="SEARCH_DOCUMENT", + **oci_params, + ) + assert len(resp.data[0]["embedding"]) == 1024 + + # --------------------------------------------------------------------------- # Env-var credential path # --------------------------------------------------------------------------- @@ -482,261 +528,3 @@ class TestOCIEnvVarCredentials: input_type="SEARCH_DOCUMENT", ) assert len(resp.data[0]["embedding"]) == 1024 - - -# --------------------------------------------------------------------------- -# Async chat tests -# --------------------------------------------------------------------------- - - -class TestOCIAsyncChat: - """Verify async completion and streaming work for each vendor family.""" - - @pytest.mark.asyncio - async def test_async_completion_meta(self, oci_params): - import litellm - - resp = await litellm.acompletion( - model="oci/meta.llama-3.3-70b-instruct", - messages=[{"role": "user", "content": "Reply with only the word: pong"}], - max_tokens=10, - **oci_params, - ) - assert resp.choices[0].message.content is not None - assert resp.usage.prompt_tokens > 0 - - @pytest.mark.asyncio - async def test_async_completion_google(self, oci_params): - import litellm - - resp = await litellm.acompletion( - model="oci/google.gemini-2.5-flash", - messages=[{"role": "user", "content": "Reply with only the word: pong"}], - max_tokens=200, - **oci_params, - ) - assert resp.choices[0].finish_reason is not None - assert resp.usage.total_tokens > 0 - - @pytest.mark.asyncio - async def test_async_completion_xai(self, oci_params): - import litellm - - resp = await litellm.acompletion( - model="oci/xai.grok-3-mini", - messages=[{"role": "user", "content": "Reply with only the word: pong"}], - max_tokens=50, - **oci_params, - ) - assert resp.choices[0].message.content is not None - - @pytest.mark.asyncio - async def test_async_completion_cohere(self, oci_params): - import litellm - - resp = await litellm.acompletion( - model="oci/cohere.command-latest", - messages=[{"role": "user", "content": "Reply with only the word: pong"}], - max_tokens=20, - **oci_params, - ) - assert resp.choices[0].message.content is not None - - @pytest.mark.asyncio - async def test_async_streaming_meta(self, oci_params): - import litellm - - chunks = [] - async for chunk in await litellm.acompletion( - model="oci/meta.llama-3.3-70b-instruct", - messages=[{"role": "user", "content": "Count to 3."}], - max_tokens=30, - stream=True, - **oci_params, - ): - chunks.append(chunk) - - assert len(chunks) > 0 - content = "".join( - c.choices[0].delta.content or "" for c in chunks if c.choices - ) - assert len(content) > 0 - - @pytest.mark.asyncio - async def test_async_streaming_google(self, oci_params): - import litellm - - chunks = [] - async for chunk in await litellm.acompletion( - model="oci/google.gemini-2.5-flash", - messages=[{"role": "user", "content": "Say the word hello."}], - max_tokens=200, - stream=True, - **oci_params, - ): - chunks.append(chunk) - - assert len(chunks) > 0 - - @pytest.mark.asyncio - async def test_async_streaming_xai(self, oci_params): - import litellm - - chunks = [] - async for chunk in await litellm.acompletion( - model="oci/xai.grok-3-mini", - messages=[{"role": "user", "content": "Count to 3."}], - max_tokens=50, - stream=True, - **oci_params, - ): - chunks.append(chunk) - - assert len(chunks) > 0 - - -# --------------------------------------------------------------------------- -# Async embedding tests -# --------------------------------------------------------------------------- - - -class TestOCIAsyncEmbeddings: - - @pytest.mark.asyncio - async def test_async_embedding_basic(self, oci_params): - import litellm - - resp = await litellm.aembedding( - model="oci/cohere.embed-english-v3.0", - input=["Hello world"], - input_type="SEARCH_DOCUMENT", - **oci_params, - ) - assert len(resp.data) == 1 - assert len(resp.data[0]["embedding"]) == 1024 - assert resp.usage.prompt_tokens > 0 - - @pytest.mark.asyncio - async def test_async_embedding_batch(self, oci_params): - import litellm - - texts = ["The quick brown fox", "jumps over the lazy dog"] - resp = await litellm.aembedding( - model="oci/cohere.embed-english-v3.0", - input=texts, - input_type="SEARCH_DOCUMENT", - **oci_params, - ) - assert len(resp.data) == 2 - assert all(len(item["embedding"]) == 1024 for item in resp.data) - - @pytest.mark.asyncio - async def test_async_embedding_multilingual(self, oci_params): - import litellm - - resp = await litellm.aembedding( - model="oci/cohere.embed-multilingual-v3.0", - input=["Bonjour le monde"], - input_type="SEARCH_DOCUMENT", - **oci_params, - ) - assert len(resp.data[0]["embedding"]) == 1024 - - -# --------------------------------------------------------------------------- -# Tool use / function calling tests -# --------------------------------------------------------------------------- - - -class TestOCIToolUse: - """Verify tool use (function calling) works for models that support it.""" - - # Simple weather tool definition - WEATHER_TOOL = { - "type": "function", - "function": { - "name": "get_weather", - "description": "Get the current weather for a city.", - "parameters": { - "type": "object", - "properties": { - "city": { - "type": "string", - "description": "The city name.", - } - }, - "required": ["city"], - }, - }, - } - - def _assert_tool_call(self, resp, expected_tool: str = "get_weather"): - """Assert the response contains a tool call.""" - choice = resp.choices[0] - assert choice.finish_reason in ("tool_calls", "stop") - if choice.finish_reason == "tool_calls": - assert choice.message.tool_calls is not None - assert len(choice.message.tool_calls) > 0 - assert choice.message.tool_calls[0].function.name == expected_tool - # stop finish_reason can happen if model answers without calling the tool — - # acceptable behaviour, not a bug. - - def test_tool_use_meta(self, oci_params): - import litellm - - resp = litellm.completion( - model="oci/meta.llama-3.3-70b-instruct", - messages=[ - {"role": "user", "content": "What is the weather in Paris?"} - ], - tools=[self.WEATHER_TOOL], - tool_choice="auto", - max_tokens=100, - **oci_params, - ) - self._assert_tool_call(resp) - - def test_tool_use_cohere(self, oci_params): - import litellm - - resp = litellm.completion( - model="oci/cohere.command-latest", - messages=[ - {"role": "user", "content": "What is the weather in Tokyo?"} - ], - tools=[self.WEATHER_TOOL], - max_tokens=200, - **oci_params, - ) - self._assert_tool_call(resp) - - def test_tool_use_google(self, oci_params): - import litellm - - resp = litellm.completion( - model="oci/google.gemini-2.5-flash", - messages=[ - {"role": "user", "content": "What is the weather in London?"} - ], - tools=[self.WEATHER_TOOL], - tool_choice="auto", - max_tokens=200, - **oci_params, - ) - self._assert_tool_call(resp) - - @pytest.mark.asyncio - async def test_async_tool_use_meta(self, oci_params): - import litellm - - resp = await litellm.acompletion( - model="oci/meta.llama-3.3-70b-instruct", - messages=[ - {"role": "user", "content": "What is the weather in Berlin?"} - ], - tools=[self.WEATHER_TOOL], - tool_choice="auto", - max_tokens=100, - **oci_params, - ) - self._assert_tool_call(resp) diff --git a/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py b/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py index 035fa7d5846..3d566dd7c64 100644 --- a/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py +++ b/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py @@ -789,9 +789,11 @@ class TestOCISplitChunks: async def _run_async_split(self, raw_chunks): """Invoke the async split_chunks logic directly.""" results = [] + async def _gen(): for c in raw_chunks: yield c + async for item in _gen(): for chunk in item.split("\n\n"): stripped = chunk.strip() @@ -877,9 +879,8 @@ class TestOCIProviderEmbeddingConfig: """ import inspect from litellm.utils import ProviderConfigManager - source = inspect.getsource( - ProviderConfigManager.get_provider_embedding_config - ) + + source = inspect.getsource(ProviderConfigManager.get_provider_embedding_config) oci_count = source.count("LlmProviders.OCI") assert oci_count == 1, ( f"Expected exactly 1 OCI branch in get_provider_embedding_config, found {oci_count}. " @@ -930,10 +931,16 @@ class TestOCICohereParamMapping: model="cohere.command-latest", drop_params=False, ) - for injected in ("maxTokens", "temperature", "topK", "topP", "frequencyPenalty"): - assert injected not in result, ( - f"'{injected}' should not be injected when user did not provide it" - ) + for injected in ( + "maxTokens", + "temperature", + "topK", + "topP", + "frequencyPenalty", + ): + assert ( + injected not in result + ), f"'{injected}' should not be injected when user did not provide it" def test_cohere_explicit_params_still_passed(self): """User-provided Cohere params must still be forwarded correctly.""" @@ -993,9 +1000,9 @@ class TestOCIStreamingSignedBody: signed_json_body=signed_bytes, ) - assert posted_data["data"] == signed_bytes, ( - "Streaming must use signed_json_body, not re-serialize data" - ) + assert ( + posted_data["data"] == signed_bytes + ), "Streaming must use signed_json_body, not re-serialize data" def test_get_custom_stream_wrapper_fallback_without_signed_body(self, monkeypatch): """When signed_json_body is None, fall back to json.dumps(data).""" @@ -1032,6 +1039,6 @@ class TestOCIStreamingSignedBody: signed_json_body=None, ) - assert posted_data["data"] == json.dumps(payload), ( - "Without signed_json_body, must fall back to json.dumps(data)" - ) + assert posted_data["data"] == json.dumps( + payload + ), "Without signed_json_body, must fall back to json.dumps(data)" diff --git a/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py b/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py index 04e2f01be41..f945923ee63 100644 --- a/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py +++ b/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py @@ -5,10 +5,14 @@ import json from unittest.mock import patch, MagicMock from litellm import ModelResponse +from litellm.llms.oci.chat.cohere import ( + adapt_messages_to_cohere_standard, + adapt_tool_definitions_to_cohere_standard, +) from litellm.llms.oci.chat.transformation import ( OCIChatConfig, - get_vendor_from_model, OCIStreamWrapper, + get_vendor_from_model, ) from litellm.types.llms.oci import OCIVendors @@ -75,7 +79,7 @@ class TestOCICohereToolCalls: ] # Transform tools - cohere_tools = config.adapt_tool_definitions_to_cohere_standard(openai_tools) + cohere_tools = adapt_tool_definitions_to_cohere_standard(openai_tools) # Verify transformation assert len(cohere_tools) == 2 @@ -95,7 +99,10 @@ class TestOCICohereToolCalls: # Check unit parameter unit_param = weather_tool.parameterDefinitions["unit"] - assert unit_param.description == "Temperature unit (celsius or fahrenheit). Allowed values: ['celsius', 'fahrenheit']" + assert ( + unit_param.description + == "Temperature unit (celsius or fahrenheit). Allowed values: ['celsius', 'fahrenheit']" + ) assert unit_param.type == "str" assert unit_param.isRequired == False @@ -156,7 +163,7 @@ class TestOCICohereToolCalls: assert chat_request["apiFormat"] == "COHERE" assert chat_request["message"] == "What's the weather like in Tokyo?" assert chat_request["chatHistory"] == [] - + # Verify tools are transformed correctly assert "tools" in chat_request assert len(chat_request["tools"]) == 1 @@ -317,7 +324,7 @@ class TestOCICohereToolCalls: }, ] - chat_history = config.adapt_messages_to_cohere_standard(messages) + chat_history = adapt_messages_to_cohere_standard(messages) # First message is the user message assert chat_history[0].role == "USER" @@ -358,7 +365,7 @@ class TestOCICohereToolCalls: }, ] - chat_history = config.adapt_messages_to_cohere_standard(messages) + chat_history = adapt_messages_to_cohere_standard(messages) # Verify chat history structure (excludes last message) assert len(chat_history) == 2 @@ -524,7 +531,7 @@ class TestOCICohereToolCalls: ] # The function should handle missing function key gracefully - cohere_tools = config.adapt_tool_definitions_to_cohere_standard(invalid_tools) + cohere_tools = adapt_tool_definitions_to_cohere_standard(invalid_tools) # Should create a tool with empty name and description assert len(cohere_tools) == 1 @@ -679,7 +686,7 @@ class TestOCICoherePreambleOverride: {"role": "user", "content": "Second question"}, ] - chat_history = config.adapt_messages_to_cohere_standard(messages) + chat_history = adapt_messages_to_cohere_standard(messages) # Should contain user and assistant only, no system # Note: adapt_messages_to_cohere_standard excludes the last message @@ -706,7 +713,7 @@ class TestOCICohereStreaming: stream_wrapper = self._create_stream_wrapper() # chunk_creator is the public dispatch entry point - assert hasattr(stream_wrapper, 'chunk_creator') + assert hasattr(stream_wrapper, "chunk_creator") assert callable(stream_wrapper.chunk_creator) def test_cohere_streaming_chunk_parsing(self): diff --git a/tests/test_litellm/llms/oci/chat/test_oci_streaming_tool_calls.py b/tests/test_litellm/llms/oci/chat/test_oci_streaming_tool_calls.py index de79793975e..fe8d6991a93 100644 --- a/tests/test_litellm/llms/oci/chat/test_oci_streaming_tool_calls.py +++ b/tests/test_litellm/llms/oci/chat/test_oci_streaming_tool_calls.py @@ -11,13 +11,10 @@ Error: ValidationError: 1 validation error for OCIStreamChunk message.toolCalls. import os import sys -import pytest -from unittest.mock import MagicMock -# Adds the parent directory to the system path sys.path.insert(0, os.path.abspath("../../../../..")) -from litellm.llms.oci.chat.transformation import OCIStreamWrapper +from litellm.llms.oci.chat.generic import handle_generic_stream_chunk from litellm.types.utils import ModelResponseStream @@ -26,12 +23,9 @@ class TestOCIStreamingToolCalls: def test_stream_chunk_with_missing_arguments_field(self): """ - Test that streaming chunks with tool calls missing 'arguments' field are handled. - OCI API can return tool calls in early chunks without the 'arguments' field, which should be filled with an empty string to satisfy Pydantic validation. """ - # Mock streaming chunk with tool call missing 'arguments' field chunk_data = { "index": 0, "finishReason": None, @@ -43,21 +37,13 @@ class TestOCIStreamingToolCalls: "type": "FUNCTION", "id": "call_abc123", "name": "get_weather", - # Note: 'arguments' field is missing + # 'arguments' field is missing } ], }, } - wrapper = OCIStreamWrapper( - completion_stream=iter([]), - model="meta.llama-3.1-405b-instruct", - custom_llm_provider="oci", - logging_obj=MagicMock(), - ) - - # This should not raise a ValidationError - result = wrapper._handle_generic_stream_chunk(chunk_data) + result = handle_generic_stream_chunk(chunk_data) assert isinstance(result, ModelResponseStream) assert len(result.choices) == 1 @@ -66,9 +52,7 @@ class TestOCIStreamingToolCalls: assert result.choices[0].delta.tool_calls[0]["function"]["arguments"] == "" def test_stream_chunk_with_missing_id_field(self): - """ - Test that streaming chunks with tool calls missing 'id' field are handled. - """ + """Missing 'id' gets a generated call_* id.""" chunk_data = { "index": 0, "finishReason": None, @@ -80,30 +64,20 @@ class TestOCIStreamingToolCalls: "type": "FUNCTION", "name": "get_weather", "arguments": '{"location": "San Francisco"}', - # Note: 'id' field is missing + # 'id' field is missing } ], }, } - wrapper = OCIStreamWrapper( - completion_stream=iter([]), - model="meta.llama-3.1-405b-instruct", - custom_llm_provider="oci", - logging_obj=MagicMock(), - ) - - result = wrapper._handle_generic_stream_chunk(chunk_data) + result = handle_generic_stream_chunk(chunk_data) assert isinstance(result, ModelResponseStream) assert result.choices[0].delta.tool_calls is not None - # Missing id is filled with a generated call_* id to avoid empty/null ids assert result.choices[0].delta.tool_calls[0]["id"].startswith("call_") def test_stream_chunk_with_missing_name_field(self): - """ - Test that streaming chunks with tool calls missing 'name' field are handled. - """ + """Missing 'name' defaults to empty string.""" chunk_data = { "index": 0, "finishReason": None, @@ -115,29 +89,20 @@ class TestOCIStreamingToolCalls: "type": "FUNCTION", "id": "call_abc123", "arguments": '{"location": "San Francisco"}', - # Note: 'name' field is missing + # 'name' field is missing } ], }, } - wrapper = OCIStreamWrapper( - completion_stream=iter([]), - model="meta.llama-3.1-405b-instruct", - custom_llm_provider="oci", - logging_obj=MagicMock(), - ) - - result = wrapper._handle_generic_stream_chunk(chunk_data) + result = handle_generic_stream_chunk(chunk_data) assert isinstance(result, ModelResponseStream) assert result.choices[0].delta.tool_calls is not None assert result.choices[0].delta.tool_calls[0]["function"]["name"] == "" def test_stream_chunk_with_all_missing_fields(self): - """ - Test that streaming chunks with tool calls missing all optional fields are handled. - """ + """All optional fields missing — all default gracefully.""" chunk_data = { "index": 0, "finishReason": None, @@ -147,32 +112,22 @@ class TestOCIStreamingToolCalls: "toolCalls": [ { "type": "FUNCTION" - # All fields missing: id, name, arguments + # id, name, arguments all missing } ], }, } - wrapper = OCIStreamWrapper( - completion_stream=iter([]), - model="meta.llama-3.1-405b-instruct", - custom_llm_provider="oci", - logging_obj=MagicMock(), - ) - - result = wrapper._handle_generic_stream_chunk(chunk_data) + result = handle_generic_stream_chunk(chunk_data) assert isinstance(result, ModelResponseStream) assert result.choices[0].delta.tool_calls is not None - # Missing id is filled with a generated call_* id to avoid empty/null ids assert result.choices[0].delta.tool_calls[0]["id"].startswith("call_") assert result.choices[0].delta.tool_calls[0]["function"]["name"] == "" assert result.choices[0].delta.tool_calls[0]["function"]["arguments"] == "" def test_stream_chunk_with_complete_tool_call(self): - """ - Test that streaming chunks with complete tool calls still work correctly. - """ + """Fully-populated tool call passes through unchanged.""" chunk_data = { "index": 0, "finishReason": None, @@ -190,14 +145,7 @@ class TestOCIStreamingToolCalls: }, } - wrapper = OCIStreamWrapper( - completion_stream=iter([]), - model="meta.llama-3.1-405b-instruct", - custom_llm_provider="oci", - logging_obj=MagicMock(), - ) - - result = wrapper._handle_generic_stream_chunk(chunk_data) + result = handle_generic_stream_chunk(chunk_data) assert isinstance(result, ModelResponseStream) assert result.choices[0].delta.tool_calls is not None @@ -212,9 +160,7 @@ class TestOCIStreamingToolCalls: ) def test_stream_chunk_with_multiple_tool_calls_missing_fields(self): - """ - Test that streaming chunks with multiple tool calls, some with missing fields, are handled. - """ + """Multiple tool calls with a mix of complete and incomplete entries.""" chunk_data = { "index": 0, "finishReason": None, @@ -222,50 +168,34 @@ class TestOCIStreamingToolCalls: "role": "ASSISTANT", "content": None, "toolCalls": [ - { - "type": "FUNCTION", - "id": "call_1", - "name": "get_weather", - # Missing arguments - }, + {"type": "FUNCTION", "id": "call_1", "name": "get_weather"}, { "type": "FUNCTION", "name": "get_time", "arguments": '{"timezone": "UTC"}', - # Missing id }, { "type": "FUNCTION", "id": "call_3", "name": "calculate", "arguments": '{"expression": "2+2"}', - # Complete }, ], }, } - wrapper = OCIStreamWrapper( - completion_stream=iter([]), - model="meta.llama-3.1-405b-instruct", - custom_llm_provider="oci", - logging_obj=MagicMock(), - ) - - result = wrapper._handle_generic_stream_chunk(chunk_data) + result = handle_generic_stream_chunk(chunk_data) assert isinstance(result, ModelResponseStream) assert result.choices[0].delta.tool_calls is not None assert len(result.choices[0].delta.tool_calls) == 3 - # First tool call - missing arguments assert result.choices[0].delta.tool_calls[0]["id"] == "call_1" assert ( result.choices[0].delta.tool_calls[0]["function"]["name"] == "get_weather" ) assert result.choices[0].delta.tool_calls[0]["function"]["arguments"] == "" - # Second tool call - missing id gets a generated call_* id assert result.choices[0].delta.tool_calls[1]["id"].startswith("call_") assert result.choices[0].delta.tool_calls[1]["function"]["name"] == "get_time" assert ( @@ -273,7 +203,6 @@ class TestOCIStreamingToolCalls: == '{"timezone": "UTC"}' ) - # Third tool call - complete assert result.choices[0].delta.tool_calls[2]["id"] == "call_3" assert result.choices[0].delta.tool_calls[2]["function"]["name"] == "calculate" assert ( @@ -282,9 +211,7 @@ class TestOCIStreamingToolCalls: ) def test_stream_chunk_without_tool_calls(self): - """ - Test that streaming chunks without tool calls continue to work as before. - """ + """Plain text chunks (no tool calls) pass through correctly.""" chunk_data = { "index": 0, "finishReason": None, @@ -294,14 +221,7 @@ class TestOCIStreamingToolCalls: }, } - wrapper = OCIStreamWrapper( - completion_stream=iter([]), - model="meta.llama-3.1-405b-instruct", - custom_llm_provider="oci", - logging_obj=MagicMock(), - ) - - result = wrapper._handle_generic_stream_chunk(chunk_data) + result = handle_generic_stream_chunk(chunk_data) assert isinstance(result, ModelResponseStream) assert result.choices[0].delta.content == "Hello, how can I help you?" 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 index 35ae23f7dfb..4f0be397fdf 100644 --- a/tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py +++ b/tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py @@ -67,7 +67,10 @@ class TestOCIEmbedConfig: optional_params={"oci_region": "us-chicago-1"}, litellm_params={}, ) - assert url == "https://inference.generativeai.us-chicago-1.oci.oraclecloud.com/20231130/actions/embedText" + assert ( + url + == "https://inference.generativeai.us-chicago-1.oci.oraclecloud.com/20231130/actions/embedText" + ) def test_get_complete_url_respects_api_base(self): """api_base is returned as-is (caller supplies complete URL for dedicated/custom endpoints)."""