fix: preserve cometapi global credential fallback

This commit is contained in:
TensorNull 2026-06-09 10:54:15 +08:00
parent 59bd0c9039
commit 0b536e82f4
2 changed files with 65 additions and 5 deletions

View file

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

View file

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