diff --git a/litellm/llms/cometapi/common_utils.py b/litellm/llms/cometapi/common_utils.py index eb2b8a35567..b99dc94497a 100644 --- a/litellm/llms/cometapi/common_utils.py +++ b/litellm/llms/cometapi/common_utils.py @@ -8,16 +8,25 @@ DEFAULT_COMETAPI_API_BASE = "https://api.cometapi.com/v1" def get_cometapi_api_key(api_key: Optional[str] = None) -> Optional[str]: + import litellm + return ( - api_key or get_secret_str("COMETAPI_KEY") or get_secret_str("COMETAPI_API_KEY") + api_key + or litellm.cometapi_key + or get_secret_str("COMETAPI_KEY") + or get_secret_str("COMETAPI_API_KEY") + or litellm.api_key ) def get_cometapi_api_base(api_base: Optional[str] = None) -> str: + import litellm + return ( api_base or get_secret_str("COMETAPI_BASE_URL") or get_secret_str("COMETAPI_API_BASE") + or litellm.api_base or DEFAULT_COMETAPI_API_BASE ) @@ -26,9 +35,18 @@ def get_cometapi_complete_url(api_base: Optional[str], endpoint: str) -> str: base_url = get_cometapi_api_base(api_base).rstrip("/") normalized_endpoint = endpoint.strip("/") parsed_endpoint = urlsplit(normalized_endpoint) - if not normalized_endpoint or parsed_endpoint.query or parsed_endpoint.fragment: + if ( + not normalized_endpoint + or parsed_endpoint.scheme + or parsed_endpoint.netloc + or parsed_endpoint.query + or parsed_endpoint.fragment + or ".." in normalized_endpoint.split("/") + ): raise ValueError("CometAPI endpoint must be a non-empty path") parsed_base_url = urlsplit(base_url) + if parsed_base_url.query or parsed_base_url.fragment: + raise ValueError("CometAPI api_base must not include query or fragment") path_segments = [segment for segment in parsed_base_url.path.split("/") if segment] endpoint_segments = normalized_endpoint.split("/") invalid_version_segments = [ @@ -57,8 +75,8 @@ def get_cometapi_complete_url(api_base: Optional[str], endpoint: str) -> str: parsed_base_url.scheme, parsed_base_url.netloc, complete_path, - parsed_base_url.query, - parsed_base_url.fragment, + "", + "", ) ) diff --git a/tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py b/tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py index b6469aacda3..8fd965c6254 100644 --- a/tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py +++ b/tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py @@ -82,6 +82,7 @@ def test_cometapi_embedding_url_normalization(api_base, expected, monkeypatch): def test_cometapi_key_and_base_precedence(monkeypatch): _clear_cometapi_env(monkeypatch) + _pollute_openai_globals(monkeypatch) monkeypatch.setenv("COMETAPI_KEY", "comet-env-key") monkeypatch.setenv("COMETAPI_BASE_URL", "https://proxy.example.com/openai/v1") @@ -93,6 +94,14 @@ def test_cometapi_key_and_base_precedence(monkeypatch): assert get_cometapi_api_base() == "https://proxy.example.com/openai/v1" +def test_cometapi_key_and_base_fall_back_to_litellm_globals(monkeypatch): + _clear_cometapi_env(monkeypatch) + _pollute_openai_globals(monkeypatch) + + assert get_cometapi_api_key() == "openai-global-key" + assert get_cometapi_api_base() == "https://openai.invalid/v1" + + def test_cometapi_complete_url_preserves_existing_v1_path(monkeypatch): _clear_cometapi_env(monkeypatch) @@ -121,7 +130,17 @@ def test_cometapi_complete_url_rejects_non_v1_paths(api_base, monkeypatch): get_cometapi_complete_url(api_base, "embeddings") -@pytest.mark.parametrize("endpoint", ["", "/", "embeddings?foo=bar", "embeddings#frag"]) +@pytest.mark.parametrize( + "endpoint", + [ + "", + "/", + "embeddings?foo=bar", + "embeddings#frag", + "http://evil.test/x", + "../embeddings", + ], +) def test_cometapi_complete_url_rejects_invalid_endpoints(endpoint, monkeypatch): _clear_cometapi_env(monkeypatch) @@ -129,6 +148,24 @@ def test_cometapi_complete_url_rejects_invalid_endpoints(endpoint, monkeypatch): get_cometapi_complete_url("https://api.cometapi.com/v1", endpoint) +@pytest.mark.parametrize( + "api_base", + [ + "https://api.cometapi.com/v1?tenant=test", + "https://api.cometapi.com/v1#fragment", + ], +) +def test_cometapi_complete_url_rejects_api_base_query_or_fragment( + api_base, monkeypatch +): + _clear_cometapi_env(monkeypatch) + + with pytest.raises( + ValueError, match="CometAPI api_base must not include query or fragment" + ): + get_cometapi_complete_url(api_base, "embeddings") + + def test_cometapi_embedding_uses_api_key_alias(monkeypatch): _clear_cometapi_env(monkeypatch) monkeypatch.setenv("COMETAPI_API_KEY", "comet-env-key") @@ -414,6 +451,7 @@ def test_cometapi_speech_uses_cometapi_key_and_base(monkeypatch): _clear_cometapi_env(monkeypatch) _pollute_openai_globals(monkeypatch) monkeypatch.setenv("COMETAPI_API_KEY", "comet-env-key") + monkeypatch.setenv("COMETAPI_BASE_URL", "https://api.cometapi.com/v1") with patch("litellm.main.openai_chat_completions.audio_speech") as mock_speech: mock_speech.return_value = b"audio" @@ -451,6 +489,7 @@ def test_cometapi_speech_requires_voice(monkeypatch): def test_cometapi_transcription_uses_cometapi_key_and_base(monkeypatch): _clear_cometapi_env(monkeypatch) _pollute_openai_globals(monkeypatch) + monkeypatch.setenv("COMETAPI_BASE_URL", "https://api.cometapi.com/v1") audio_file = io.BytesIO(b"not-real-audio") audio_file.name = "sample.wav" @@ -565,6 +604,7 @@ class _OpenAIClient: def test_cometapi_moderation_uses_cometapi_key_and_base(monkeypatch): _clear_cometapi_env(monkeypatch) _pollute_openai_globals(monkeypatch) + monkeypatch.setenv("COMETAPI_BASE_URL", "https://api.cometapi.com/v1") _OpenAIClient.instances = [] with patch("litellm.main.openai.OpenAI", _OpenAIClient): @@ -586,6 +626,7 @@ def test_cometapi_moderation_uses_cometapi_key_and_base(monkeypatch): def test_cometapi_moderation_accepts_bare_model_with_custom_provider(monkeypatch): _clear_cometapi_env(monkeypatch) _pollute_openai_globals(monkeypatch) + monkeypatch.setenv("COMETAPI_BASE_URL", "https://api.cometapi.com/v1") _OpenAIClient.instances = [] with patch("litellm.main.openai.OpenAI", _OpenAIClient): @@ -653,6 +694,7 @@ class _AsyncOpenAIClient: async def test_cometapi_amoderation_uses_cometapi_key_and_base(monkeypatch): _clear_cometapi_env(monkeypatch) _pollute_openai_globals(monkeypatch) + monkeypatch.setenv("COMETAPI_BASE_URL", "https://api.cometapi.com/v1") fake_client = _AsyncOpenAIClient() with patch(