From 30ecbf1dad6fb4557e9cd27473725961a8805659 Mon Sep 17 00:00:00 2001 From: Federico Kamelhar Date: Sun, 5 Apr 2026 01:03:50 -0400 Subject: [PATCH] =?UTF-8?q?fix(oci):=20address=20remaining=20bugs=20from?= =?UTF-8?q?=20issue=20#25082=20=E2=80=94=20streaming=20signed=20body,=20Co?= =?UTF-8?q?here=20stop=20sequences,=20hardcoded=20defaults?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Bug 1: sync and async streaming paths now use signed_json_body when provided instead of re-serializing data with json.dumps() — the OCI RSA-SHA256 signature covers the exact request body bytes, so re-serializing produces an invalid sig - Bug 3: Cohere stop sequences now map to 'stopSequences' (was incorrectly 'stop') - Bug 4: removed hardcoded Cohere defaults (maxTokens=600, temperature=1, topK=0, topP=0.75, frequencyPenalty=0) that silently overrode user intent on every call - Added 6 unit tests covering all three fixes --- litellm/llms/oci/chat/transformation.py | 18 +-- .../oci/chat/test_oci_chat_transformation.py | 150 ++++++++++++++++++ 2 files changed, 156 insertions(+), 12 deletions(-) diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index 467814c7622..4cb8b71856d 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -145,10 +145,12 @@ class OCIChatConfig(BaseConfig): "response_format": "responseFormat", } - # Cohere uses the same parameter keys as GENERIC except tool_choice is unsupported. + # Cohere uses the same parameter keys as GENERIC with two differences: + # - tool_choice is unsupported + # - stop sequences are named "stopSequences" not "stop" # Build a *separate* frozen reference map so callers never mutate the canonical dict. self._openai_to_oci_cohere_param_map = { - k: v + k: ("stopSequences" if k == "stop" else v) for k, v in self.openai_to_oci_generic_param_map.items() if k not in ("tool_choice", "max_retries") } @@ -281,14 +283,6 @@ class OCIChatConfig(BaseConfig): selected_params: Dict = {} if vendor == OCIVendors.COHERE: open_ai_to_oci_param_map = self._openai_to_oci_cohere_param_map - # Add default values for Cohere API - selected_params = { - "maxTokens": 600, - "temperature": 1, - "topK": 0, - "topP": 0.75, - "frequencyPenalty": 0, - } else: open_ai_to_oci_param_map = self.openai_to_oci_generic_param_map @@ -717,7 +711,7 @@ class OCIChatConfig(BaseConfig): response = client.post( api_base, headers=headers, - data=json.dumps(data), + data=signed_json_body if signed_json_body is not None else json.dumps(data), stream=True, logging_obj=logging_obj, timeout=STREAMING_TIMEOUT, @@ -768,7 +762,7 @@ class OCIChatConfig(BaseConfig): response = await client.post( api_base, headers=headers, - data=json.dumps(data), + data=signed_json_body if signed_json_body is not None else json.dumps(data), stream=True, logging_obj=logging_obj, timeout=STREAMING_TIMEOUT, 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 01c75de1bfb..035fa7d5846 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 @@ -885,3 +885,153 @@ class TestOCIProviderEmbeddingConfig: f"Expected exactly 1 OCI branch in get_provider_embedding_config, found {oci_count}. " "A duplicate dead-code branch may have been reintroduced." ) + + +class TestOCICohereParamMapping: + """ + Unit tests for Bug 3 (stop → stopSequences) and Bug 4 (hardcoded defaults removed). + """ + + def _make_config(self): + return OCIChatConfig() + + def test_cohere_stop_maps_to_stop_sequences(self): + """Bug 3: Cohere API uses 'stopSequences', not 'stop'.""" + config = self._make_config() + result = config.map_openai_params( + non_default_params={"stop": ["END", "STOP"]}, + optional_params={}, + model="cohere.command-latest", + drop_params=False, + ) + assert "stopSequences" in result, "stop should map to stopSequences for Cohere" + assert result["stopSequences"] == ["END", "STOP"] + assert "stop" not in result + + def test_generic_stop_maps_to_stop(self): + """GENERIC vendors (Meta, Google, xAI) keep 'stop' as-is.""" + config = self._make_config() + result = config.map_openai_params( + non_default_params={"stop": ["END"]}, + optional_params={}, + model="meta.llama-3.3-70b-instruct", + drop_params=False, + ) + assert result.get("stop") == ["END"] + assert "stopSequences" not in result + + def test_cohere_no_hardcoded_defaults(self): + """Bug 4: Cohere calls must not inject maxTokens/temperature/topK/topP/frequencyPenalty + when the user hasn't provided them.""" + config = self._make_config() + result = config.map_openai_params( + non_default_params={}, + optional_params={}, + model="cohere.command-latest", + drop_params=False, + ) + for injected in ("maxTokens", "temperature", "topK", "topP", "frequencyPenalty"): + assert injected not in result, ( + f"'{injected}' should not be injected when user did not provide it" + ) + + def test_cohere_explicit_params_still_passed(self): + """User-provided Cohere params must still be forwarded correctly.""" + config = self._make_config() + result = config.map_openai_params( + non_default_params={"max_tokens": 200, "temperature": 0.5}, + optional_params={}, + model="cohere.command-latest", + drop_params=False, + ) + assert result.get("maxTokens") == 200 + assert result.get("temperature") == 0.5 + + +class TestOCIStreamingSignedBody: + """ + Unit test for Bug 1: sync and async streaming paths must use signed_json_body + when provided, not re-serialize data with json.dumps(). + """ + + def test_get_custom_stream_wrapper_uses_signed_body(self, monkeypatch): + """ + When signed_json_body is provided, the POST must use that exact bytes object, + not json.dumps(data) — otherwise the RSA-SHA256 signature is invalid. + """ + import httpx + from unittest.mock import MagicMock, patch + + config = OCIChatConfig() + signed_bytes = b'{"signed": true}' + posted_data = {} + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.iter_text.return_value = iter([]) + + mock_client = MagicMock() + mock_client.post.return_value = mock_response + + def capture_post(url, **kwargs): + posted_data["data"] = kwargs.get("data") + return mock_response + + mock_client.post.side_effect = capture_post + + mock_logging = MagicMock() + + config.get_sync_custom_stream_wrapper( + api_base="https://example.com", + headers={}, + data={"key": "value"}, + messages=[], + model="meta.llama-3.3-70b-instruct", + custom_llm_provider="oci", + logging_obj=mock_logging, + client=mock_client, + signed_json_body=signed_bytes, + ) + + assert posted_data["data"] == signed_bytes, ( + "Streaming must use signed_json_body, not re-serialize data" + ) + + def test_get_custom_stream_wrapper_fallback_without_signed_body(self, monkeypatch): + """When signed_json_body is None, fall back to json.dumps(data).""" + import json + from unittest.mock import MagicMock + + config = OCIChatConfig() + posted_data = {} + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.iter_text.return_value = iter([]) + + mock_client = MagicMock() + + def capture_post(url, **kwargs): + posted_data["data"] = kwargs.get("data") + return mock_response + + mock_client.post.side_effect = capture_post + + mock_logging = MagicMock() + payload = {"key": "value"} + + config.get_sync_custom_stream_wrapper( + api_base="https://example.com", + headers={}, + data=payload, + messages=[], + model="meta.llama-3.3-70b-instruct", + custom_llm_provider="oci", + logging_obj=mock_logging, + client=mock_client, + signed_json_body=None, + ) + + assert posted_data["data"] == json.dumps(payload), ( + "Without signed_json_body, must fall back to json.dumps(data)" + )