mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
test(oci): add unit tests to improve patch coverage
- test_oci_common_utils.py (new): covers sha256_base64, build_signature_string, OCIRequestWrapper.path_url, resolve_oci_credentials, get_oci_base_url, validate_oci_environment, sign_with_oci_signer error paths, sign_oci_request routing, load_private_key_from_file error paths, resolve_oci_schema_refs (including circular ref and external $ref), resolve_oci_schema_anyof, sanitize_oci_schema (all branches), enrich_cohere_param_description - test_oci_generic_chat.py (new): covers content-message error paths (non-dict item, unsupported type, non-string text, invalid image_url), tool-call validation error paths, adapt_messages_to_generic_oci_standard error paths, handle_generic_response (None message, text content, tool calls), handle_generic_stream_chunk (finish reasons, streaming tool calls), OCIStreamWrapper non-string chunk error - test_oci_chat_transformation.py: add error paths for validate_environment (empty messages), transform_request (missing compartment_id, Cohere without user messages), transform_response (error key), map_openai_params (unsupported param with and without drop_params), tool_choice string mapping - test_oci_cohere_tool_calls.py: add edge cases for stream chunk finish reasons (TOOL_CALL, MAX_TOKENS, unknown), _extract_text_content with non-dict list items and non-string input, adapt_messages_to_cohere_standard with malformed JSON tool arguments
This commit is contained in:
parent
3262d3ff48
commit
345a9d3574
4 changed files with 964 additions and 0 deletions
|
|
@ -1042,3 +1042,114 @@ class TestOCIStreamingSignedBody:
|
|||
assert posted_data["data"] == json.dumps(
|
||||
payload
|
||||
), "Without signed_json_body, must fall back to json.dumps(data)"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Additional coverage: error paths in validate_environment, transform_request,
|
||||
# transform_response, and map_openai_params
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestOCIChatConfigErrorPaths:
|
||||
def test_validate_environment_empty_messages_raises(self):
|
||||
config = OCIChatConfig()
|
||||
with pytest.raises(Exception, match="messages"):
|
||||
config.validate_environment(
|
||||
headers={},
|
||||
model=TEST_MODEL_NAME,
|
||||
messages=[],
|
||||
optional_params={
|
||||
"oci_signer": MagicMock(),
|
||||
"oci_compartment_id": TEST_COMPARTMENT_ID,
|
||||
},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
def test_transform_request_missing_compartment_id_raises(self):
|
||||
config = OCIChatConfig()
|
||||
with pytest.raises(Exception, match="oci_compartment_id"):
|
||||
config.transform_request(
|
||||
model=TEST_MODEL_NAME,
|
||||
messages=TEST_MESSAGES, # type: ignore
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
def test_transform_request_cohere_no_user_message_raises(self):
|
||||
config = OCIChatConfig()
|
||||
with pytest.raises(Exception, match="user message"):
|
||||
config.transform_request(
|
||||
model="cohere.command-latest",
|
||||
messages=[{"role": "system", "content": "You are helpful."}], # type: ignore
|
||||
optional_params={"oci_compartment_id": TEST_COMPARTMENT_ID},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
def test_transform_response_error_key_raises(self):
|
||||
config = OCIChatConfig()
|
||||
response = httpx.Response(
|
||||
status_code=400,
|
||||
json={"error": "model not found"},
|
||||
)
|
||||
with pytest.raises(Exception, match="model not found"):
|
||||
config.transform_response(
|
||||
model=TEST_MODEL_NAME,
|
||||
raw_response=response,
|
||||
model_response=ModelResponse(),
|
||||
logging_obj={}, # type: ignore
|
||||
request_data={},
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding={},
|
||||
)
|
||||
|
||||
def test_map_openai_params_unsupported_param_raises_without_drop(self):
|
||||
config = OCIChatConfig()
|
||||
with pytest.raises(Exception, match="not supported on OCI"):
|
||||
config.map_openai_params(
|
||||
non_default_params={"audio": {"voice": "alloy"}},
|
||||
optional_params={},
|
||||
model=TEST_MODEL_NAME,
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
def test_map_openai_params_unsupported_param_dropped(self):
|
||||
config = OCIChatConfig()
|
||||
result = config.map_openai_params(
|
||||
non_default_params={"audio": {"voice": "alloy"}},
|
||||
optional_params={},
|
||||
model=TEST_MODEL_NAME,
|
||||
drop_params=True,
|
||||
)
|
||||
assert "audio" not in result
|
||||
|
||||
def test_transform_request_tool_choice_string_mapped(self):
|
||||
config = OCIChatConfig()
|
||||
result = config.transform_request(
|
||||
model=TEST_MODEL_NAME,
|
||||
messages=TEST_MESSAGES, # type: ignore
|
||||
optional_params={
|
||||
"oci_compartment_id": TEST_COMPARTMENT_ID,
|
||||
"tool_choice": "auto",
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "fn",
|
||||
"description": "d",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert result["chatRequest"]["toolChoice"] == {"type": "AUTO"}
|
||||
|
||||
|
||||
import pytest
|
||||
from unittest.mock import MagicMock
|
||||
|
|
|
|||
|
|
@ -741,6 +741,88 @@ class TestOCICoherePreambleOverride:
|
|||
assert roles == ["USER", "CHATBOT"]
|
||||
|
||||
|
||||
class TestCohereStreamChunkEdgeCases:
|
||||
"""Additional coverage for handle_cohere_stream_chunk error/edge paths."""
|
||||
|
||||
def _wrapper(self):
|
||||
from litellm.llms.oci.chat.transformation import OCIStreamWrapper
|
||||
|
||||
return OCIStreamWrapper(
|
||||
completion_stream=MagicMock(),
|
||||
model="cohere.command-latest",
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
def test_stream_chunk_tool_call_finish_reason(self):
|
||||
wrapper = self._wrapper()
|
||||
chunk = {
|
||||
"apiFormat": "COHERE",
|
||||
"text": "",
|
||||
"index": 0,
|
||||
"finishReason": "TOOL_CALL",
|
||||
}
|
||||
result = wrapper.chunk_creator(f"data: {json.dumps(chunk)}")
|
||||
assert result.choices[0].finish_reason == "tool_calls"
|
||||
|
||||
def test_stream_chunk_max_tokens_finish_reason(self):
|
||||
wrapper = self._wrapper()
|
||||
chunk = {"apiFormat": "COHERE", "text": "truncated", "index": 0, "finishReason": "MAX_TOKENS"}
|
||||
result = wrapper.chunk_creator(f"data: {json.dumps(chunk)}")
|
||||
assert result.choices[0].finish_reason == "length"
|
||||
|
||||
def test_stream_chunk_unknown_finish_reason_does_not_raise(self):
|
||||
from litellm.llms.oci.chat.cohere import handle_cohere_stream_chunk
|
||||
|
||||
chunk = {"apiFormat": "COHERE", "text": "", "index": 0, "finishReason": "FUTURE_REASON"}
|
||||
# Should not raise — unknown reasons fall through the elif chain unchanged
|
||||
result = handle_cohere_stream_chunk(chunk)
|
||||
assert result.choices[0] is not None
|
||||
|
||||
def test_stream_chunk_null_index_defaults_to_zero(self):
|
||||
wrapper = self._wrapper()
|
||||
chunk = {"apiFormat": "COHERE", "text": "hi", "index": None}
|
||||
result = wrapper.chunk_creator(f"data: {json.dumps(chunk)}")
|
||||
assert result.choices[0].index == 0
|
||||
|
||||
|
||||
class TestCohereMessageAdaptationEdgeCases:
|
||||
"""Coverage for adapt_messages_to_cohere_standard error paths."""
|
||||
|
||||
def test_json_decode_error_in_tool_args_defaults_to_empty(self):
|
||||
from litellm.llms.oci.chat.cohere import adapt_messages_to_cohere_standard
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "calling",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "c1",
|
||||
"type": "function",
|
||||
"function": {"name": "fn", "arguments": "NOT JSON {{{"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "follow up"},
|
||||
]
|
||||
# Should not raise — bad JSON defaults to empty params {}
|
||||
history = adapt_messages_to_cohere_standard(messages)
|
||||
assert history[0].toolCalls[0].parameters == {}
|
||||
|
||||
def test_extract_text_content_list_with_non_dict_items(self):
|
||||
from litellm.llms.oci.chat.cohere import _extract_text_content
|
||||
|
||||
# List with a non-dict item — should be silently skipped
|
||||
result = _extract_text_content([{"type": "text", "text": "hello"}, "bad_item"])
|
||||
assert result == "hello"
|
||||
|
||||
def test_extract_text_content_non_string_non_list(self):
|
||||
from litellm.llms.oci.chat.cohere import _extract_text_content
|
||||
|
||||
result = _extract_text_content(12345)
|
||||
assert result == "12345"
|
||||
|
||||
|
||||
class TestOCICohereStreaming:
|
||||
"""Test Cohere streaming functionality"""
|
||||
|
||||
|
|
|
|||
302
tests/test_litellm/llms/oci/chat/test_oci_generic_chat.py
Normal file
302
tests/test_litellm/llms/oci/chat/test_oci_generic_chat.py
Normal file
|
|
@ -0,0 +1,302 @@
|
|||
"""
|
||||
Unit tests for litellm/llms/oci/chat/generic.py — error paths and stream handling.
|
||||
"""
|
||||
|
||||
import json
|
||||
import pytest
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm import ModelResponse
|
||||
from litellm.llms.oci.chat.generic import (
|
||||
adapt_messages_to_generic_oci_standard,
|
||||
adapt_messages_to_generic_oci_standard_content_message,
|
||||
adapt_messages_to_generic_oci_standard_tool_call,
|
||||
handle_generic_response,
|
||||
handle_generic_stream_chunk,
|
||||
)
|
||||
from litellm.llms.oci.chat.transformation import OCIChatConfig, OCIStreamWrapper
|
||||
from litellm.llms.oci.common_utils import OCIError
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# adapt_messages_to_generic_oci_standard_content_message — error paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGenericContentMessageErrors:
|
||||
def test_non_dict_content_item_raises(self):
|
||||
with pytest.raises(OCIError, match="must be a dictionary"):
|
||||
adapt_messages_to_generic_oci_standard_content_message(
|
||||
"user", ["not a dict"]
|
||||
)
|
||||
|
||||
def test_non_string_type_field_raises(self):
|
||||
with pytest.raises(OCIError, match="string `type` field"):
|
||||
adapt_messages_to_generic_oci_standard_content_message(
|
||||
"user", [{"type": 123, "text": "hi"}]
|
||||
)
|
||||
|
||||
def test_unsupported_content_type_raises(self):
|
||||
with pytest.raises(OCIError, match="not supported by OCI"):
|
||||
adapt_messages_to_generic_oci_standard_content_message(
|
||||
"user", [{"type": "video_url", "url": "https://example.com/v.mp4"}]
|
||||
)
|
||||
|
||||
def test_non_string_text_raises(self):
|
||||
with pytest.raises(OCIError, match="must have a string `text` field"):
|
||||
adapt_messages_to_generic_oci_standard_content_message(
|
||||
"user", [{"type": "text", "text": 42}]
|
||||
)
|
||||
|
||||
def test_image_url_as_invalid_type_raises(self):
|
||||
with pytest.raises(OCIError, match="must be a string or an object"):
|
||||
adapt_messages_to_generic_oci_standard_content_message(
|
||||
"user", [{"type": "image_url", "image_url": 99}]
|
||||
)
|
||||
|
||||
def test_image_url_as_string(self):
|
||||
msg = adapt_messages_to_generic_oci_standard_content_message(
|
||||
"user", [{"type": "image_url", "image_url": "https://example.com/img.png"}]
|
||||
)
|
||||
assert msg.content[0].imageUrl.url == "https://example.com/img.png"
|
||||
|
||||
def test_image_url_as_dict(self):
|
||||
msg = adapt_messages_to_generic_oci_standard_content_message(
|
||||
"user",
|
||||
[{"type": "image_url", "image_url": {"url": "https://example.com/img.png"}}],
|
||||
)
|
||||
assert msg.content[0].imageUrl.url == "https://example.com/img.png"
|
||||
|
||||
def test_text_content_string(self):
|
||||
msg = adapt_messages_to_generic_oci_standard_content_message("user", "hello")
|
||||
assert msg.content[0].text == "hello"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# adapt_messages_to_generic_oci_standard_tool_call — error paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGenericToolCallErrors:
|
||||
def test_non_dict_tool_call_raises(self):
|
||||
with pytest.raises(OCIError, match="must be a dictionary"):
|
||||
adapt_messages_to_generic_oci_standard_tool_call("assistant", ["bad"])
|
||||
|
||||
def test_non_function_type_raises(self):
|
||||
with pytest.raises(OCIError, match="only supports function tool calls"):
|
||||
adapt_messages_to_generic_oci_standard_tool_call(
|
||||
"assistant",
|
||||
[{"type": "database", "id": "x", "function": {"name": "f", "arguments": "{}"}}],
|
||||
)
|
||||
|
||||
def test_non_string_id_raises(self):
|
||||
with pytest.raises(OCIError, match="id.*must be a string"):
|
||||
adapt_messages_to_generic_oci_standard_tool_call(
|
||||
"assistant",
|
||||
[{"type": "function", "id": 123, "function": {"name": "f", "arguments": "{}"}}],
|
||||
)
|
||||
|
||||
def test_non_dict_function_raises(self):
|
||||
with pytest.raises(OCIError, match="`function` must be a dictionary"):
|
||||
adapt_messages_to_generic_oci_standard_tool_call(
|
||||
"assistant",
|
||||
[{"type": "function", "id": "c1", "function": "not_a_dict"}],
|
||||
)
|
||||
|
||||
def test_non_string_function_name_raises(self):
|
||||
with pytest.raises(OCIError, match="function.name.*must be a string"):
|
||||
adapt_messages_to_generic_oci_standard_tool_call(
|
||||
"assistant",
|
||||
[{"type": "function", "id": "c1", "function": {"name": 5, "arguments": "{}"}}],
|
||||
)
|
||||
|
||||
def test_non_string_arguments_raises(self):
|
||||
with pytest.raises(OCIError, match="arguments.*must be a JSON string"):
|
||||
adapt_messages_to_generic_oci_standard_tool_call(
|
||||
"assistant",
|
||||
[
|
||||
{
|
||||
"type": "function",
|
||||
"id": "c1",
|
||||
"function": {"name": "fn", "arguments": {"key": "val"}},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# adapt_messages_to_generic_oci_standard — combined paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGenericMessageAdaptation:
|
||||
def test_tool_calls_not_list_raises(self):
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": "not_a_list",
|
||||
}
|
||||
]
|
||||
with pytest.raises(OCIError, match="`tool_calls` must be a list"):
|
||||
adapt_messages_to_generic_oci_standard(messages)
|
||||
|
||||
def test_tool_result_non_string_tool_call_id_raises(self):
|
||||
messages = [
|
||||
{"role": "tool", "content": "result", "tool_call_id": 999}
|
||||
]
|
||||
with pytest.raises(OCIError, match="string `tool_call_id`"):
|
||||
adapt_messages_to_generic_oci_standard(messages)
|
||||
|
||||
def test_tool_result_non_string_content_raises(self):
|
||||
messages = [
|
||||
{"role": "tool", "content": {"structured": "data"}, "tool_call_id": "c1"}
|
||||
]
|
||||
with pytest.raises(OCIError, match="`content` must be a string"):
|
||||
adapt_messages_to_generic_oci_standard(messages)
|
||||
|
||||
def test_non_string_non_list_content_raises(self):
|
||||
messages = [{"role": "user", "content": 42}]
|
||||
with pytest.raises(OCIError, match="`content` must be a string or list"):
|
||||
adapt_messages_to_generic_oci_standard(messages)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# handle_generic_response — error and None message paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestHandleGenericResponse:
|
||||
def _make_response(self, body: dict, status: int = 200) -> httpx.Response:
|
||||
return httpx.Response(status_code=status, json=body)
|
||||
|
||||
def _valid_body(self, message=None):
|
||||
return {
|
||||
"modelId": "xai.grok-4",
|
||||
"modelVersion": "1",
|
||||
"chatResponse": {
|
||||
"apiFormat": "GENERIC",
|
||||
"timeCreated": "2024-01-01T00:00:00Z",
|
||||
"choices": [{"message": message, "finishReason": "COMPLETE", "index": 0}],
|
||||
"usage": {"promptTokens": 5, "completionTokens": 5, "totalTokens": 10},
|
||||
},
|
||||
}
|
||||
|
||||
def test_none_response_message(self):
|
||||
body = self._valid_body(message=None)
|
||||
raw = self._make_response(body)
|
||||
# Should not raise — None message means no content set
|
||||
result = handle_generic_response(body, "xai.grok-4", ModelResponse(), raw)
|
||||
assert result.model == "xai.grok-4"
|
||||
|
||||
def test_response_with_text_content(self):
|
||||
body = self._valid_body(
|
||||
message={
|
||||
"role": "ASSISTANT",
|
||||
"content": [{"type": "TEXT", "text": "Hello!"}],
|
||||
}
|
||||
)
|
||||
raw = self._make_response(body)
|
||||
result = handle_generic_response(body, "xai.grok-4", ModelResponse(), raw)
|
||||
assert result.choices[0].message.content == "Hello!"
|
||||
|
||||
def test_response_with_tool_calls(self):
|
||||
body = self._valid_body(
|
||||
message={
|
||||
"role": "ASSISTANT",
|
||||
"content": [],
|
||||
"toolCalls": [
|
||||
{
|
||||
"id": "call_abc",
|
||||
"type": "FUNCTION",
|
||||
"name": "get_weather",
|
||||
"arguments": '{"location": "Tokyo"}',
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
raw = self._make_response(body)
|
||||
result = handle_generic_response(body, "xai.grok-4", ModelResponse(), raw)
|
||||
assert result.choices[0].message.tool_calls is not None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# handle_generic_stream_chunk — finish reasons and error paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestHandleGenericStreamChunk:
|
||||
def test_max_tokens_finish_reason(self):
|
||||
chunk = {"apiFormat": "GENERIC", "index": 0, "finishReason": "MAX_TOKENS"}
|
||||
result = handle_generic_stream_chunk(chunk)
|
||||
assert result.choices[0].finish_reason == "length"
|
||||
|
||||
def test_tool_calls_finish_reason(self):
|
||||
chunk = {"apiFormat": "GENERIC", "index": 0, "finishReason": "TOOL_CALLS"}
|
||||
result = handle_generic_stream_chunk(chunk)
|
||||
assert result.choices[0].finish_reason == "tool_calls"
|
||||
|
||||
def test_unknown_finish_reason_does_not_raise(self):
|
||||
chunk = {"apiFormat": "GENERIC", "index": 0, "finishReason": "SOME_NEW_REASON"}
|
||||
result = handle_generic_stream_chunk(chunk)
|
||||
assert result.choices[0] is not None
|
||||
|
||||
def test_null_index_defaults_to_zero(self):
|
||||
chunk = {"apiFormat": "GENERIC", "index": None, "finishReason": None}
|
||||
result = handle_generic_stream_chunk(chunk)
|
||||
assert result.choices[0].index == 0
|
||||
|
||||
def test_image_content_in_stream_raises(self):
|
||||
from litellm.types.llms.oci import OCIImageContentPart, OCIImageUrl, OCIMessage
|
||||
|
||||
chunk = {
|
||||
"apiFormat": "GENERIC",
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "ASSISTANT",
|
||||
"content": [{"type": "IMAGE", "imageUrl": {"url": "https://example.com/img.png"}}],
|
||||
},
|
||||
}
|
||||
with pytest.raises(OCIError, match="image content"):
|
||||
handle_generic_stream_chunk(chunk)
|
||||
|
||||
def test_stream_chunk_with_tool_calls(self):
|
||||
chunk = {
|
||||
"apiFormat": "GENERIC",
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "ASSISTANT",
|
||||
"content": [],
|
||||
"toolCalls": [
|
||||
{
|
||||
"id": "call_abc",
|
||||
"type": "FUNCTION",
|
||||
"name": "get_weather",
|
||||
"arguments": '{"location": "Tokyo"}',
|
||||
}
|
||||
],
|
||||
},
|
||||
}
|
||||
result = handle_generic_stream_chunk(chunk)
|
||||
assert result.choices[0].delta.tool_calls is not None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OCIStreamWrapper.chunk_creator — non-string chunk
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestOCIStreamWrapperChunkCreator:
|
||||
def _wrapper(self):
|
||||
return OCIStreamWrapper(
|
||||
completion_stream=MagicMock(),
|
||||
model="xai.grok-4",
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
def test_non_string_chunk_raises(self):
|
||||
w = self._wrapper()
|
||||
with pytest.raises(ValueError, match="not a string"):
|
||||
w.chunk_creator({"already": "parsed"})
|
||||
469
tests/test_litellm/llms/oci/test_oci_common_utils.py
Normal file
469
tests/test_litellm/llms/oci/test_oci_common_utils.py
Normal file
|
|
@ -0,0 +1,469 @@
|
|||
"""
|
||||
Unit tests for litellm/llms/oci/common_utils.py.
|
||||
|
||||
Covers schema utilities, signing helpers, and credential resolution paths
|
||||
that require no real OCI credentials or network calls.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from litellm.llms.oci.common_utils import (
|
||||
OCI_API_VERSION,
|
||||
OCIError,
|
||||
OCIRequestWrapper,
|
||||
build_signature_string,
|
||||
enrich_cohere_param_description,
|
||||
get_oci_base_url,
|
||||
resolve_oci_credentials,
|
||||
resolve_oci_schema_anyof,
|
||||
resolve_oci_schema_refs,
|
||||
sanitize_oci_schema,
|
||||
sha256_base64,
|
||||
sign_oci_request,
|
||||
sign_with_oci_signer,
|
||||
validate_oci_environment,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OCI_API_VERSION
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_oci_api_version_constant():
|
||||
assert OCI_API_VERSION == "20231130"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# sha256_base64
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_sha256_base64_known_value():
|
||||
import base64, hashlib
|
||||
|
||||
data = b"hello"
|
||||
expected = base64.b64encode(hashlib.sha256(data).digest()).decode()
|
||||
assert sha256_base64(data) == expected
|
||||
|
||||
|
||||
def test_sha256_base64_empty():
|
||||
result = sha256_base64(b"")
|
||||
assert isinstance(result, str)
|
||||
assert len(result) > 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# build_signature_string
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_build_signature_string_request_target():
|
||||
headers = {"host": "example.com", "date": "Mon, 01 Jan 2024 00:00:00 GMT"}
|
||||
result = build_signature_string(
|
||||
"POST", "/20231130/actions/chat", headers, ["(request-target)", "host", "date"]
|
||||
)
|
||||
lines = result.split("\n")
|
||||
assert lines[0] == "(request-target): post /20231130/actions/chat"
|
||||
assert lines[1] == "host: example.com"
|
||||
assert lines[2] == "date: Mon, 01 Jan 2024 00:00:00 GMT"
|
||||
|
||||
|
||||
def test_build_signature_string_method_lowercased():
|
||||
headers = {"host": "h"}
|
||||
result = build_signature_string("GET", "/path", headers, ["(request-target)"])
|
||||
assert result == "(request-target): get /path"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OCIRequestWrapper.path_url
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_request_wrapper_path_url_no_query():
|
||||
w = OCIRequestWrapper(
|
||||
method="POST",
|
||||
url="https://inference.generativeai.us-ashburn-1.oci.oraclecloud.com/20231130/actions/chat",
|
||||
headers={},
|
||||
body=b"",
|
||||
)
|
||||
assert w.path_url == "/20231130/actions/chat"
|
||||
|
||||
|
||||
def test_request_wrapper_path_url_with_query():
|
||||
w = OCIRequestWrapper(
|
||||
method="GET",
|
||||
url="https://example.com/path?foo=bar&baz=1",
|
||||
headers={},
|
||||
body=b"",
|
||||
)
|
||||
assert w.path_url == "/path?foo=bar&baz=1"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# resolve_oci_credentials
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resolve_credentials_from_params():
|
||||
params = {
|
||||
"oci_region": "eu-frankfurt-1",
|
||||
"oci_user": "user1",
|
||||
"oci_fingerprint": "fp1",
|
||||
"oci_tenancy": "tenant1",
|
||||
"oci_key": "key_content",
|
||||
"oci_compartment_id": "comp1",
|
||||
}
|
||||
result = resolve_oci_credentials(params)
|
||||
assert result["oci_region"] == "eu-frankfurt-1"
|
||||
assert result["oci_user"] == "user1"
|
||||
assert result["oci_compartment_id"] == "comp1"
|
||||
|
||||
|
||||
def test_resolve_credentials_env_fallback(monkeypatch):
|
||||
monkeypatch.setenv("OCI_REGION", "ap-tokyo-1")
|
||||
monkeypatch.setenv("OCI_USER", "env_user")
|
||||
monkeypatch.setenv("OCI_COMPARTMENT_ID", "env_comp")
|
||||
result = resolve_oci_credentials({})
|
||||
assert result["oci_region"] == "ap-tokyo-1"
|
||||
assert result["oci_user"] == "env_user"
|
||||
assert result["oci_compartment_id"] == "env_comp"
|
||||
|
||||
|
||||
def test_resolve_credentials_region_default(monkeypatch):
|
||||
monkeypatch.delenv("OCI_REGION", raising=False)
|
||||
result = resolve_oci_credentials({})
|
||||
assert result["oci_region"] == "us-ashburn-1"
|
||||
|
||||
|
||||
def test_resolve_credentials_params_override_env(monkeypatch):
|
||||
monkeypatch.setenv("OCI_REGION", "ap-tokyo-1")
|
||||
result = resolve_oci_credentials({"oci_region": "us-phoenix-1"})
|
||||
assert result["oci_region"] == "us-phoenix-1"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_oci_base_url
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_get_oci_base_url_explicit_api_base():
|
||||
url = get_oci_base_url({}, api_base="https://custom.endpoint.com/")
|
||||
assert url == "https://custom.endpoint.com"
|
||||
|
||||
|
||||
def test_get_oci_base_url_from_region():
|
||||
url = get_oci_base_url({"oci_region": "eu-frankfurt-1"})
|
||||
assert url == "https://inference.generativeai.eu-frankfurt-1.oci.oraclecloud.com"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# validate_oci_environment
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_validate_oci_environment_sets_defaults():
|
||||
headers = {}
|
||||
result = validate_oci_environment(headers, {})
|
||||
assert result["content-type"] == "application/json"
|
||||
assert "user-agent" in result
|
||||
|
||||
|
||||
def test_validate_oci_environment_does_not_overwrite_existing():
|
||||
headers = {"content-type": "text/plain", "user-agent": "my-agent"}
|
||||
result = validate_oci_environment(headers, {})
|
||||
assert result["content-type"] == "text/plain"
|
||||
assert result["user-agent"] == "my-agent"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# sign_with_oci_signer — error paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_sign_with_oci_signer_none_raises():
|
||||
with pytest.raises(ValueError, match="oci_signer cannot be None"):
|
||||
sign_with_oci_signer({}, {"oci_signer": None}, {}, "https://example.com")
|
||||
|
||||
|
||||
def test_sign_with_oci_signer_exception_wrapped():
|
||||
bad_signer = MagicMock()
|
||||
bad_signer.do_request_sign.side_effect = RuntimeError("signing failed")
|
||||
with pytest.raises(OCIError, match="Failed to sign request"):
|
||||
sign_with_oci_signer(
|
||||
{}, {"oci_signer": bad_signer}, {"key": "val"}, "https://example.com"
|
||||
)
|
||||
|
||||
|
||||
def test_sign_with_oci_signer_success():
|
||||
signer = MagicMock()
|
||||
signer.do_request_sign.return_value = None
|
||||
headers, body = sign_with_oci_signer(
|
||||
{}, {"oci_signer": signer}, {"key": "val"}, "https://example.com"
|
||||
)
|
||||
assert isinstance(body, bytes)
|
||||
signer.do_request_sign.assert_called_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# sign_oci_request — routing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_sign_oci_request_routes_to_signer():
|
||||
signer = MagicMock()
|
||||
signer.do_request_sign.return_value = None
|
||||
headers, body = sign_oci_request(
|
||||
{}, {"oci_signer": signer}, {}, "https://example.com"
|
||||
)
|
||||
signer.do_request_sign.assert_called_once()
|
||||
|
||||
|
||||
def test_sign_oci_request_routes_to_manual_missing_creds():
|
||||
with pytest.raises(OCIError, match="Missing required OCI credentials"):
|
||||
sign_oci_request({}, {}, {}, "https://example.com")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# load_private_key_from_file — error paths (no real key needed)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_load_private_key_from_file_not_found():
|
||||
from litellm.llms.oci.common_utils import load_private_key_from_file
|
||||
|
||||
with pytest.raises(FileNotFoundError, match="Private key file not found"):
|
||||
load_private_key_from_file("/nonexistent/path/key.pem")
|
||||
|
||||
|
||||
def test_load_private_key_from_file_empty(tmp_path):
|
||||
from litellm.llms.oci.common_utils import load_private_key_from_file
|
||||
|
||||
empty = tmp_path / "empty.pem"
|
||||
empty.write_text("")
|
||||
with pytest.raises(ValueError, match="Private key file is empty"):
|
||||
load_private_key_from_file(str(empty))
|
||||
|
||||
|
||||
def test_load_private_key_from_file_os_error():
|
||||
from litellm.llms.oci.common_utils import load_private_key_from_file
|
||||
|
||||
with patch("builtins.open", side_effect=OSError("permission denied")):
|
||||
with pytest.raises(OSError, match="Failed to read private key file"):
|
||||
load_private_key_from_file("/some/path/key.pem")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# resolve_oci_schema_refs
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resolve_schema_refs_basic():
|
||||
schema = {
|
||||
"$defs": {"Foo": {"type": "string"}},
|
||||
"properties": {"x": {"$ref": "#/$defs/Foo"}},
|
||||
}
|
||||
result = resolve_oci_schema_refs(schema)
|
||||
assert result["properties"]["x"] == {"type": "string"}
|
||||
assert "$defs" not in result
|
||||
|
||||
|
||||
def test_resolve_schema_refs_external_ref_unchanged():
|
||||
schema = {"properties": {"x": {"$ref": "https://example.com/schema"}}}
|
||||
result = resolve_oci_schema_refs(schema)
|
||||
assert result["properties"]["x"] == {"$ref": "https://example.com/schema"}
|
||||
|
||||
|
||||
def test_resolve_schema_refs_circular_breaks_cycle():
|
||||
schema = {
|
||||
"$defs": {"Node": {"properties": {"child": {"$ref": "#/$defs/Node"}}}},
|
||||
"properties": {"root": {"$ref": "#/$defs/Node"}},
|
||||
}
|
||||
result = resolve_oci_schema_refs(schema)
|
||||
# Should not raise; circular ref replaced with {"type": "object"}
|
||||
child = result["properties"]["root"]["properties"]["child"]
|
||||
assert child == {"type": "object"}
|
||||
|
||||
|
||||
def test_resolve_schema_refs_no_defs():
|
||||
schema = {"type": "object", "properties": {"x": {"type": "string"}}}
|
||||
result = resolve_oci_schema_refs(schema)
|
||||
assert result == schema
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# resolve_oci_schema_anyof
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resolve_schema_anyof_optional_field():
|
||||
schema = {"anyOf": [{"type": "string"}, {"type": "null"}]}
|
||||
result = resolve_oci_schema_anyof(schema)
|
||||
assert result["type"] == "string"
|
||||
assert "anyOf" not in result
|
||||
|
||||
|
||||
def test_resolve_schema_anyof_all_null_returns_empty():
|
||||
schema = {"anyOf": [{"type": "null"}, {"type": "null"}]}
|
||||
result = resolve_oci_schema_anyof(schema)
|
||||
# No non-null branch — anyOf stays or schema unchanged
|
||||
# The function only strips anyOf when there IS a non-null branch
|
||||
assert "anyOf" in result
|
||||
|
||||
|
||||
def test_resolve_schema_anyof_no_anyof_unchanged():
|
||||
schema = {"type": "string", "description": "A name"}
|
||||
assert resolve_oci_schema_anyof(schema) == schema
|
||||
|
||||
|
||||
def test_resolve_schema_anyof_nested():
|
||||
schema = {
|
||||
"properties": {
|
||||
"age": {"anyOf": [{"type": "integer"}, {"type": "null"}]}
|
||||
}
|
||||
}
|
||||
result = resolve_oci_schema_anyof(schema)
|
||||
assert result["properties"]["age"]["type"] == "integer"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# sanitize_oci_schema
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_sanitize_schema_removes_title():
|
||||
schema = {"title": "MyModel", "type": "object", "properties": {}}
|
||||
result = sanitize_oci_schema(schema)
|
||||
assert "title" not in result
|
||||
|
||||
|
||||
def test_sanitize_schema_removes_null_default():
|
||||
schema = {"type": "string", "default": None}
|
||||
result = sanitize_oci_schema(schema)
|
||||
assert "default" not in result
|
||||
|
||||
|
||||
def test_sanitize_schema_keeps_non_null_default():
|
||||
schema = {"type": "string", "default": "hello"}
|
||||
result = sanitize_oci_schema(schema)
|
||||
assert result["default"] == "hello"
|
||||
|
||||
|
||||
def test_sanitize_schema_type_any_becomes_object():
|
||||
schema = {"type": "any"}
|
||||
result = sanitize_oci_schema(schema)
|
||||
assert result["type"] == "object"
|
||||
|
||||
|
||||
def test_sanitize_schema_type_list_picks_non_null():
|
||||
schema = {"type": ["string", "null"]}
|
||||
result = sanitize_oci_schema(schema)
|
||||
assert result["type"] == "string"
|
||||
|
||||
|
||||
def test_sanitize_schema_type_list_all_null_becomes_string():
|
||||
schema = {"type": ["null"]}
|
||||
result = sanitize_oci_schema(schema)
|
||||
assert result["type"] == "string"
|
||||
|
||||
|
||||
def test_sanitize_schema_array_gets_items():
|
||||
schema = {"type": "array"}
|
||||
result = sanitize_oci_schema(schema)
|
||||
assert result["items"] == {"type": "object"}
|
||||
|
||||
|
||||
def test_sanitize_schema_array_keeps_existing_items():
|
||||
schema = {"type": "array", "items": {"type": "string"}}
|
||||
result = sanitize_oci_schema(schema)
|
||||
assert result["items"] == {"type": "string"}
|
||||
|
||||
|
||||
def test_sanitize_schema_required_filters_missing_properties():
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {"a": {"type": "string"}},
|
||||
"required": ["a", "b"], # "b" not in properties
|
||||
}
|
||||
result = sanitize_oci_schema(schema)
|
||||
assert result["required"] == ["a"]
|
||||
|
||||
|
||||
def test_sanitize_schema_required_non_list_becomes_empty():
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {"a": {"type": "string"}},
|
||||
"required": "a", # invalid: string instead of list
|
||||
}
|
||||
result = sanitize_oci_schema(schema)
|
||||
assert result["required"] == []
|
||||
|
||||
|
||||
def test_sanitize_schema_list_input():
|
||||
schemas = [{"title": "A", "type": "string"}, {"title": "B", "type": "integer"}]
|
||||
result = sanitize_oci_schema(schemas)
|
||||
assert all("title" not in s for s in result)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# enrich_cohere_param_description
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_enrich_description_enum():
|
||||
result = enrich_cohere_param_description("A color", {"enum": ["red", "blue"]})
|
||||
assert "Allowed values: ['red', 'blue']" in result
|
||||
|
||||
|
||||
def test_enrich_description_format():
|
||||
result = enrich_cohere_param_description("A date", {"format": "date-time"})
|
||||
assert "Format: date-time" in result
|
||||
|
||||
|
||||
def test_enrich_description_range_both():
|
||||
result = enrich_cohere_param_description("A number", {"minimum": 0, "maximum": 100})
|
||||
assert "Range: min=0, max=100" in result
|
||||
|
||||
|
||||
def test_enrich_description_range_min_only():
|
||||
result = enrich_cohere_param_description("A number", {"minimum": 1})
|
||||
assert "Range: min=1" in result
|
||||
assert "max" not in result
|
||||
|
||||
|
||||
def test_enrich_description_range_max_only():
|
||||
result = enrich_cohere_param_description("", {"maximum": 10})
|
||||
assert "Range: max=10" in result
|
||||
|
||||
|
||||
def test_enrich_description_pattern():
|
||||
result = enrich_cohere_param_description("An ID", {"pattern": "^[a-z]+$"})
|
||||
assert "Pattern: ^[a-z]+$" in result
|
||||
|
||||
|
||||
def test_enrich_description_all_constraints():
|
||||
result = enrich_cohere_param_description(
|
||||
"Val",
|
||||
{
|
||||
"enum": ["a"],
|
||||
"format": "uuid",
|
||||
"minimum": 0,
|
||||
"maximum": 1,
|
||||
"pattern": ".*",
|
||||
},
|
||||
)
|
||||
assert "Allowed values" in result
|
||||
assert "Format" in result
|
||||
assert "Range" in result
|
||||
assert "Pattern" in result
|
||||
|
||||
|
||||
def test_enrich_description_no_constraints():
|
||||
result = enrich_cohere_param_description("Just a description", {})
|
||||
assert result == "Just a description"
|
||||
|
||||
|
||||
def test_enrich_description_empty_description_no_constraints():
|
||||
result = enrich_cohere_param_description("", {})
|
||||
assert result == ""
|
||||
Loading…
Add table
Reference in a new issue