diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 1479c378014..eadc087383a 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -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)}" diff --git a/litellm/llms/openai/workload_identity.py b/litellm/llms/openai/workload_identity.py index 15105d67957..be4d015af8a 100644 --- a/litellm/llms/openai/workload_identity.py +++ b/litellm/llms/openai/workload_identity.py @@ -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") diff --git a/tests/test_litellm/llms/openai/test_openai_workload_identity.py b/tests/test_litellm/llms/openai/test_openai_workload_identity.py index b8cc8c9c80f..d7257d6af89 100644 --- a/tests/test_litellm/llms/openai/test_openai_workload_identity.py +++ b/tests/test_litellm/llms/openai/test_openai_workload_identity.py @@ -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"