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