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:
Federico Kamelhar 2026-04-06 22:49:08 -04:00
parent 3262d3ff48
commit 345a9d3574
4 changed files with 964 additions and 0 deletions

View file

@ -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

View file

@ -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"""

View 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"})

View 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 == ""