diff --git a/litellm/llms/openai/workload_identity.py b/litellm/llms/openai/workload_identity.py index 283fdfb92c2..369b4f1e3f7 100644 --- a/litellm/llms/openai/workload_identity.py +++ b/litellm/llms/openai/workload_identity.py @@ -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 diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 0cff558985b..923a6cc5743 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -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, diff --git a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index cad065c64d1..ac010cd90a0 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -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.""" diff --git a/tests/unit/proxy/test_openai_ws_passthrough_routes.py b/tests/unit/proxy/test_openai_ws_passthrough_routes.py index 7d79192b884..6a9cd972dc2 100644 --- a/tests/unit/proxy/test_openai_ws_passthrough_routes.py +++ b/tests/unit/proxy/test_openai_ws_passthrough_routes.py @@ -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")