fix(openai): scope workload identity to the openai provider and env-resolved base/key

This commit is contained in:
mateo-berri 2026-08-31 12:15:52 -07:00
parent ae945f4fa3
commit 72adeda9ce
3 changed files with 44 additions and 8 deletions

View file

@ -393,12 +393,10 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
litellm_params = litellm_params or GenericLiteLLMParams()
api_key = litellm_params.api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY")
headers.setdefault("Content-Type", "application/json")
workload_identity_config: Final = resolve_openai_workload_identity_config(
api_key=api_key,
api_base=litellm_params.api_base
or litellm.api_base
or get_secret_str("OPENAI_BASE_URL")
or get_secret_str("OPENAI_API_BASE"),
workload_identity_config: Final = (
resolve_openai_workload_identity_config(api_key=api_key, api_base=litellm_params.api_base)
if self.custom_llm_provider is LlmProviders.OPENAI
else None
)
if workload_identity_config is not None:
headers["Authorization"] = f"Bearer {get_workload_identity_bearer_token(workload_identity_config)}"

View file

@ -5,6 +5,7 @@ from functools import lru_cache
from typing import TYPE_CHECKING, Final
from urllib.parse import urlparse
import litellm
from litellm.secret_managers.main import get_secret_str
from .common_utils import OpenAIError
@ -44,9 +45,12 @@ def resolve_openai_workload_identity_config(
api_key: str | None,
api_base: str | None,
) -> OpenAIWorkloadIdentityConfig | None:
if api_key is not None:
if api_key is not None or get_secret_str("OPENAI_API_KEY") is not None:
return None
if not _targets_openai_api(api_base):
effective_api_base: Final = (
api_base or litellm.api_base or get_secret_str("OPENAI_BASE_URL") or get_secret_str("OPENAI_API_BASE")
)
if not _targets_openai_api(effective_api_base):
return None
identity_provider_id: Final = get_secret_str("OPENAI_IDENTITY_PROVIDER_ID")
service_account_id: Final = get_secret_str("OPENAI_SERVICE_ACCOUNT_ID")

View file

@ -9,6 +9,7 @@ import respx
from openai import AsyncOpenAI, OpenAI
import litellm
from litellm.llms.litellm_proxy.responses.transformation import LiteLLMProxyResponsesAPIConfig
from litellm.llms.openai.common_utils import BaseOpenAILLM, OpenAIError
from litellm.llms.openai.openai import OpenAIChatCompletion
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
@ -28,6 +29,9 @@ def wif_env(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> OpenAIWorkloadId
token_file: Final = tmp_path / "subject_token.jwt"
token_file.write_text("subject-token-from-file")
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
monkeypatch.delenv("OPENAI_API_BASE", raising=False)
monkeypatch.setattr(litellm, "api_base", None)
monkeypatch.setenv("OPENAI_IDENTITY_PROVIDER_ID", "idp_test123")
monkeypatch.setenv("OPENAI_SERVICE_ACCOUNT_ID", "user-test456")
monkeypatch.setenv("OPENAI_IDENTITY_TOKEN_FILE", str(token_file))
@ -53,12 +57,36 @@ class TestResolveConfig:
def test_static_api_key_wins(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
assert resolve_openai_workload_identity_config(api_key="sk-static", api_base=None) is None
def test_env_openai_api_key_wins(
self, wif_env: OpenAIWorkloadIdentityConfig, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setenv("OPENAI_API_KEY", "sk-from-env")
assert resolve_openai_workload_identity_config(api_key=None, api_base=None) is None
def test_foreign_api_base_disables(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
assert resolve_openai_workload_identity_config(api_key=None, api_base="https://my-vllm.internal/v1") is None
def test_openai_api_base_allows(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
assert resolve_openai_workload_identity_config(api_key=None, api_base="https://api.openai.com/v1") == wif_env
def test_foreign_env_base_url_disables(
self, wif_env: OpenAIWorkloadIdentityConfig, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setenv("OPENAI_BASE_URL", "https://my-vllm.internal/v1")
assert resolve_openai_workload_identity_config(api_key=None, api_base=None) is None
def test_openai_env_base_url_allows(
self, wif_env: OpenAIWorkloadIdentityConfig, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setenv("OPENAI_BASE_URL", "https://api.openai.com/v1")
assert resolve_openai_workload_identity_config(api_key=None, api_base=None) == wif_env
def test_foreign_litellm_api_base_disables(
self, wif_env: OpenAIWorkloadIdentityConfig, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(litellm, "api_base", "https://my-vllm.internal/v1")
assert resolve_openai_workload_identity_config(api_key=None, api_base=None) is None
@pytest.mark.parametrize(
"missing_var",
["OPENAI_IDENTITY_PROVIDER_ID", "OPENAI_SERVICE_ACCOUNT_ID", "OPENAI_IDENTITY_TOKEN_FILE"],
@ -186,3 +214,9 @@ class TestResponsesValidateEnvironment:
litellm_params=GenericLiteLLMParams(api_base="https://my-vllm.internal/v1"),
)
assert headers["Authorization"] == "Bearer None"
def test_litellm_proxy_subclass_never_mints_wif(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
headers: Final = LiteLLMProxyResponsesAPIConfig().validate_environment(
headers={}, model="gpt-4o-mini", litellm_params=GenericLiteLLMParams()
)
assert headers["Authorization"] == "Bearer None"