fix(openai): mint workload identity tokens for PrivateLink and regional api.openai.com hosts

This commit is contained in:
mateo-berri 2026-09-03 14:38:47 -07:00
parent 8699998c9e
commit 19da217167
2 changed files with 45 additions and 3 deletions

View file

@ -8,7 +8,7 @@ from urllib.parse import urlparse
import litellm
from litellm.secret_managers.main import get_secret_str, normalize_nonempty_secret_str
from .common_utils import OpenAIError
from .common_utils import OpenAIError, is_openai_backed_api_base
if TYPE_CHECKING:
from collections.abc import Callable
@ -16,7 +16,6 @@ if TYPE_CHECKING:
from openai.auth import SubjectTokenProvider, WorkloadIdentity, WorkloadIdentityAuth
OPENAI_WIF_CLIENT_ID: Final = "litellm"
_OPENAI_API_HOST: Final = "api.openai.com"
_SDK_UPGRADE_MESSAGE: Final = (
"OpenAI workload identity federation requires openai>=2.32.0. "
"Upgrade the installed openai package to use OPENAI_IDENTITY_PROVIDER_ID / "
@ -75,7 +74,7 @@ def _targets_openai_api(api_base: str | None) -> bool:
if api_base is None:
return True
parsed: Final = urlparse(api_base)
return parsed.scheme == "https" and parsed.hostname == _OPENAI_API_HOST
return parsed.scheme == "https" and is_openai_backed_api_base(api_base)
@lru_cache(maxsize=16)

View file

@ -85,6 +85,31 @@ class TestResolveConfig:
def test_plaintext_http_api_base_disables(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
assert resolve_openai_workload_identity_config(api_key=None, api_base="http://api.openai.com/v1") is None
@pytest.mark.parametrize(
"api_base",
(
"https://southcentralus.privatelink.api.openai.com/v1",
"https://eu.api.openai.com/v1",
"https://us.api.openai.com/v1",
),
)
def test_openai_backed_api_base_allows(self, wif_env: OpenAIWorkloadIdentityConfig, api_base: str) -> None:
assert resolve_openai_workload_identity_config(api_key=None, api_base=api_base) == wif_env
@pytest.mark.parametrize(
"api_base",
(
"https://api.openai.com.evil.example/v1",
"https://openai.com/v1",
"https://euapi.openai.com/v1",
"http://southcentralus.privatelink.api.openai.com/v1",
),
)
def test_lookalike_or_plaintext_api_base_disables(
self, wif_env: OpenAIWorkloadIdentityConfig, api_base: str
) -> None:
assert resolve_openai_workload_identity_config(api_key=None, api_base=api_base) is None
def test_foreign_env_base_url_disables(
self, wif_env: OpenAIWorkloadIdentityConfig, monkeypatch: pytest.MonkeyPatch
) -> None:
@ -158,6 +183,14 @@ class TestClientConstruction:
assert client.api_key == "workload-identity-auth"
assert client._workload_identity_auth is not None
def test_privatelink_client_uses_workload_identity(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
client: Final = OpenAIChatCompletion()._get_openai_client(
is_async=False, api_key=None, api_base="https://southcentralus.privatelink.api.openai.com/v1"
)
assert isinstance(client, OpenAI)
assert client.api_key == "workload-identity-auth"
assert client._workload_identity_auth is not None
def test_static_key_client_unaffected(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
client: Final = OpenAIChatCompletion()._get_openai_client(is_async=False, api_key="sk-static", api_base=None)
assert isinstance(client, OpenAI)
@ -231,6 +264,16 @@ class TestResponsesValidateEnvironment:
)
assert headers["Authorization"] == "Bearer None"
@respx.mock
def test_privatelink_api_base_mints_bearer(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
mock_token_exchange()
headers: Final = OpenAIResponsesAPIConfig().validate_environment(
headers={},
model="gpt-4o-mini",
litellm_params=GenericLiteLLMParams(api_base="https://southcentralus.privatelink.api.openai.com/v1"),
)
assert headers["Authorization"] == "Bearer exchanged-bearer-token"
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()