diff --git a/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py b/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py index 3d566dd7c64..8b866a8da44 100644 --- a/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py +++ b/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py @@ -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 diff --git a/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py b/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py index 252404759c5..7ef69a4a2e5 100644 --- a/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py +++ b/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py @@ -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""" diff --git a/tests/test_litellm/llms/oci/chat/test_oci_generic_chat.py b/tests/test_litellm/llms/oci/chat/test_oci_generic_chat.py new file mode 100644 index 00000000000..4f02ce8f504 --- /dev/null +++ b/tests/test_litellm/llms/oci/chat/test_oci_generic_chat.py @@ -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"}) diff --git a/tests/test_litellm/llms/oci/test_oci_common_utils.py b/tests/test_litellm/llms/oci/test_oci_common_utils.py new file mode 100644 index 00000000000..5d4e2602f5a --- /dev/null +++ b/tests/test_litellm/llms/oci/test_oci_common_utils.py @@ -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 == ""