fix(proxy): use OpenAI workload identity federation tokens on /openai_passthrough (#44140)

* fix(proxy): use OpenAI workload identity federation tokens on /openai_passthrough

The OpenAI passthrough routes (HTTP and websocket) only looked up a static
OpenAI API key, so proxies authenticating to OpenAI through workload identity
federation failed with "Required 'OPENAI_API_KEY'". When no non-empty static
key is configured, exchange the workload identity subject token for a bearer
token, scoped to the passthrough's own OPENAI_API_BASE so the token is never
sent to a non-OpenAI host

* fix(proxy): close the OpenAI websocket passthrough cleanly when the workload identity exchange fails

A rejected, unreachable or unreadable workload identity exchange raised out of
the websocket route before the handshake was accepted, so clients saw a bare
handshake failure. Log the cause and close with 1011 and a fixed reason instead

* refactor(openai): resolve workload identity bearer tokens for an api base in the OpenAI provider module

The passthrough route only owns the static key lookup now and asks the OpenAI
provider module for a workload identity bearer token scoped to its api base

---------

Co-authored-by: Krrish Dholakia <krrish+github@berri.ai>
This commit is contained in:
Matthew Lapointe 2026-10-03 00:07:09 -04:00 • committed by GitHub
parent 80a2f4d8a8
commit ad566c90dd
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 162 additions and 10 deletions

View file

@ -70,6 +70,13 @@ def get_workload_identity_bearer_token(config: OpenAIWorkloadIdentityConfig) ->
return _workload_identity_auth(config).get_token()
async def get_workload_identity_bearer_token_for_api_base(api_base: str) -> str | None:
config: Final = resolve_openai_workload_identity_config(api_key=None, api_base=api_base)
if config is None:
return None
return await _workload_identity_auth(config).get_token_async()
def _targets_openai_api(api_base: str | None) -> bool:
if api_base is None:
return True

View file

@ -22,6 +22,7 @@ from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
import httpx
import openai
from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket
from fastapi.responses import StreamingResponse
from pydantic import ConfigDict, TypeAdapter
@ -59,6 +60,8 @@ from litellm.llms.deepgram.common_utils import (
)
from litellm.llms.fal_ai.cost_calculator import fal_ai_passthrough_cost, fal_ai_queue_base
from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path
from litellm.llms.openai.common_utils import OpenAIError as LiteLLMOpenAIError
from litellm.llms.openai.workload_identity import get_workload_identity_bearer_token_for_api_base
from litellm.llms.oss_decision import OssDecisionProvider, oss_connection, validate_oss_request
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
@ -101,7 +104,7 @@ from litellm.proxy.vector_store_endpoints.utils import (
get_litellm_managed_vector_store,
is_allowed_to_call_vector_store_endpoint,
)
from litellm.secret_managers.main import get_secret_str, str_to_bool
from litellm.secret_managers.main import get_secret_str, normalize_nonempty_secret_str, str_to_bool
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY,
@ -2978,6 +2981,21 @@ async def vertex_proxy_route(
)
_OPENAI_WS_TOKEN_EXCHANGE_FAILED_REASON: Final = "OpenAI workload identity token exchange failed"
async def _openai_passthrough_credential(base_target_url: str) -> str | None:
static_api_key: Final = normalize_nonempty_secret_str(
passthrough_endpoint_router.get_credentials(
custom_llm_provider=litellm.LlmProviders.OPENAI.value,
region_name=None,
)
)
if static_api_key is not None:
return static_api_key
return await get_workload_identity_bearer_token_for_api_base(base_target_url)
@router.api_route(
"/openai/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
@ -3013,11 +3031,7 @@ async def openai_proxy_route(
[Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)
"""
base_target_url: Final = os.getenv("OPENAI_API_BASE") or "https://api.openai.com/"
# Add or update query parameters
openai_api_key: Final = passthrough_endpoint_router.get_credentials(
custom_llm_provider=litellm.LlmProviders.OPENAI.value,
region_name=None,
)
openai_api_key: Final = await _openai_passthrough_credential(base_target_url)
if openai_api_key is None:
raise Exception("Required 'OPENAI_API_KEY' in environment to make pass-through calls to OpenAI.")
@ -3185,10 +3199,12 @@ async def openai_websocket_proxy_route(
return
base_target_url: Final = os.getenv("OPENAI_API_BASE") or "https://api.openai.com/"
openai_api_key: Final = passthrough_endpoint_router.get_credentials(
custom_llm_provider=litellm.LlmProviders.OPENAI.value,
region_name=None,
)
try:
openai_api_key: Final = await _openai_passthrough_credential(base_target_url)
except (openai.OpenAIError, httpx.HTTPError, LiteLLMOpenAIError):
verbose_proxy_logger.exception("OpenAI workload identity token exchange failed for websocket passthrough")
await websocket.close(code=1011, reason=_OPENAI_WS_TOKEN_EXCHANGE_FAILED_REASON)
return
if openai_api_key is None:
await websocket.close(
code=1011,

View file

@ -3585,6 +3585,72 @@ def test_openai_passthrough_forwards_verbatim_to_openai(
assert route.calls.last.request.headers["authorization"] == "Bearer sk-upstream"
@pytest.fixture
def openai_wif_env(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None:
from litellm.llms.openai.workload_identity import _workload_identity_auth
token_file: Final = tmp_path / "subject_token.jwt"
token_file.write_text("subject-token-from-file")
monkeypatch.delenv("OPENAI_API_BASE", raising=False)
monkeypatch.delenv("OPENAI_BASE_URL", 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))
_workload_identity_auth.cache_clear()
@pytest.mark.parametrize("static_key", [None, "", " "])
def test_openai_passthrough_uses_workload_identity_token_without_static_key(
openai_passthrough_client: TestClient,
openai_wif_env: None,
monkeypatch: pytest.MonkeyPatch,
static_key: str | None,
) -> None:
if static_key is None:
monkeypatch.delenv("OPENAI_API_KEY")
else:
monkeypatch.setenv("OPENAI_API_KEY", static_key)
with respx.mock(assert_all_called=True) as upstream:
token_exchange = upstream.post("https://auth.openai.com/oauth/token").mock(
return_value=httpx.Response(200, json={"access_token": "wif-bearer", "expires_in": 3600})
)
route = upstream.post("https://api.openai.com/v1/responses").mock(
return_value=httpx.Response(200, json={"id": "upstream_123"})
)
response = openai_passthrough_client.post(
"/openai_passthrough/v1/responses", json={"model": "gpt-5.1", "input": "hi"}
)
assert (response.status_code, response.json()) == (200, {"id": "upstream_123"})
assert route.calls.last.request.headers["authorization"] == "Bearer wif-bearer"
assert json.loads(token_exchange.calls.last.request.content)["subject_token"] == "subject-token-from-file"
@pytest.mark.asyncio
async def test_openai_passthrough_never_sends_workload_identity_token_to_foreign_api_base(
openai_wif_env: None, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
monkeypatch.setenv("OPENAI_API_BASE", "https://my-vllm.internal/")
monkeypatch.setenv("OPENAI_BASE_URL", "https://api.openai.com/v1")
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value=None,
),
respx.mock(assert_all_mocked=True) as upstream,
pytest.raises(Exception, match="Required 'OPENAI_API_KEY'"),
):
await openai_proxy_route(
endpoint="v1/responses",
request=MagicMock(spec=Request),
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=MagicMock(),
)
assert upstream.calls.call_count == 0
class TestCursorProxyRoute:
"""Tests for the Cursor Cloud Agents pass-through route."""

View file

@ -7,9 +7,13 @@ from types import MappingProxyType, SimpleNamespace
from typing import Final
from unittest.mock import patch
import httpx
import pytest
import respx
from starlette.routing import WebSocketRoute
import litellm
from litellm.llms.openai.workload_identity import _workload_identity_auth
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
_OPENAI_WS_DISABLED_REFUSAL,
@ -174,6 +178,65 @@ async def test_openai_websocket_accepts_first_client_subprotocol():
assert websocket.closed is None
TOKEN_EXCHANGE_URL: Final = "https://auth.openai.com/oauth/token"
@pytest.fixture
def openai_wif_token_file(monkeypatch, tmp_path):
token_file = tmp_path / "subject_token.jwt"
token_file.write_text("subject-token-from-file")
monkeypatch.delenv("OPENAI_API_BASE", raising=False)
monkeypatch.delenv("OPENAI_BASE_URL", 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))
_workload_identity_auth.cache_clear()
return token_file
@pytest.mark.asyncio
async def test_openai_websocket_uses_workload_identity_token_without_static_key(openai_wif_token_file):
websocket = _FakeWebSocket("/openai_passthrough/v1/realtime", "model=gpt-realtime")
with patch(GET_CREDENTIALS, return_value=None), respx.mock(assert_all_called=True) as upstream:
upstream.post(TOKEN_EXCHANGE_URL).mock(
return_value=httpx.Response(200, json={"access_token": "wif-bearer", "expires_in": 3600})
)
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), ENABLED)
assert [call.custom_headers for call in served.relay.calls] == [
MappingProxyType({"Authorization": "Bearer wif-bearer"})
]
assert websocket.closed is None
@pytest.mark.asyncio
@pytest.mark.parametrize(
"subject_token_present, exchange_outcome",
[
(True, httpx.Response(401, json={"error": "invalid_grant"})),
(True, httpx.ConnectError("auth.openai.com unreachable")),
(False, httpx.Response(200, json={"access_token": "wif-bearer", "expires_in": 3600})),
],
ids=["rejected", "unreachable", "missing_subject_token"],
)
async def test_openai_websocket_closes_cleanly_when_workload_identity_exchange_fails(
openai_wif_token_file, subject_token_present, exchange_outcome
):
if not subject_token_present:
openai_wif_token_file.unlink()
websocket = _FakeWebSocket("/openai_passthrough/v1/realtime", "model=gpt-realtime")
with patch(GET_CREDENTIALS, return_value=None), respx.mock(assert_all_called=False) as upstream:
upstream.post(TOKEN_EXCHANGE_URL).mock(side_effect=exchange_outcome)
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), ENABLED)
assert websocket.closed == (1011, "OpenAI workload identity token exchange failed")
assert websocket.accepts == []
assert served.relay.calls == []
@pytest.mark.asyncio
async def test_openai_websocket_closes_cleanly_when_provider_credentials_missing():
websocket = _FakeWebSocket("/openai/v1/realtime", "model=gpt-4o-realtime-preview")