From e1bd3dd1f89c5a15d14f77583d872220f41a75cd Mon Sep 17 00:00:00 2001 From: TensorNull Date: Tue, 9 Jun 2026 12:07:11 +0800 Subject: [PATCH] fix: guard cometapi custom base routing --- litellm/images/main.py | 8 +- litellm/llms/cometapi/chat/transformation.py | 2 +- litellm/llms/cometapi/common_utils.py | 13 +- litellm/llms/cometapi/embed/transformation.py | 2 +- .../image_generation/transformation.py | 4 +- litellm/main.py | 30 +++-- .../chat/test_cometapi_chat_transformation.py | 2 +- .../llms/cometapi/test_cometapi_endpoints.py | 118 ++++++++++++++++-- 8 files changed, 151 insertions(+), 28 deletions(-) diff --git a/litellm/images/main.py b/litellm/images/main.py index 770635baacd..fda773b2b61 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -441,12 +441,16 @@ def image_generation( # noqa: PLR0915 _api_base = api_base or litellm.api_base if custom_llm_provider == litellm.LlmProviders.COMETAPI: from litellm.llms.cometapi.common_utils import ( + get_cometapi_api_base, require_cometapi_api_key, ) - _api_base = api_base + cometapi_request_api_key = api_key or dynamic_api_key + _api_base = get_cometapi_api_base( + api_base, api_key=cometapi_request_api_key + ) api_key = require_cometapi_api_key( - api_key or dynamic_api_key or litellm.cometapi_key + cometapi_request_api_key or litellm.cometapi_key ) litellm_params_dict.pop("api_key", None) litellm_params_dict["api_base"] = _api_base diff --git a/litellm/llms/cometapi/chat/transformation.py b/litellm/llms/cometapi/chat/transformation.py index 3bd9334732e..b3a01258243 100644 --- a/litellm/llms/cometapi/chat/transformation.py +++ b/litellm/llms/cometapi/chat/transformation.py @@ -62,7 +62,7 @@ class CometAPIConfig(OpenAIGPTConfig): Returns: str: The complete URL for the API call. """ - return get_cometapi_complete_url(api_base, "chat/completions") + return get_cometapi_complete_url(api_base, "chat/completions", api_key=api_key) def get_error_class( self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] diff --git a/litellm/llms/cometapi/common_utils.py b/litellm/llms/cometapi/common_utils.py index 4ac71b63ffb..d679056ff3e 100644 --- a/litellm/llms/cometapi/common_utils.py +++ b/litellm/llms/cometapi/common_utils.py @@ -18,7 +18,12 @@ def get_cometapi_api_key(api_key: Optional[str] = None) -> Optional[str]: ) -def get_cometapi_api_base(api_base: Optional[str] = None) -> str: +def get_cometapi_api_base( + api_base: Optional[str] = None, api_key: Optional[str] = None +) -> str: + if api_base and not api_key: + raise ValueError("CometAPI api_base requires an explicit api_key") + return ( api_base or get_secret_str("COMETAPI_BASE_URL") @@ -27,8 +32,10 @@ def get_cometapi_api_base(api_base: Optional[str] = None) -> str: ) -def get_cometapi_complete_url(api_base: Optional[str], endpoint: str) -> str: - base_url = get_cometapi_api_base(api_base).rstrip("/") +def get_cometapi_complete_url( + api_base: Optional[str], endpoint: str, api_key: Optional[str] = None +) -> str: + base_url = get_cometapi_api_base(api_base, api_key=api_key).rstrip("/") normalized_endpoint = endpoint.strip("/") parsed_endpoint = urlsplit(normalized_endpoint) if ( diff --git a/litellm/llms/cometapi/embed/transformation.py b/litellm/llms/cometapi/embed/transformation.py index 23e1b6d440f..41e1dc46684 100644 --- a/litellm/llms/cometapi/embed/transformation.py +++ b/litellm/llms/cometapi/embed/transformation.py @@ -42,7 +42,7 @@ class CometAPIEmbeddingConfig(BaseEmbeddingConfig): """ Get the complete URL for the CometAPI embedding endpoint. """ - return get_cometapi_complete_url(api_base, "embeddings") + return get_cometapi_complete_url(api_base, "embeddings", api_key=api_key) def validate_environment( self, diff --git a/litellm/llms/cometapi/image_generation/transformation.py b/litellm/llms/cometapi/image_generation/transformation.py index 4e9358a1acb..5cc37200805 100644 --- a/litellm/llms/cometapi/image_generation/transformation.py +++ b/litellm/llms/cometapi/image_generation/transformation.py @@ -104,7 +104,9 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig): """ Get the complete url for the request """ - return get_cometapi_complete_url(api_base, "images/generations") + return get_cometapi_complete_url( + api_base, "images/generations", api_key=api_key + ) def validate_environment( self, diff --git a/litellm/main.py b/litellm/main.py index 2b7b9dd81de..381912004d0 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -2543,8 +2543,9 @@ def completion( # type: ignore # noqa: PLR0915 stream=stream, ) elif custom_llm_provider == "cometapi": + cometapi_request_api_key = api_key api_key = require_cometapi_api_key(api_key or litellm.cometapi_key) - api_base = get_cometapi_api_base(api_base) + api_base = get_cometapi_api_base(api_base, api_key=cometapi_request_api_key) ## COMPLETION CALL response = base_llm_http_handler.completion( @@ -5826,8 +5827,9 @@ def embedding( # noqa: PLR0915 litellm_params={}, ) elif custom_llm_provider == "cometapi": + cometapi_request_api_key = api_key api_key = require_cometapi_api_key(api_key or litellm.cometapi_key) - api_base = get_cometapi_api_base(api_base) + api_base = get_cometapi_api_base(api_base, api_key=cometapi_request_api_key) response = base_llm_http_handler.embedding( model=model, input=input, @@ -6415,10 +6417,14 @@ def moderation( pass if custom_llm_provider == "cometapi": + cometapi_request_api_key = api_key or _dynamic_api_key + cometapi_request_api_base = api_base or _dynamic_api_base api_key = require_cometapi_api_key( - api_key or _dynamic_api_key or litellm.cometapi_key + cometapi_request_api_key or litellm.cometapi_key + ) + api_base = get_cometapi_api_base( + cometapi_request_api_base, api_key=cometapi_request_api_key ) - api_base = get_cometapi_api_base(api_base or _dynamic_api_base) else: api_key = ( api_key @@ -6481,10 +6487,14 @@ async def amoderation( pass if custom_llm_provider == "cometapi": + cometapi_request_api_key = api_key or _dynamic_api_key + cometapi_request_api_base = optional_params.api_base or _dynamic_api_base api_key = require_cometapi_api_key( - api_key or _dynamic_api_key or litellm.cometapi_key + cometapi_request_api_key or litellm.cometapi_key + ) + api_base = get_cometapi_api_base( + cometapi_request_api_base, api_key=cometapi_request_api_key ) - api_base = get_cometapi_api_base(optional_params.api_base or _dynamic_api_base) else: api_key = ( api_key @@ -6743,8 +6753,9 @@ def transcription( # noqa: PLR0915 litellm_params=litellm_params_dict, ) elif custom_llm_provider == "cometapi": + cometapi_request_api_key = api_key api_key = require_cometapi_api_key(api_key or litellm.cometapi_key) - api_base = get_cometapi_api_base(api_base) + api_base = get_cometapi_api_base(api_base, api_key=cometapi_request_api_key) response = openai_audio_transcriptions.audio_transcriptions( model=model, audio_file=file, @@ -6994,10 +7005,11 @@ def speech( # noqa: PLR0915 model=model, llm_provider=custom_llm_provider, ) + cometapi_request_api_key = api_key or dynamic_api_key api_key = require_cometapi_api_key( - api_key or dynamic_api_key or litellm.cometapi_key + cometapi_request_api_key or litellm.cometapi_key ) - api_base = get_cometapi_api_base(api_base) + api_base = get_cometapi_api_base(api_base, api_key=cometapi_request_api_key) headers = headers or litellm.headers response = openai_chat_completions.audio_speech( model=model, diff --git a/tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py b/tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py index 7e5483dc8f8..28b383b7b2e 100644 --- a/tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py +++ b/tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py @@ -216,7 +216,7 @@ class TestCometAPIConfig: assert ( config.get_complete_url( api_base="https://api.cometapi.com/v1", - api_key=None, + api_key="comet-explicit-key", model="gpt-5.5", optional_params={}, litellm_params={}, diff --git a/tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py b/tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py index d6bfa936c20..a47455330a8 100644 --- a/tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py +++ b/tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py @@ -67,11 +67,12 @@ def _pollute_openai_globals(monkeypatch): ) def test_cometapi_embedding_url_normalization(api_base, expected, monkeypatch): _clear_cometapi_env(monkeypatch) + request_api_key = "comet-explicit-key" if api_base else None assert ( CometAPIEmbeddingConfig().get_complete_url( api_base=api_base, - api_key=None, + api_key=request_api_key, model="text-embedding-3-small", optional_params={}, litellm_params={}, @@ -88,12 +89,22 @@ def test_cometapi_key_and_base_precedence(monkeypatch): assert get_cometapi_api_key("explicit-key") == "explicit-key" assert get_cometapi_api_key() == "comet-env-key" - assert get_cometapi_api_base("https://explicit.example.com/v1") == ( - "https://explicit.example.com/v1" + assert ( + get_cometapi_api_base("https://explicit.example.com/v1", api_key="explicit-key") + == "https://explicit.example.com/v1" ) assert get_cometapi_api_base() == "https://proxy.example.com/openai/v1" +def test_cometapi_custom_base_requires_request_key(monkeypatch): + _clear_cometapi_env(monkeypatch) + _pollute_openai_globals(monkeypatch) + monkeypatch.setenv("COMETAPI_KEY", "comet-env-key") + + with pytest.raises(ValueError, match="api_base requires an explicit api_key"): + get_cometapi_api_base("https://attacker.example.com/v1") + + def test_cometapi_key_and_base_ignore_litellm_globals(monkeypatch): _clear_cometapi_env(monkeypatch) _pollute_openai_globals(monkeypatch) @@ -106,7 +117,11 @@ def test_cometapi_complete_url_preserves_existing_v1_path(monkeypatch): _clear_cometapi_env(monkeypatch) assert ( - get_cometapi_complete_url("https://proxy.example.com/openai/v1", "embeddings") + get_cometapi_complete_url( + "https://proxy.example.com/openai/v1", + "embeddings", + api_key="comet-explicit-key", + ) == "https://proxy.example.com/openai/v1/embeddings" ) @@ -127,7 +142,7 @@ def test_cometapi_complete_url_rejects_non_v1_paths(api_base, monkeypatch): ValueError, match="CometAPI OpenAI-compatible endpoints require a /v1 api_base", ): - get_cometapi_complete_url(api_base, "embeddings") + get_cometapi_complete_url(api_base, "embeddings", api_key="comet-explicit-key") @pytest.mark.parametrize( @@ -145,7 +160,11 @@ def test_cometapi_complete_url_rejects_invalid_endpoints(endpoint, monkeypatch): _clear_cometapi_env(monkeypatch) with pytest.raises(ValueError, match="CometAPI endpoint must be a non-empty path"): - get_cometapi_complete_url("https://api.cometapi.com/v1", endpoint) + get_cometapi_complete_url( + "https://api.cometapi.com/v1", + endpoint, + api_key="comet-explicit-key", + ) @pytest.mark.parametrize( @@ -163,7 +182,7 @@ def test_cometapi_complete_url_rejects_api_base_query_or_fragment( with pytest.raises( ValueError, match="CometAPI api_base must not include query or fragment" ): - get_cometapi_complete_url(api_base, "embeddings") + get_cometapi_complete_url(api_base, "embeddings", api_key="comet-explicit-key") def test_cometapi_embedding_uses_api_key_alias(monkeypatch): @@ -242,7 +261,7 @@ def test_cometapi_image_generation_url_normalization(): assert ( config.get_complete_url( api_base="https://api.cometapi.com/v1", - api_key=None, + api_key="comet-explicit-key", model="gpt-image-2", optional_params={}, litellm_params={}, @@ -252,7 +271,7 @@ def test_cometapi_image_generation_url_normalization(): assert ( config.get_complete_url( api_base="https://api.cometapi.com/v1/images/generations", - api_key=None, + api_key="comet-explicit-key", model="gpt-image-2", optional_params={}, litellm_params={}, @@ -309,6 +328,27 @@ def test_cometapi_image_generation_missing_key_does_not_call_handler(monkeypatch mock_image_handler.assert_not_called() +def test_cometapi_image_generation_custom_base_requires_request_key(monkeypatch): + _clear_cometapi_env(monkeypatch) + _pollute_openai_globals(monkeypatch) + monkeypatch.setenv("COMETAPI_KEY", "comet-env-key") + + with patch( + "litellm.images.main.llm_http_handler.image_generation_handler" + ) as mock_image_handler: + with pytest.raises( + litellm.APIConnectionError, + match="api_base requires an explicit api_key", + ): + image_generation( + model="cometapi/gpt-image-2", + prompt="A small comet over a clean API diagram", + api_base="https://attacker.example.com/v1", + ) + + mock_image_handler.assert_not_called() + + def test_cometapi_image_generation_maps_new_openai_params(monkeypatch): _clear_cometapi_env(monkeypatch) _pollute_openai_globals(monkeypatch) @@ -320,6 +360,7 @@ def test_cometapi_image_generation_maps_new_openai_params(monkeypatch): image_generation( model="cometapi/gpt-image-2", prompt="A small comet over a clean API diagram", + api_base="https://proxy.example.com/openai/v1", api_key="comet-explicit-key", output_compression=70, output_format="png", @@ -329,7 +370,9 @@ def test_cometapi_image_generation_maps_new_openai_params(monkeypatch): call_kwargs = mock_image_handler.call_args.kwargs assert call_kwargs["api_key"] == "comet-explicit-key" assert call_kwargs["custom_llm_provider"] == "cometapi" - assert call_kwargs["litellm_params"]["api_base"] is None + assert call_kwargs["litellm_params"]["api_base"] == ( + "https://proxy.example.com/openai/v1" + ) assert "api_key" not in call_kwargs["litellm_params"] assert ( call_kwargs["image_generation_optional_request_params"]["output_compression"] @@ -522,6 +565,23 @@ def test_cometapi_speech_missing_key_does_not_call_openai_audio(monkeypatch): mock_speech.assert_not_called() +def test_cometapi_speech_custom_base_requires_request_key(monkeypatch): + _clear_cometapi_env(monkeypatch) + _pollute_openai_globals(monkeypatch) + monkeypatch.setenv("COMETAPI_KEY", "comet-env-key") + + with patch("litellm.main.openai_chat_completions.audio_speech") as mock_speech: + with pytest.raises(ValueError, match="api_base requires an explicit api_key"): + speech( + model="cometapi/tts-1", + input="hello", + voice="alloy", + api_base="https://attacker.example.com/v1", + ) + + mock_speech.assert_not_called() + + def test_cometapi_speech_requires_voice(monkeypatch): _clear_cometapi_env(monkeypatch) @@ -553,6 +613,27 @@ def test_cometapi_transcription_uses_cometapi_key_and_base(monkeypatch): assert call_kwargs["api_base"] == "https://api.cometapi.com/v1" +def test_cometapi_transcription_custom_base_requires_request_key(monkeypatch): + _clear_cometapi_env(monkeypatch) + _pollute_openai_globals(monkeypatch) + monkeypatch.setenv("COMETAPI_KEY", "comet-env-key") + + audio_file = io.BytesIO(b"not-real-audio") + audio_file.name = "sample.wav" + + with patch( + "litellm.main.openai_audio_transcriptions.audio_transcriptions" + ) as mock_transcription: + with pytest.raises(ValueError, match="api_base requires an explicit api_key"): + transcription( + model="cometapi/whisper-1", + file=audio_file, + api_base="https://attacker.example.com/v1", + ) + + mock_transcription.assert_not_called() + + def test_openai_transcription_fallback_is_unchanged(monkeypatch): _clear_cometapi_env(monkeypatch) _pollute_openai_globals(monkeypatch) @@ -701,6 +782,23 @@ def test_cometapi_moderation_missing_key_does_not_create_openai_client(monkeypat assert _OpenAIClient.instances == [] +def test_cometapi_moderation_custom_base_requires_request_key(monkeypatch): + _clear_cometapi_env(monkeypatch) + _pollute_openai_globals(monkeypatch) + monkeypatch.setenv("COMETAPI_KEY", "comet-env-key") + + _OpenAIClient.instances = [] + with patch("litellm.main.openai.OpenAI", _OpenAIClient): + with pytest.raises(ValueError, match="api_base requires an explicit api_key"): + moderation( + model="cometapi/omni-moderation-latest", + input="hello", + api_base="https://attacker.example.com/v1", + ) + + assert _OpenAIClient.instances == [] + + def test_openai_moderation_fallback_is_unchanged(monkeypatch): _clear_cometapi_env(monkeypatch) _pollute_openai_globals(monkeypatch)