From bf71e09ba11f2e0cab9e6f1794726f02b0dd7440 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 3 Sep 2026 12:18:01 -0700 Subject: [PATCH 1/3] feat(proxy): configure OpenAI workload identity federation from the Add Model and LLM Credentials forms --- .../litellm_core_utils/get_litellm_params.py | 13 +- litellm/llms/azure_ai/embed/handler.py | 2 + .../llms/openai/chat/gpt_transformation.py | 49 ++- litellm/llms/openai/openai.py | 16 +- .../llms/openai/responses/transformation.py | 7 +- litellm/llms/openai/workload_identity.py | 19 +- litellm/main.py | 1 + litellm/proxy/auth/auth_utils.py | 4 +- .../common_utils/credential_hydration.py | 16 +- .../proxy/credential_endpoints/endpoints.py | 8 +- .../model_management_endpoints.py | 6 +- .../provider_create_fields.json | 86 +++++ .../clientside_credential_handler.py | 10 +- litellm/types/router.py | 33 +- litellm/types/utils.py | 5 +- .../test_get_litellm_params.py | 38 ++ .../anthropic/test_anthropic_common_utils.py | 4 +- .../openai/test_openai_workload_identity.py | 332 ++++++++++++++++++ .../proxy/auth/test_auth_utils.py | 16 +- .../credential_endpoints/test_endpoints.py | 34 ++ .../test_model_management_endpoints.py | 119 ++++++- .../public_endpoints/test_public_endpoints.py | 51 ++- .../test_fallback_event_handlers.py | 17 + tests/test_litellm/types/test_router.py | 79 ++++- ui/litellm-dashboard/src/lib/http/schema.d.ts | 12 + 25 files changed, 881 insertions(+), 96 deletions(-) diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index 685fc1d915e..7fbab6e59ee 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -58,10 +58,20 @@ ANTHROPIC_WIF_KWARGS_KEYS: Final = frozenset( } ) +OPENAI_WIF_KWARGS_KEYS: Final = frozenset( + { + "openai_identity_provider_id", + "openai_service_account_id", + "openai_identity_token_file", + } +) + # Keys `completion()` forwards from its own kwargs into `get_litellm_params`, # which are otherwise invisible to it because that call site passes explicit # named arguments rather than `**kwargs`. -FORWARDED_KWARGS_KEYS: Final = AWS_CREDENTIAL_KWARGS_KEYS | ANTHROPIC_WIF_KWARGS_KEYS | frozenset({RUST_KWARG_KEY}) +FORWARDED_KWARGS_KEYS: Final = ( + AWS_CREDENTIAL_KWARGS_KEYS | ANTHROPIC_WIF_KWARGS_KEYS | OPENAI_WIF_KWARGS_KEYS | frozenset({RUST_KWARG_KEY}) +) # Pre-define optional kwargs keys as frozenset for O(1) lookups # These are extracted from kwargs only if present, avoiding unnecessary .get() calls @@ -100,6 +110,7 @@ OPTIONAL_KWARGS_KEYS: Final = ( ) | AWS_CREDENTIAL_KWARGS_KEYS | ANTHROPIC_WIF_KWARGS_KEYS + | OPENAI_WIF_KWARGS_KEYS ) # Backward-compatible alias for existing imports/tests. diff --git a/litellm/llms/azure_ai/embed/handler.py b/litellm/llms/azure_ai/embed/handler.py index 65c3997c099..bcdd060e7bf 100644 --- a/litellm/llms/azure_ai/embed/handler.py +++ b/litellm/llms/azure_ai/embed/handler.py @@ -1,3 +1,4 @@ +from collections.abc import Mapping from typing import Final from openai import OpenAI @@ -206,6 +207,7 @@ class AzureAIEmbedding(OpenAIChatCompletion): aembedding=None, max_retries: int | None = None, shared_session=None, + litellm_params: Mapping[str, object] | None = None, ) -> EmbeddingResponse: """ - Separate image url from text diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 9afc6331d96..5fdce1ed283 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -54,6 +54,7 @@ from litellm.types.utils import ( from litellm.utils import convert_to_model_response_object from ..common_utils import OpenAIError +from ..workload_identity import get_workload_identity_bearer_token, resolve_openai_workload_identity_config if TYPE_CHECKING: import tiktoken @@ -70,6 +71,11 @@ else: _NO_TOOLS_UPDATE: Final[Mapping[str, object]] = MappingProxyType({}) +def _litellm_params_str(litellm_params: Mapping[str, object] | None, key: str) -> str | None: + value: Final = litellm_params.get(key) if litellm_params is not None else None + return value if isinstance(value, str) else None + + class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): """ Reference: https://platform.openai.com/docs/api-reference/chat/create @@ -747,28 +753,37 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): """ Calls OpenAI's `/v1/models` endpoint and returns the list of models. """ + return self._fetch_model_ids(api_base=api_base, bearer_token=api_key or get_secret_str("OPENAI_API_KEY")) - if api_base is None: - api_base = "https://api.openai.com" - if api_key is None: - api_key = get_secret_str("OPENAI_API_KEY") - - # Strip api_base to just the base URL (scheme + host + port) - parsed_url: Final = httpx.URL(api_base) - base_url = f"{parsed_url.scheme}://{parsed_url.host}" - if parsed_url.port: - base_url += f":{parsed_url.port}" - - response: Final = litellm.module_level_client.get( - url=f"{base_url}/v1/models", - headers={"Authorization": f"Bearer {api_key}"}, + def discover_models( + self, litellm_params: Mapping[str, object] | None = None + ) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override + if type(self) is not OpenAIGPTConfig: + return super().discover_models(litellm_params) + api_key: Final = _litellm_params_str(litellm_params, "api_key") + api_base: Final = _litellm_params_str(litellm_params, "api_base") + workload_identity_config: Final = resolve_openai_workload_identity_config( + api_key=api_key, api_base=api_base, litellm_params=litellm_params + ) + if workload_identity_config is None: + return self.get_models(api_key=api_key, api_base=api_base) + return self._fetch_model_ids( + api_base=api_base, bearer_token=get_workload_identity_bearer_token(workload_identity_config) ) + @staticmethod + def _fetch_model_ids( + api_base: str | None, bearer_token: str | None + ) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override + parsed_url: Final = httpx.URL(api_base or "https://api.openai.com") + port_suffix: Final = f":{parsed_url.port}" if parsed_url.port else "" + response: Final = litellm.module_level_client.get( + url=f"{parsed_url.scheme}://{parsed_url.host}{port_suffix}/v1/models", + headers={"Authorization": f"Bearer {bearer_token}"}, + ) if response.status_code != 200: raise Exception(f"Failed to get models: {response.text}") - - models: Final = response.json()["data"] - return [model["id"] for model in models] + return [model["id"] for model in response.json()["data"]] @staticmethod def get_api_key(api_key: str | None = None) -> str | None: diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 1cfc6e06ee9..f5ad6f6a95d 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -382,8 +382,11 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): organization: str | None = None, client: OpenAI | AsyncOpenAI | None = None, shared_session: Optional["ClientSession"] = None, + litellm_params: Mapping[str, object] | None = None, ) -> OpenAI | AsyncOpenAI | None: - workload_identity_config: Final = resolve_openai_workload_identity_config(api_key=api_key, api_base=api_base) + workload_identity_config: Final = resolve_openai_workload_identity_config( + api_key=api_key, api_base=api_base, litellm_params=litellm_params + ) client_initialization_params: Final[dict] = locals() if client is None: if not isinstance(max_retries, int): @@ -773,6 +776,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): max_retries=max_retries, organization=organization, stream_options=stream_options, + litellm_params=litellm_params, ) else: if not isinstance(max_retries, int): @@ -786,6 +790,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): max_retries=max_retries, organization=organization, client=client, + litellm_params=litellm_params, ) ## LOGGING @@ -927,6 +932,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): organization=organization, client=client, shared_session=shared_session, + litellm_params=litellm_params, ) ## LOGGING @@ -1022,6 +1028,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): max_retries=None, headers=None, stream_options: dict | None = None, + litellm_params: Mapping[str, object] | None = None, ): data["stream"] = True data.update(self.get_stream_options(stream_options=stream_options, api_base=api_base)) @@ -1035,6 +1042,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): max_retries=max_retries, organization=organization, client=client, + litellm_params=litellm_params, ) ## LOGGING logging_obj.pre_call( @@ -1107,6 +1115,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): organization=organization, client=client, shared_session=shared_session, + litellm_params=litellm_params, ) ## LOGGING logging_obj.pre_call( @@ -1243,6 +1252,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): client: AsyncOpenAI | None = None, max_retries=None, shared_session: Optional["ClientSession"] = None, + litellm_params: Mapping[str, object] | None = None, ): try: openai_aclient: Final[AsyncOpenAI] = self._get_openai_client( @@ -1253,6 +1263,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): max_retries=max_retries, client=client, shared_session=shared_session, + litellm_params=litellm_params, ) raw_response: Final = await self.make_openai_embedding_request( openai_aclient=openai_aclient, @@ -1316,6 +1327,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): aembedding=None, max_retries: int | None = None, shared_session: Optional["ClientSession"] = None, + litellm_params: Mapping[str, object] | None = None, ) -> EmbeddingResponse: super().embedding() try: @@ -1342,6 +1354,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): client=client, max_retries=max_retries, shared_session=shared_session, + litellm_params=litellm_params, ) openai_client: Final[OpenAI] = self._get_openai_client( @@ -1351,6 +1364,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): timeout=timeout, max_retries=max_retries, client=client, + litellm_params=litellm_params, ) ## embedding CALL diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 01313e95878..a30458aca73 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -10,6 +10,7 @@ from typing_extensions import ReadOnly, TypedDict import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import process_response_headers +from litellm.litellm_core_utils.get_litellm_params import OPENAI_WIF_KWARGS_KEYS from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( _safe_convert_created_field, ) @@ -445,7 +446,11 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): 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) + resolve_openai_workload_identity_config( + api_key=api_key, + api_base=litellm_params.api_base, + litellm_params=litellm_params.model_dump(include=set(OPENAI_WIF_KWARGS_KEYS)), + ) if self.custom_llm_provider is LlmProviders.OPENAI else None ) diff --git a/litellm/llms/openai/workload_identity.py b/litellm/llms/openai/workload_identity.py index ecec161ed46..48e9b779a34 100644 --- a/litellm/llms/openai/workload_identity.py +++ b/litellm/llms/openai/workload_identity.py @@ -1,5 +1,6 @@ from __future__ import annotations +from collections.abc import Mapping from dataclasses import dataclass from functools import lru_cache from typing import TYPE_CHECKING, Final @@ -44,6 +45,7 @@ class OpenAIWorkloadIdentityConfig: def resolve_openai_workload_identity_config( api_key: str | None, api_base: str | None, + litellm_params: Mapping[str, object] | None = None, ) -> OpenAIWorkloadIdentityConfig | None: static_api_key: Final = normalize_nonempty_secret_str(api_key) or normalize_nonempty_secret_str( get_secret_str("OPENAI_API_KEY") @@ -55,10 +57,12 @@ def resolve_openai_workload_identity_config( ) 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") - token_file: Final = get_secret_str("OPENAI_IDENTITY_TOKEN_FILE") - if not identity_provider_id or not service_account_id or not token_file: + identity_provider_id: Final = _config_value( + litellm_params, "openai_identity_provider_id", "OPENAI_IDENTITY_PROVIDER_ID" + ) + service_account_id: Final = _config_value(litellm_params, "openai_service_account_id", "OPENAI_SERVICE_ACCOUNT_ID") + token_file: Final = _config_value(litellm_params, "openai_identity_token_file", "OPENAI_IDENTITY_TOKEN_FILE") + if identity_provider_id is None or service_account_id is None or token_file is None: return None return OpenAIWorkloadIdentityConfig( identity_provider_id=identity_provider_id, @@ -71,6 +75,13 @@ def get_workload_identity_bearer_token(config: OpenAIWorkloadIdentityConfig) -> return _workload_identity_auth(config).get_token() +def _config_value(litellm_params: Mapping[str, object] | None, param_key: str, env_name: str) -> str | None: + param_value: Final = litellm_params.get(param_key) if litellm_params is not None else None + if isinstance(param_value, str) and param_value: + return param_value + return normalize_nonempty_secret_str(get_secret_str(env_name)) + + def _targets_openai_api(api_base: str | None) -> bool: if api_base is None: return True diff --git a/litellm/main.py b/litellm/main.py index 0bca4a7350e..cf8572d42d3 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -6317,6 +6317,7 @@ def embedding( aembedding=aembedding, max_retries=max_retries, shared_session=shared_session, + litellm_params=litellm_params_dict, ) elif custom_llm_provider == "databricks": api_base = api_base or litellm.api_base or get_secret("DATABRICKS_API_BASE") diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 47f29095c81..accfd6a5f0d 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -34,7 +34,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import ( ) from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS from litellm.types.router import reject_server_owned_wif_params as _reject_server_owned_wif_params -from litellm.types.utils import CustomPricingLiteLLMParams, anthropic_wif_litellm_params +from litellm.types.utils import CustomPricingLiteLLMParams, server_owned_wif_litellm_params def is_invalid_virtual_key_error(exception: BaseException | None) -> bool: @@ -227,7 +227,7 @@ def _allow_model_level_clientside_configurable_parameters( # ``extra_body.aws_web_identity_token``) without re-validating, so the # banned-key check has to descend into it the same way it descends into # ``litellm_embedding_config``. -_ANTHROPIC_WIF_UNCONDITIONAL_BANNED: Final[tuple[str, ...]] = anthropic_wif_litellm_params +_SERVER_OWNED_WIF_UNCONDITIONAL_BANNED: Final[tuple[str, ...]] = server_owned_wif_litellm_params # The Bedrock Claude Platform route reads a workspace from workspace_id or aws_workspace_id as # well, and neither is a federation parameter, so say so rather than leaving that caller stuck. diff --git a/litellm/proxy/common_utils/credential_hydration.py b/litellm/proxy/common_utils/credential_hydration.py index 473d8a99c9d..1497de36d24 100644 --- a/litellm/proxy/common_utils/credential_hydration.py +++ b/litellm/proxy/common_utils/credential_hydration.py @@ -17,8 +17,8 @@ from litellm.proxy.utils import PrismaClient from litellm.repositories.credentials_repository import CredentialsRepository from litellm.types.router import ( GenericLiteLLMParams, - anthropic_wif_fields_named, - anthropic_wif_fields_present, + server_owned_wif_fields_named, + server_owned_wif_fields_present, ) from litellm.types.utils import CredentialItem @@ -87,16 +87,16 @@ async def named_credential_wif_fields( name for credential in litellm.credential_list if credential.credential_name == credential_name - for name in anthropic_wif_fields_named(credential.credential_values) + for name in server_owned_wif_fields_named(credential.credential_values) ) if prisma_client is None: return in_memory db_credential: Final = await CredentialsRepository(prisma_client).find_by_name(credential_name) - stored: Final = () if db_credential is None else anthropic_wif_fields_named(db_credential.credential_values) + stored: Final = () if db_credential is None else server_owned_wif_fields_named(db_credential.credential_values) return tuple(dict.fromkeys(in_memory + stored)) -async def effective_anthropic_wif_fields( +async def effective_server_owned_wif_fields( stored: Mapping[str, object] | None, incoming: GenericLiteLLMParams | None, prisma_client: PrismaClient | None, @@ -108,12 +108,12 @@ async def effective_anthropic_wif_fields( attaches ``litellm_credential_name`` inherits whatever that credential holds. The two sides are matched differently on purpose. ``stored`` is matched by VALUE, because - ``GenericLiteLLMParams`` declares every ``anthropic_*`` field, so matching it by key would + ``GenericLiteLLMParams`` declares every federation field, so matching it by key would report every deployment on the proxy as federated. ``incoming`` is matched by the keys the write actually set, so an explicit null still counts as touching the field. """ - from_stored: Final = () if stored is None else anthropic_wif_fields_present(stored) - from_incoming: Final = () if incoming is None else anthropic_wif_fields_named(incoming.model_fields_set) + from_stored: Final = () if stored is None else server_owned_wif_fields_present(stored) + from_incoming: Final = () if incoming is None else server_owned_wif_fields_named(incoming.model_fields_set) from_credential: Final = tuple( chain.from_iterable( await asyncio.gather( diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index 35b5f9b8d61..06192b0ec7a 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -33,7 +33,7 @@ from litellm.proxy.common_utils.credential_hydration import ( from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper from litellm.proxy.utils import handle_exception_on_proxy, jsonify_object from litellm.repositories.credentials_repository import CredentialsRepository -from litellm.types.router import anthropic_wif_fields_named +from litellm.types.router import server_owned_wif_fields_named from litellm.types.utils import CreateCredentialItem, CredentialItem router: Final = APIRouter() @@ -67,13 +67,13 @@ def _incoming_wif_fields(credential: CredentialItem) -> tuple[str, ...]: names in ``credential_values_to_delete``, since dropping a federation field off the stored credential breaks every deployment referencing it just as installing one would redirect them. """ - return anthropic_wif_fields_named(credential.credential_values) + anthropic_wif_fields_named( + return server_owned_wif_fields_named(credential.credential_values) + server_owned_wif_fields_named( credential.credential_values_to_delete or () ) def _stored_wif_fields(stored_credential: CredentialItem) -> tuple[str, ...]: - return anthropic_wif_fields_named(stored_credential.credential_values) + return server_owned_wif_fields_named(stored_credential.credential_values) def _reject_overlapping_credential_values(credential: CredentialItem) -> None: @@ -178,7 +178,7 @@ async def create_credential( status_code=400, detail="Credential values are required. Unable to infer credential values from model ID.", ) - _reject_non_admin_wif_fields(anthropic_wif_fields_named(credential.credential_values), user_api_key_dict) + _reject_non_admin_wif_fields(server_owned_wif_fields_named(credential.credential_values), user_api_key_dict) _reject_non_admin_wif_fields( await named_credential_wif_fields(credential.credential_name, prisma_client), user_api_key_dict ) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index c0c342442eb..a35919f3e8f 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -56,7 +56,7 @@ from litellm.proxy.common_utils.config_sync_pubsub import ( publish_config_change, ) from litellm.proxy.common_utils.credential_hydration import ( - effective_anthropic_wif_fields, + effective_server_owned_wif_fields, hydrate_named_credential, hydrate_named_credential_authoritative, ) @@ -1592,13 +1592,13 @@ class ModelManagementAuthChecks: if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN: return stored: Final = model_params.litellm_params.model_dump(exclude_none=True) - wif_fields: Final = await effective_anthropic_wif_fields(stored, incoming_params, prisma_client) + wif_fields: Final = await effective_server_owned_wif_fields(stored, incoming_params, prisma_client) if wif_fields: # ProxyException rather than HTTPException so the offending field stays a structured # `param`, which is the contract the narrower gate this replaced already published. raise ProxyException( message=( - f"Only proxy admins can modify a deployment configured for Anthropic workload identity " + f"Only proxy admins can modify a deployment configured for workload identity " f"federation ({wif_fields[0]!r})." ), type=ProxyErrorTypes.auth_error.value, diff --git a/litellm/proxy/public_endpoints/provider_create_fields.json b/litellm/proxy/public_endpoints/provider_create_fields.json index 30786dc7456..39d0f81055f 100644 --- a/litellm/proxy/public_endpoints/provider_create_fields.json +++ b/litellm/proxy/public_endpoints/provider_create_fields.json @@ -2617,6 +2617,92 @@ "default_value": null } ], + "credential_variants": { + "selector_label": "Authentication method", + "default_variant": "api_key", + "field_definitions": [ + { + "key": "api_base", + "label": "API Base", + "placeholder": "https://api.openai.com/v1", + "tooltip": "Common endpoints: https://api.openai.com/v1, https://eu.api.openai.com, https://us.api.openai.com", + "required": false, + "field_type": "text", + "options": null, + "default_value": "https://api.openai.com/v1" + }, + { + "key": "organization", + "label": "OpenAI Organization ID", + "placeholder": "[OPTIONAL] my-unique-org", + "tooltip": null, + "required": false, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "api_key", + "label": "OpenAI API Key", + "placeholder": null, + "tooltip": null, + "required": true, + "field_type": "password", + "options": null, + "default_value": null + }, + { + "key": "openai_identity_provider_id", + "label": "Identity Provider ID", + "placeholder": "idp_...", + "tooltip": "The identity provider id from the OpenAI platform's workload identity federation settings.", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "openai_service_account_id", + "label": "Service Account ID", + "placeholder": "user-...", + "tooltip": "The OpenAI service account the federated workload authenticates as.", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "openai_identity_token_file", + "label": "Identity Token File Path", + "placeholder": "/var/run/secrets/tokens/openai", + "tooltip": "Absolute path to a mounted file containing the workload's OIDC identity token. Federation only activates when no static OpenAI API key is set on the deployment or in OPENAI_API_KEY.", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + } + ], + "variants": [ + { + "id": "api_key", + "label": "API Key", + "field_keys": ["api_base", "organization", "api_key"], + "fixed_values": {} + }, + { + "id": "wif_token_file", + "label": "Workload Identity Federation (token file)", + "field_keys": [ + "api_base", + "organization", + "openai_identity_provider_id", + "openai_service_account_id", + "openai_identity_token_file" + ], + "fixed_values": {} + } + ] + }, "default_model_placeholder": "gpt-3.5-turbo" }, { diff --git a/litellm/router_utils/clientside_credential_handler.py b/litellm/router_utils/clientside_credential_handler.py index 8c7c2fabdd1..186772925a9 100644 --- a/litellm/router_utils/clientside_credential_handler.py +++ b/litellm/router_utils/clientside_credential_handler.py @@ -13,7 +13,7 @@ Ensures cooldowns are applied correctly. from typing import Final -from litellm.types.utils import anthropic_wif_litellm_params +from litellm.types.utils import server_owned_wif_litellm_params clientside_credential_keys: Final = ["api_key", "api_base", "base_url"] @@ -21,7 +21,7 @@ clientside_credential_keys: Final = ["api_key", "api_base", "base_url"] # mint a federation token there even when WIF is configured only through ANTHROPIC_* env vars (which # cannot be cleared from litellm_params). DISABLE_WORKLOAD_IDENTITY_PARAM: Final = "anthropic_disable_workload_identity_federation" -_ANTHROPIC_WIF_CLEAR_ON_BASE_OVERRIDE: Final = tuple(sorted(anthropic_wif_litellm_params)) +_WIF_CLEAR_ON_BASE_OVERRIDE: Final = tuple(sorted(server_owned_wif_litellm_params)) def _admin_config_fields_to_clear_on_base_override() -> list[str]: @@ -67,14 +67,14 @@ def _admin_config_fields_to_clear_on_base_override() -> list[str]: # ``api_base`` for the same reason as the OCI entries above. "nvcf_function_id", "use_ssl", - # Anthropic workload-identity federation minting fields, restated here from - # anthropic_wif_litellm_params the same way azure_ad_token above is restated + # Workload-identity federation minting fields, restated here from + # server_owned_wif_litellm_params the same way azure_ad_token above is restated # despite also being declared on CredentialLiteLLMParams (hence covered by # typed_fields too): a federation token minted for a client-redirected api_base # would send the workload's OIDC assertion, and then the minted bearer, to the # caller-chosen host, so this list must stay correct even if a field is ever # dropped from the typed model. - *_ANTHROPIC_WIF_CLEAR_ON_BASE_OVERRIDE, + *_WIF_CLEAR_ON_BASE_OVERRIDE, ] return typed_fields + kwargs_only_fields diff --git a/litellm/types/router.py b/litellm/types/router.py index bd31d228a38..96daf208091 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -29,7 +29,7 @@ from .utils import ( ) from .utils import ( # private alias: `from .types.router import *` would rebind a public Final in litellm/__init__.py - anthropic_wif_litellm_params as _anthropic_wif_litellm_params, + server_owned_wif_litellm_params as _server_owned_wif_litellm_params, ) @@ -324,40 +324,43 @@ class CredentialLiteLLMParams(BaseModel): # other federation fields above are declared for, rather than being rebuilt away in transit. anthropic_disable_workload_identity_federation: bool | None = None + ## OPENAI WORKLOAD IDENTITY FEDERATION ## + openai_identity_provider_id: str | None = None + openai_service_account_id: str | None = None + openai_identity_token_file: str | None = None -def anthropic_wif_fields_present(fields: Mapping[str, object]) -> tuple[str, ...]: - """Server-owned Anthropic workload identity federation field names set in ``fields``. + +def server_owned_wif_fields_present(fields: Mapping[str, object]) -> tuple[str, ...]: + """Server-owned workload identity federation field names set in ``fields``. ``fields`` is a ``litellm_params`` dict (or a credential's ``credential_values`` mapping, which feeds the same resolution when referenced by name). Derived from - ``anthropic_wif_litellm_params`` rather than hand-copied, so a persistence gate built on + ``server_owned_wif_litellm_params`` rather than hand-copied, so a persistence gate built on this stays correct when a new WIF field is added there. """ - return tuple(name for name in _anthropic_wif_litellm_params if fields.get(name) is not None) + return tuple(name for name in _server_owned_wif_litellm_params if fields.get(name) is not None) -def anthropic_wif_fields_named(keys: Container[str]) -> tuple[str, ...]: - """Server-owned Anthropic workload identity federation field names that appear in ``keys``, - whatever value they carry. +def server_owned_wif_fields_named(keys: Container[str]) -> tuple[str, ...]: + """Server-owned workload identity federation field names that appear in ``keys``, whatever + value they carry. - The write gates on credentials need this key-based sibling of ``anthropic_wif_fields_present``: + The write gates on credentials need this key-based sibling of ``server_owned_wif_fields_present``: ``get_litellm_params`` forwards a WIF kwarg on key presence and the federation resolver rejects a foreign variant's field by key, so a persisted ``{"anthropic_issuer_url": None}`` wedges every deployment that references the credential even though no value is set. Pass a mapping (its keys are tested) or a plain collection of key names. """ - return tuple(name for name in _anthropic_wif_litellm_params if name in keys) + return tuple(name for name in _server_owned_wif_litellm_params if name in keys) -_ANTHROPIC_WIF_POINTER_FIELDS: Final = frozenset( - name for name in _anthropic_wif_litellm_params if name.endswith("_ref") -) +_WIF_POINTER_FIELDS: Final = frozenset(name for name in _server_owned_wif_litellm_params if name.endswith("_ref")) def holds_secret_pointer(param_name: str) -> bool: """A ``*_ref`` federation field is a secret POINTER the identity source dereferences at use time, so a loader expanding ``os.environ/`` values must leave it as written.""" - return param_name in _ANTHROPIC_WIF_POINTER_FIELDS + return param_name in _WIF_POINTER_FIELDS _RESERVED_INIT_KEYS: Final = frozenset({"self", "params", "__class__"}) @@ -1130,7 +1133,7 @@ def reject_server_owned_wif_params(body: Mapping[str, object]) -> None: opt-in. This lives here rather than under ``litellm.proxy`` so the router can call it on a post-authentication merge without core importing from the proxy package. """ - for param in _anthropic_wif_litellm_params: + for param in _server_owned_wif_litellm_params: if param in body: raise ValueError( f"Rejected Request: {param} is a server-owned workload identity federation parameter " diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 282a73630da..5e6574f5a68 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3595,13 +3595,16 @@ bedrock_batch_litellm_params: Final = ( # already bound on the partially-initialized module. from ..litellm_core_utils.get_litellm_params import ( # noqa: E402 # deferred past CallTypes to break the import cycle ANTHROPIC_WIF_KWARGS_KEYS, + OPENAI_WIF_KWARGS_KEYS, ) anthropic_wif_litellm_params: Final = tuple(sorted(ANTHROPIC_WIF_KWARGS_KEYS)) +openai_wif_litellm_params: Final = tuple(sorted(OPENAI_WIF_KWARGS_KEYS)) +server_owned_wif_litellm_params: Final = anthropic_wif_litellm_params + openai_wif_litellm_params all_litellm_params = ( agentic_loop_internal_litellm_params - + [TRUSTED_CALLBACK_VARS_FIELD, *bedrock_batch_litellm_params, *anthropic_wif_litellm_params] + + [TRUSTED_CALLBACK_VARS_FIELD, *bedrock_batch_litellm_params, *server_owned_wif_litellm_params] + [ "metadata", "litellm_metadata", diff --git a/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py b/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py index d84bb54eb1a..ad4e417a29f 100644 --- a/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py +++ b/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py @@ -332,3 +332,41 @@ class TestAnthropicWifIdentitySourceKeys: params = get_litellm_params() for key in self.NEW_KEYS: assert key not in params + + +class TestOpenAIWifKeys: + """The three openai_* WIF keys carry a deployment's federation identity through the kwargs + funnel into litellm_params (where the OpenAI client factory reads them) and stay out of the + provider body, exactly like the anthropic_* keys above.""" + + THREE_KEYS = { + "openai_identity_provider_id": "idp_1", + "openai_service_account_id": "user-1", + "openai_identity_token_file": "/var/run/secrets/tokens/openai", + } + + def test_keys_are_exactly_the_registered_set(self): + from litellm.litellm_core_utils.get_litellm_params import OPENAI_WIF_KWARGS_KEYS + + assert set(self.THREE_KEYS) == OPENAI_WIF_KWARGS_KEYS + + def test_keys_survive_into_litellm_params(self): + params = get_litellm_params(**self.THREE_KEYS) + for key, value in self.THREE_KEYS.items(): + assert params[key] == value + + def test_keys_are_forwarded_from_completion_kwargs(self): + from litellm.litellm_core_utils.get_litellm_params import FORWARDED_KWARGS_KEYS + + assert set(self.THREE_KEYS) <= FORWARDED_KWARGS_KEYS + + def test_keys_stay_out_of_the_provider_body(self): + from litellm.types.utils import all_litellm_params + + for key in self.THREE_KEYS: + assert key in all_litellm_params + + def test_keys_absent_when_not_configured(self): + params = get_litellm_params() + for key in self.THREE_KEYS: + assert key not in params diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index 3e616d4d31d..7ede9480346 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -3509,10 +3509,10 @@ class TestWifParamsAreNotClientSettable: The workspace id was once carved out here as inert; it is not. It is the scope of the minted org credential, and the router merges request kwargs over deployment params, so a caller who set it picked the scope instead of the administrator.""" - from litellm.proxy.auth.auth_utils import _ANTHROPIC_WIF_UNCONDITIONAL_BANNED + from litellm.proxy.auth.auth_utils import _SERVER_OWNED_WIF_UNCONDITIONAL_BANNED from litellm.types.utils import anthropic_wif_litellm_params - assert set(_ANTHROPIC_WIF_UNCONDITIONAL_BANNED) == set(anthropic_wif_litellm_params) + assert set(anthropic_wif_litellm_params) <= set(_SERVER_OWNED_WIF_UNCONDITIONAL_BANNED) class TestWifServerOwnedParamsAreUnconditional: 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 d8d9936e9a1..a152a70bcec 100644 --- a/tests/test_litellm/llms/openai/test_openai_workload_identity.py +++ b/tests/test_litellm/llms/openai/test_openai_workload_identity.py @@ -10,6 +10,7 @@ from openai import AsyncOpenAI, OpenAI import litellm from litellm.llms.litellm_proxy.responses.transformation import LiteLLMProxyResponsesAPIConfig +from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm.llms.openai.common_utils import BaseOpenAILLM, OpenAIError from litellm.llms.openai.openai import OpenAIChatCompletion from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig @@ -22,6 +23,17 @@ from litellm.llms.openai.workload_identity import ( from litellm.types.router import GenericLiteLLMParams TOKEN_EXCHANGE_URL: Final = "https://auth.openai.com/oauth/token" +CHAT_COMPLETIONS_URL: Final = "https://api.openai.com/v1/chat/completions" +EMBEDDINGS_URL: Final = "https://api.openai.com/v1/embeddings" +MODELS_URL: Final = "https://api.openai.com/v1/models" +CHAT_COMPLETION_BODY: Final = { + "id": "chatcmpl-wif", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, +} @pytest.fixture @@ -236,3 +248,323 @@ class TestResponsesValidateEnvironment: headers={}, model="gpt-4o-mini", litellm_params=GenericLiteLLMParams() ) assert headers["Authorization"] == "Bearer None" + + +@pytest.fixture +def deployment_wif(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> dict[str, str]: + token_file: Final = tmp_path / "deployment_subject_token.jwt" + token_file.write_text("subject-token-from-deployment-file") + for name in ( + "OPENAI_API_KEY", + "OPENAI_BASE_URL", + "OPENAI_API_BASE", + "OPENAI_IDENTITY_PROVIDER_ID", + "OPENAI_SERVICE_ACCOUNT_ID", + "OPENAI_IDENTITY_TOKEN_FILE", + ): + monkeypatch.delenv(name, raising=False) + monkeypatch.setattr(litellm, "api_base", None) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + _workload_identity_auth.cache_clear() + litellm.in_memory_llm_clients_cache.flush_cache() + return { + "openai_identity_provider_id": "idp_deployment", + "openai_service_account_id": "user-deployment", + "openai_identity_token_file": str(token_file), + } + + +def deployment_config(deployment_wif: dict[str, str]) -> OpenAIWorkloadIdentityConfig: + return OpenAIWorkloadIdentityConfig( + identity_provider_id="idp_deployment", + service_account_id="user-deployment", + token_file=deployment_wif["openai_identity_token_file"], + ) + + +def mock_chat_completions() -> respx.Route: + return respx.post(CHAT_COMPLETIONS_URL).mock(return_value=httpx.Response(200, json=CHAT_COMPLETION_BODY)) + + +def mock_streaming_chat_completions() -> respx.Route: + chunk: Final = {"id": "chatcmpl-1", "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini"} + events: Final = ( + {**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", "content": "ok"}, "finish_reason": None}]}, + {**chunk, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}, + ) + body: Final = "".join(f"data: {json.dumps(event)}\n\n" for event in events) + "data: [DONE]\n\n" + return respx.post(CHAT_COMPLETIONS_URL).mock( + return_value=httpx.Response(200, headers={"content-type": "text/event-stream"}, content=body) + ) + + +class TestResolveConfigFromDeployment: + def test_resolves_from_litellm_params_without_env(self, deployment_wif: dict[str, str]) -> None: + assert resolve_openai_workload_identity_config( + api_key=None, api_base=None, litellm_params=deployment_wif + ) == deployment_config(deployment_wif) + + def test_env_alone_disables_nothing_when_params_are_absent(self, deployment_wif: dict[str, str]) -> None: + assert resolve_openai_workload_identity_config(api_key=None, api_base=None, litellm_params=None) is None + + def test_unrelated_litellm_params_do_not_resolve(self, deployment_wif: dict[str, str]) -> None: + assert ( + resolve_openai_workload_identity_config(api_key=None, api_base=None, litellm_params={"model": "gpt-4o"}) + is None + ) + + def test_litellm_params_beat_env(self, wif_env: OpenAIWorkloadIdentityConfig) -> None: + config: Final = resolve_openai_workload_identity_config( + api_key=None, + api_base=None, + litellm_params={ + "openai_identity_provider_id": "idp_deployment", + "openai_service_account_id": "user-deployment", + "openai_identity_token_file": wif_env.token_file, + }, + ) + assert config == OpenAIWorkloadIdentityConfig( + identity_provider_id="idp_deployment", + service_account_id="user-deployment", + token_file=wif_env.token_file, + ) + + def test_partial_litellm_params_fill_from_env_per_field(self, wif_env: OpenAIWorkloadIdentityConfig) -> None: + config: Final = resolve_openai_workload_identity_config( + api_key=None, api_base=None, litellm_params={"openai_identity_provider_id": "idp_deployment"} + ) + assert config == OpenAIWorkloadIdentityConfig( + identity_provider_id="idp_deployment", + service_account_id=wif_env.service_account_id, + token_file=wif_env.token_file, + ) + + @pytest.mark.parametrize("blank", ["", None, 7]) + def test_blank_or_non_string_param_falls_back_to_env( + self, wif_env: OpenAIWorkloadIdentityConfig, blank: object + ) -> None: + config: Final = resolve_openai_workload_identity_config( + api_key=None, api_base=None, litellm_params={"openai_identity_provider_id": blank} + ) + assert config == wif_env + + def test_partial_litellm_params_without_env_disable(self, deployment_wif: dict[str, str]) -> None: + partial: Final = {key: value for key, value in deployment_wif.items() if key != "openai_identity_token_file"} + assert resolve_openai_workload_identity_config(api_key=None, api_base=None, litellm_params=partial) is None + + def test_static_api_key_beats_litellm_params(self, deployment_wif: dict[str, str]) -> None: + assert ( + resolve_openai_workload_identity_config(api_key="sk-static", api_base=None, litellm_params=deployment_wif) + is None + ) + + def test_env_openai_api_key_beats_litellm_params( + self, deployment_wif: dict[str, str], monkeypatch: pytest.MonkeyPatch + ) -> None: + monkeypatch.setenv("OPENAI_API_KEY", "sk-from-env") + assert ( + resolve_openai_workload_identity_config(api_key=None, api_base=None, litellm_params=deployment_wif) is None + ) + + def test_foreign_api_base_disables_deployment_wif(self, deployment_wif: dict[str, str]) -> None: + assert ( + resolve_openai_workload_identity_config( + api_key=None, api_base="https://my-vllm.internal/v1", litellm_params=deployment_wif + ) + is None + ) + + +class TestDeploymentClientConstruction: + def test_sync_client_from_deployment_params(self, deployment_wif: dict[str, str]) -> None: + client: Final = OpenAIChatCompletion()._get_openai_client( + is_async=False, api_key=None, api_base=None, litellm_params=deployment_wif + ) + assert isinstance(client, OpenAI) + assert client.api_key == "workload-identity-auth" + assert client._workload_identity_auth is not None + + def test_async_client_from_deployment_params(self, deployment_wif: dict[str, str]) -> None: + client: Final = OpenAIChatCompletion()._get_openai_client( + is_async=True, api_key=None, api_base=None, litellm_params=deployment_wif + ) + assert isinstance(client, AsyncOpenAI) + assert client._workload_identity_auth is not None + + def test_distinct_deployments_get_distinct_cached_clients(self, deployment_wif: dict[str, str]) -> None: + other_deployment: Final = {**deployment_wif, "openai_service_account_id": "user-other"} + handler: Final = OpenAIChatCompletion() + first: Final = handler._get_openai_client( + is_async=False, api_key=None, api_base=None, litellm_params=deployment_wif + ) + second: Final = handler._get_openai_client( + is_async=False, api_key=None, api_base=None, litellm_params=other_deployment + ) + again: Final = handler._get_openai_client( + is_async=False, api_key=None, api_base=None, litellm_params=dict(deployment_wif) + ) + assert first is not second + assert again is first + + @respx.mock + def test_completion_kwargs_carry_exchanged_bearer(self, deployment_wif: dict[str, str]) -> None: + mock_token_exchange("deployment-bearer") + completion_route: Final = mock_chat_completions() + + response: Final = litellm.completion( + model="openai/gpt-4o-mini", messages=[{"role": "user", "content": "hi"}], **deployment_wif + ) + + assert response.choices[0].message.content == "ok" + request: Final = completion_route.calls.last.request + assert request.headers["Authorization"] == "Bearer deployment-bearer" + assert not any(key.startswith("openai_") for key in json.loads(request.content)) + + @respx.mock + def test_streaming_completion_kwargs_carry_exchanged_bearer(self, deployment_wif: dict[str, str]) -> None: + mock_token_exchange("stream-bearer") + stream_route: Final = mock_streaming_chat_completions() + + chunks: Final = tuple( + litellm.completion( + model="openai/gpt-4o-mini", messages=[{"role": "user", "content": "hi"}], stream=True, **deployment_wif + ) + ) + + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "ok" + assert stream_route.calls.last.request.headers["Authorization"] == "Bearer stream-bearer" + + @respx.mock + @pytest.mark.asyncio + async def test_async_streaming_completion_kwargs_carry_exchanged_bearer( + self, deployment_wif: dict[str, str] + ) -> None: + mock_token_exchange("async-stream-bearer") + stream_route: Final = mock_streaming_chat_completions() + + stream: Final = await litellm.acompletion( + model="openai/gpt-4o-mini", messages=[{"role": "user", "content": "hi"}], stream=True, **deployment_wif + ) + chunks: Final = tuple([chunk async for chunk in stream]) + + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "ok" + assert stream_route.calls.last.request.headers["Authorization"] == "Bearer async-stream-bearer" + + @respx.mock + def test_router_deployment_without_api_key_authenticates_via_token_exchange( + self, deployment_wif: dict[str, str] + ) -> None: + exchange_route: Final = mock_token_exchange("router-bearer") + completion_route: Final = mock_chat_completions() + router: Final = litellm.Router( + model_list=[{"model_name": "wif-gpt", "litellm_params": {"model": "openai/gpt-4o-mini", **deployment_wif}}] + ) + + response: Final = router.completion(model="wif-gpt", messages=[{"role": "user", "content": "hi"}]) + + assert response.choices[0].message.content == "ok" + assert exchange_route.called + assert completion_route.calls.last.request.headers["Authorization"] == "Bearer router-bearer" + + @respx.mock + def test_embedding_kwargs_carry_exchanged_bearer(self, deployment_wif: dict[str, str]) -> None: + mock_token_exchange("embedding-bearer") + embeddings_route: Final = respx.post(EMBEDDINGS_URL).mock( + return_value=httpx.Response( + 200, + json={ + "object": "list", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 1, "total_tokens": 1}, + }, + ) + ) + + litellm.embedding(model="openai/text-embedding-3-small", input=["hi"], **deployment_wif) + + assert embeddings_route.calls.last.request.headers["Authorization"] == "Bearer embedding-bearer" + + +class TestResponsesValidateEnvironmentFromDeployment: + @respx.mock + def test_mints_bearer_from_litellm_params(self, deployment_wif: dict[str, str]) -> None: + mock_token_exchange("responses-bearer") + headers: Final = OpenAIResponsesAPIConfig().validate_environment( + headers={}, model="gpt-4o-mini", litellm_params=GenericLiteLLMParams(**deployment_wif) + ) + assert headers["Authorization"] == "Bearer responses-bearer" + + def test_static_key_in_litellm_params_wins(self, deployment_wif: dict[str, str]) -> None: + headers: Final = OpenAIResponsesAPIConfig().validate_environment( + headers={}, + model="gpt-4o-mini", + litellm_params=GenericLiteLLMParams(api_key="sk-responses", **deployment_wif), + ) + assert headers["Authorization"] == "Bearer sk-responses" + + +class TestDiscoverModels: + @staticmethod + def mock_models() -> respx.Route: + return respx.get(MODELS_URL).mock( + return_value=httpx.Response(200, json={"data": [{"id": "gpt-4o-mini"}, {"id": "gpt-4.1"}]}) + ) + + @respx.mock + def test_discovers_with_exchanged_bearer_from_litellm_params(self, deployment_wif: dict[str, str]) -> None: + mock_token_exchange("discovery-bearer") + models_route: Final = self.mock_models() + + assert OpenAIGPTConfig().discover_models(deployment_wif) == ["gpt-4o-mini", "gpt-4.1"] + assert models_route.calls.last.request.headers["Authorization"] == "Bearer discovery-bearer" + + @respx.mock + def test_discovers_with_env_wif_when_params_carry_no_key(self, wif_env: OpenAIWorkloadIdentityConfig) -> None: + mock_token_exchange("env-discovery-bearer") + models_route: Final = self.mock_models() + + OpenAIGPTConfig().discover_models({}) + + assert models_route.calls.last.request.headers["Authorization"] == "Bearer env-discovery-bearer" + + @respx.mock + def test_static_api_key_in_params_skips_token_exchange(self, deployment_wif: dict[str, str]) -> None: + exchange_route: Final = mock_token_exchange() + models_route: Final = self.mock_models() + + OpenAIGPTConfig().discover_models({**deployment_wif, "api_key": "sk-discovery"}) + + assert models_route.calls.last.request.headers["Authorization"] == "Bearer sk-discovery" + assert not exchange_route.called + + @respx.mock + def test_openai_compatible_subclass_never_mints_wif(self, deployment_wif: dict[str, str]) -> None: + exchange_route: Final = mock_token_exchange() + models_route: Final = self.mock_models() + + class CompatibleConfig(OpenAIGPTConfig): + pass + + CompatibleConfig().discover_models(deployment_wif) + + assert models_route.calls.last.request.headers["Authorization"] == "Bearer None" + assert not exchange_route.called + + +class TestClientsideBaseOverride: + def test_client_api_base_override_clears_deployment_wif(self, deployment_wif: dict[str, str]) -> None: + from litellm.router_utils.clientside_credential_handler import get_dynamic_litellm_params + + redirected: Final = get_dynamic_litellm_params( + litellm_params={"model": "openai/gpt-4o-mini", **deployment_wif}, + request_kwargs={"api_base": "https://not-openai.example/v1"}, + ) + + assert not any(key in redirected for key in deployment_wif) + assert ( + resolve_openai_workload_identity_config( + api_key=None, api_base=redirected["api_base"], litellm_params=redirected + ) + is None + ) diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 8c0a95d3f31..f5c7ac713c1 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -29,16 +29,20 @@ from litellm.proxy.auth.auth_utils import ( ) -def test_every_anthropic_wif_kwarg_key_is_request_banned(): - """anthropic_wif_litellm_params (types/utils.py) is derived from ANTHROPIC_WIF_KWARGS_KEYS - (get_litellm_params.py) precisely so a new WIF field can never be added to the kwargs funnel +def test_every_server_owned_wif_kwarg_key_is_request_banned(): + """server_owned_wif_litellm_params (types/utils.py) is derived from ANTHROPIC_WIF_KWARGS_KEYS + and OPENAI_WIF_KWARGS_KEYS (get_litellm_params.py) precisely so a new WIF field can never be + added to the kwargs funnel without automatically joining the request-body ban list; this guards that invariant itself, independent of today's field count, so it fails if the derivation is ever reverted to a hand-typed list that drifts.""" - from litellm.litellm_core_utils.get_litellm_params import ANTHROPIC_WIF_KWARGS_KEYS - from litellm.proxy.auth.auth_utils import _ANTHROPIC_WIF_UNCONDITIONAL_BANNED + from litellm.litellm_core_utils.get_litellm_params import ( + ANTHROPIC_WIF_KWARGS_KEYS, + OPENAI_WIF_KWARGS_KEYS, + ) + from litellm.proxy.auth.auth_utils import _SERVER_OWNED_WIF_UNCONDITIONAL_BANNED - assert ANTHROPIC_WIF_KWARGS_KEYS == set(_ANTHROPIC_WIF_UNCONDITIONAL_BANNED) + assert ANTHROPIC_WIF_KWARGS_KEYS | OPENAI_WIF_KWARGS_KEYS == set(_SERVER_OWNED_WIF_UNCONDITIONAL_BANNED) class TestCustomAuthCommonChecksWarning: diff --git a/tests/test_litellm/proxy/credential_endpoints/test_endpoints.py b/tests/test_litellm/proxy/credential_endpoints/test_endpoints.py index 0c0306fe624..6530de0fd2e 100644 --- a/tests/test_litellm/proxy/credential_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/credential_endpoints/test_endpoints.py @@ -502,6 +502,40 @@ class TestNonAdminCannotPersistWifFieldsOnCredential: assert response.status_code == 200, response.text repository.create.assert_awaited_once() + def test_non_admin_cannot_create_a_credential_with_an_openai_token_file(self): + with patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ): + response = _post_credential( + { + "credential_name": "attacker-cred", + "credential_values": {"openai_identity_token_file": "/var/run/secrets/tokens/attacker"}, + "credential_info": {"custom_llm_provider": "openai"}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + assert "openai_identity_token_file" in response.json()["error"]["message"] + + def test_proxy_admin_can_create_a_credential_with_the_openai_identity_trio(self, restore_credential_list): + with _repository_holding(None) as repository: + response = _post_credential( + { + "credential_name": "openai-wif", + "credential_values": { + "openai_identity_provider_id": "idp_1", + "openai_service_account_id": "user-1", + "openai_identity_token_file": "/var/run/secrets/tokens/openai", + }, + "credential_info": {"custom_llm_provider": "openai"}, + }, + auth=_as_admin, + ) + + assert response.status_code == 200, response.text + repository.create.assert_awaited_once() + def test_non_admin_cannot_update_a_credential_to_add_a_wif_destination(self): stored = CredentialItem( credential_name="existing", diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index af653cd90fb..1fc9fef7552 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -4619,7 +4619,7 @@ class TestNonAdminCannotPersistWifFieldsOnModel: ), ): with pytest.raises( - Exception, match="Only proxy admins can modify a deployment configured for Anthropic" + Exception, match="Only proxy admins can modify a deployment configured for workload identity" ) as exc_info: await patch_model( model_id="m1", @@ -4634,6 +4634,57 @@ class TestNonAdminCannotPersistWifFieldsOnModel: assert getattr(err, "param", "") == "anthropic_keycloak_token_url" mock_prisma.db.litellm_proxymodeltable.update.assert_not_called() + @pytest.mark.asyncio + async def test_patch_model_non_admin_cannot_set_openai_wif_field(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + patch_model, + ) + + non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER) + existing_row = MagicMock() + existing_row.litellm_params = {"model": "openai/gpt-4o-mini"} + existing_row.model_dump.return_value = { + "model_name": "gpt", + "litellm_params": existing_row.litellm_params, + "model_info": {"id": "m1"}, + } + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", + mock_prisma, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.llm_router", + MagicMock(**{"get_model_ids.return_value": ["m1"]}), + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", + True, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", + True, + ), + ): + with pytest.raises( + Exception, match="Only proxy admins can modify a deployment configured for workload identity" + ) as exc_info: + await patch_model( + model_id="m1", + patch_data=updateDeployment( + litellm_params=updateLiteLLMParams( + openai_identity_token_file="/var/run/secrets/tokens/attacker", + ) + ), + user_api_key_dict=non_admin, + ) + assert getattr(exc_info.value, "param", "") == "openai_identity_token_file" + mock_prisma.db.litellm_proxymodeltable.update.assert_not_called() + @pytest.mark.asyncio async def test_patch_model_admin_can_set_wif_field(self): from litellm.proxy.management_endpoints.model_management_endpoints import ( @@ -4824,7 +4875,7 @@ class TestNonAdminCannotPersistWifFieldsOnModel: ), ): with pytest.raises( - Exception, match="Only proxy admins can modify a deployment configured for Anthropic" + Exception, match="Only proxy admins can modify a deployment configured for workload identity" ) as exc_info: await update_model( model_params=updateDeployment( @@ -5100,6 +5151,58 @@ class TestDiscoverProviderModels: assert called_params["anthropic_federation_rule_id"] == "rule-1" assert called_params["anthropic_identity_token"] == "oidc/env/TOK" + @pytest.mark.asyncio + async def test_discovery_success_via_named_openai_wif_credential(self, monkeypatch): + """An OpenAI credential holding the per-deployment identity trio discovers models + keyless: only the credential name travels in the request and the hydrated trio + reaches the OpenAI discovery path.""" + import litellm + from litellm.proxy.management_endpoints.model_management_endpoints import ( + discover_provider_models, + ) + from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + ProviderModelDiscoveryRequest, + ) + from litellm.types.utils import CredentialItem + + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="openai-wif", + credential_values={ + "openai_identity_provider_id": "idp_1", + "openai_service_account_id": "user-1", + "openai_identity_token_file": "/var/run/secrets/tokens/openai", + }, + credential_info={"custom_llm_provider": "openai"}, + ) + ], + ) + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", _prisma_without_stored_credentials() + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.llms.openai.chat.gpt_transformation.OpenAIGPTConfig.discover_models", + return_value=["gpt-4o-mini"], + ) as discover_mock, + ): + result = await discover_provider_models( + data=ProviderModelDiscoveryRequest(custom_llm_provider="openai", litellm_credential_name="openai-wif"), + user_api_key_dict=self._admin(), + ) + assert result.models == ["gpt-4o-mini"] + called_params = ( + discover_mock.call_args.args[0] + if discover_mock.call_args.args + else discover_mock.call_args.kwargs["litellm_params"] + ) + assert called_params["openai_identity_provider_id"] == "idp_1" + assert called_params["openai_identity_token_file"] == "/var/run/secrets/tokens/openai" + assert called_params.get("api_key") is None + @pytest.mark.asyncio async def test_discovery_failure_surfaces_a_sanitized_error_never_a_silent_empty_list(self): from litellm.proxy.management_endpoints.model_management_endpoints import ( @@ -5208,7 +5311,9 @@ class TestWifBoundaryReadsTheResultingDeployment: patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: proxy wiring under test patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test ): - with pytest.raises(Exception, match="Only proxy admins can modify a deployment configured for Anthropic"): + with pytest.raises( + Exception, match="Only proxy admins can modify a deployment configured for workload identity" + ): await patch_model( model_id="m1", patch_data=updateDeployment( @@ -5258,7 +5363,9 @@ class TestWifBoundaryReadsTheResultingDeployment: patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: proxy wiring under test patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test ): - with pytest.raises(Exception, match="Only proxy admins can modify a deployment configured for Anthropic"): + with pytest.raises( + Exception, match="Only proxy admins can modify a deployment configured for workload identity" + ): await patch_model( model_id="m1", patch_data=updateDeployment( @@ -5306,7 +5413,9 @@ class TestWifBoundaryReadsTheResultingDeployment: patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: proxy wiring under test patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test ): - with pytest.raises(Exception, match="Only proxy admins can modify a deployment configured for Anthropic"): + with pytest.raises( + Exception, match="Only proxy admins can modify a deployment configured for workload identity" + ): await patch_model( model_id="m1", patch_data=updateDeployment( diff --git a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py index 15ed1368115..0b7c348356d 100644 --- a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py @@ -1196,6 +1196,51 @@ def test_anthropic_provider_fields_expose_credential_variants(): assert key in field_defs_by_key, f"variant {variant['id']} references undefined field {key}" +def test_openai_provider_fields_expose_credential_variants(): + """The OpenAI provider publishes an API-key variant and a workload-identity-federation + variant whose fields are the per-deployment identity trio, api_key not among them, while + the legacy credential_fields stays exactly api_base + organization + api_key.""" + app_instance = FastAPI() + app_instance.include_router(router) + test_client = TestClient(app_instance) + + response = test_client.get("/public/providers/fields") + assert response.status_code == 200 + providers = response.json() + + openai = next((p for p in providers if p["provider"] == "OpenAI"), None) + assert openai is not None + + legacy_fields_by_key = {f["key"]: f for f in openai["credential_fields"]} + assert set(legacy_fields_by_key) == {"api_base", "organization", "api_key"} + assert legacy_fields_by_key["api_key"]["required"] is True + + variants_block = openai["credential_variants"] + assert variants_block["default_variant"] == "api_key" + variants_by_id = {v["id"]: v for v in variants_block["variants"]} + assert set(variants_by_id) == {"api_key", "wif_token_file"} + + field_defs_by_key = {f["key"]: f for f in variants_block["field_definitions"]} + assert "api_key" in variants_by_id["api_key"]["field_keys"] + assert field_defs_by_key["api_key"]["required"] is True + + wif_field_keys = variants_by_id["wif_token_file"]["field_keys"] + assert "api_key" not in wif_field_keys + for wif_key in ( + "openai_identity_provider_id", + "openai_service_account_id", + "openai_identity_token_file", + ): + assert wif_key in wif_field_keys + assert field_defs_by_key[wif_key]["required"] is True + assert field_defs_by_key[wif_key]["field_type"] == "text" + assert variants_by_id["wif_token_file"]["fixed_values"] == {} + + for variant in variants_block["variants"]: + for key in variant["field_keys"]: + assert key in field_defs_by_key, f"variant {variant['id']} references undefined field {key}" + + def test_provider_fields_without_credential_variants_still_parse(): """A provider with no credential_variants block (the overwhelming majority) must keep parsing with the field simply absent, so old dashboards that only read credential_fields @@ -1208,9 +1253,9 @@ def test_provider_fields_without_credential_variants_still_parse(): assert response.status_code == 200 providers = response.json() - openai = next((p for p in providers if p["provider"] == "OpenAI"), None) - assert openai is not None - assert openai.get("credential_variants") is None + groq = next((p for p in providers if p["provider"] == "Groq"), None) + assert groq is not None + assert groq.get("credential_variants") is None def test_credential_variants_rejects_optional_field_keys_the_variant_does_not_mount(): diff --git a/tests/test_litellm/router_utils/test_fallback_event_handlers.py b/tests/test_litellm/router_utils/test_fallback_event_handlers.py index 45349408ee1..bff1be3ff1d 100644 --- a/tests/test_litellm/router_utils/test_fallback_event_handlers.py +++ b/tests/test_litellm/router_utils/test_fallback_event_handlers.py @@ -1089,6 +1089,23 @@ async def test_a_stored_fallback_target_cannot_carry_a_federation_field(): ) +@pytest.mark.asyncio +async def test_a_stored_fallback_target_cannot_carry_an_openai_federation_field(): + """The OpenAI identity trio is server-owned for the same reason: a stored fallback target + naming a token file would pick which workload assertion is exchanged for the bearer.""" + with pytest.raises(ValueError, match="openai_identity_token_file"): + await run_async_fallback( + litellm_router=FakeRouter(), + fallback_model_group=[ + {"model": "openai-backup", "openai_identity_token_file": "/var/run/secrets/tokens/other"} + ], + original_model_group="primary-model", + original_exception=RuntimeError("upstream limited request"), + max_fallbacks=3, + fallback_depth=0, + ) + + @pytest.mark.asyncio async def test_the_refusal_is_not_swallowed_as_a_fallback_error(): """Checked before the per-target loop on purpose: inside it, the refusal would be caught as diff --git a/tests/test_litellm/types/test_router.py b/tests/test_litellm/types/test_router.py index f8c822971cd..5f9fb6c57b5 100644 --- a/tests/test_litellm/types/test_router.py +++ b/tests/test_litellm/types/test_router.py @@ -6,14 +6,17 @@ from litellm.types.router import ( Deployment, LiteLLM_Params, ModelInfo, - anthropic_wif_fields_named, - anthropic_wif_fields_present, holds_secret_pointer, + reject_server_owned_wif_params, + server_owned_wif_fields_named, + server_owned_wif_fields_present, ) from litellm.types.utils import ( CustomPricingLiteLLMParams, MirroredPricingParams, anthropic_wif_litellm_params, + openai_wif_litellm_params, + server_owned_wif_litellm_params, ) @@ -118,35 +121,39 @@ def test_anthropic_wif_fields_round_trip_through_model_dump(): assert dumped[field] == value, field -def test_anthropic_wif_fields_present_reports_only_set_fields(): - assert anthropic_wif_fields_present({}) == () - assert anthropic_wif_fields_present({"model": "gpt-4o"}) == () - assert anthropic_wif_fields_present( +def test_server_owned_wif_fields_present_reports_only_set_fields(): + assert server_owned_wif_fields_present({}) == () + assert server_owned_wif_fields_present({"model": "gpt-4o"}) == () + assert server_owned_wif_fields_present( {"anthropic_keycloak_token_url": "https://idp.example/token", "model": "gpt-4o"} ) == ("anthropic_keycloak_token_url",) -def test_anthropic_wif_fields_present_is_derived_from_the_shared_list(): +def test_server_owned_wif_fields_present_is_derived_from_the_shared_list(): """A non-admin persistence gate built on this must automatically cover a field added - later to anthropic_wif_litellm_params, not just the fields known when the gate was + later to server_owned_wif_litellm_params, not just the fields known when the gate was written -- so this must read the shared list rather than a hand-copied one.""" - values = {field: "set" for field in anthropic_wif_litellm_params} - assert set(anthropic_wif_fields_present(values)) == set(anthropic_wif_litellm_params) + values = {field: "set" for field in server_owned_wif_litellm_params} + assert set(server_owned_wif_fields_present(values)) == set(server_owned_wif_litellm_params) -def test_anthropic_wif_fields_named_reports_keys_whatever_their_value(): +def test_server_owned_wif_fields_named_reports_keys_whatever_their_value(): """The credential write gates must see a key a caller sets to ``None``: the federation resolver reacts to the key's presence, not its value, so ``{"anthropic_issuer_url": None}`` wedges every deployment referencing the credential once persisted.""" - assert anthropic_wif_fields_named({}) == () - assert anthropic_wif_fields_named({"model": "gpt-4o"}) == () - assert anthropic_wif_fields_named({"anthropic_issuer_url": None}) == ("anthropic_issuer_url",) - assert anthropic_wif_fields_present({"anthropic_issuer_url": None}) == () - assert anthropic_wif_fields_named(("anthropic_keycloak_token_url", "api_key")) == ("anthropic_keycloak_token_url",) + assert server_owned_wif_fields_named({}) == () + assert server_owned_wif_fields_named({"model": "gpt-4o"}) == () + assert server_owned_wif_fields_named({"anthropic_issuer_url": None}) == ("anthropic_issuer_url",) + assert server_owned_wif_fields_present({"anthropic_issuer_url": None}) == () + assert server_owned_wif_fields_named(("anthropic_keycloak_token_url", "api_key")) == ( + "anthropic_keycloak_token_url", + ) -def test_anthropic_wif_fields_named_is_derived_from_the_shared_list(): - assert set(anthropic_wif_fields_named(frozenset(anthropic_wif_litellm_params))) == set(anthropic_wif_litellm_params) +def test_server_owned_wif_fields_named_is_derived_from_the_shared_list(): + assert set(server_owned_wif_fields_named(frozenset(server_owned_wif_litellm_params))) == set( + server_owned_wif_litellm_params + ) @pytest.mark.parametrize("param_name", ["anthropic_issuer_signing_key_ref", "anthropic_keycloak_client_secret_ref"]) @@ -157,3 +164,39 @@ def test_wif_ref_fields_hold_secret_pointers(param_name: str): @pytest.mark.parametrize("param_name", ["api_key", "anthropic_federation_rule_id", "anthropic_identity_token"]) def test_dereferenced_fields_do_not_hold_secret_pointers(param_name: str): assert not holds_secret_pointer(param_name) + + +def test_credential_litellm_params_declares_every_openai_wif_field(): + for field in openai_wif_litellm_params: + assert field in CredentialLiteLLMParams.model_fields, field + + +def test_openai_wif_fields_round_trip_through_model_dump(): + values = {field: f"value-for-{field}" for field in openai_wif_litellm_params} + + dumped = CredentialLiteLLMParams(**values).model_dump(exclude_none=True) + + for field, value in values.items(): + assert dumped[field] == value, field + + +def test_server_owned_registry_is_anthropic_plus_openai(): + assert server_owned_wif_litellm_params == anthropic_wif_litellm_params + openai_wif_litellm_params + assert set(openai_wif_litellm_params) == { + "openai_identity_provider_id", + "openai_service_account_id", + "openai_identity_token_file", + } + + +def test_server_owned_wif_fields_present_reports_openai_fields(): + assert server_owned_wif_fields_present( + {"openai_identity_token_file": "/var/run/secrets/tokens/openai", "model": "gpt-4o"} + ) == ("openai_identity_token_file",) + assert server_owned_wif_fields_named({"openai_service_account_id": None}) == ("openai_service_account_id",) + + +@pytest.mark.parametrize("param_name", openai_wif_litellm_params) +def test_reject_server_owned_wif_params_names_each_openai_field(param_name: str): + with pytest.raises(ValueError, match=param_name): + reject_server_owned_wif_params({param_name: "client-supplied"}) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 880173a9ef8..678999b2fed 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -29417,6 +29417,12 @@ export interface components { ocr_cost_per_credit?: number | null; /** Ocr Cost Per Page */ ocr_cost_per_page?: number | null; + /** Openai Identity Provider Id */ + openai_identity_provider_id?: string | null; + /** Openai Identity Token File */ + openai_identity_token_file?: string | null; + /** Openai Service Account Id */ + openai_service_account_id?: string | null; /** Organization */ organization?: string | null; /** Otpm */ @@ -39520,6 +39526,12 @@ export interface components { ocr_cost_per_credit?: number | null; /** Ocr Cost Per Page */ ocr_cost_per_page?: number | null; + /** Openai Identity Provider Id */ + openai_identity_provider_id?: string | null; + /** Openai Identity Token File */ + openai_identity_token_file?: string | null; + /** Openai Service Account Id */ + openai_service_account_id?: string | null; /** Organization */ organization?: string | null; /** Otpm */ From 994c6f0d554e7ae2e252b5f171879881d683f2e9 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 3 Sep 2026 12:54:31 -0700 Subject: [PATCH 2/3] fix(openai): keep the get_models env fallback to a missing key only and pin the WIF ban list --- litellm/llms/openai/chat/gpt_transformation.py | 6 ++++-- .../llms/anthropic/test_anthropic_common_utils.py | 6 ++++-- .../llms/openai/test_openai_workload_identity.py | 11 +++++++++++ 3 files changed, 19 insertions(+), 4 deletions(-) diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 5fdce1ed283..3c8e26ff443 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -753,7 +753,9 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): """ Calls OpenAI's `/v1/models` endpoint and returns the list of models. """ - return self._fetch_model_ids(api_base=api_base, bearer_token=api_key or get_secret_str("OPENAI_API_KEY")) + return self._fetch_model_ids( + api_base=api_base, bearer_token=get_secret_str("OPENAI_API_KEY") if api_key is None else api_key + ) def discover_models( self, litellm_params: Mapping[str, object] | None = None @@ -775,7 +777,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): def _fetch_model_ids( api_base: str | None, bearer_token: str | None ) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override - parsed_url: Final = httpx.URL(api_base or "https://api.openai.com") + parsed_url: Final = httpx.URL("https://api.openai.com" if api_base is None else api_base) port_suffix: Final = f":{parsed_url.port}" if parsed_url.port else "" response: Final = litellm.module_level_client.get( url=f"{parsed_url.scheme}://{parsed_url.host}{port_suffix}/v1/models", diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index 7ede9480346..c2dfb91591b 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -3510,9 +3510,11 @@ class TestWifParamsAreNotClientSettable: minted org credential, and the router merges request kwargs over deployment params, so a caller who set it picked the scope instead of the administrator.""" from litellm.proxy.auth.auth_utils import _SERVER_OWNED_WIF_UNCONDITIONAL_BANNED - from litellm.types.utils import anthropic_wif_litellm_params + from litellm.types.utils import anthropic_wif_litellm_params, openai_wif_litellm_params - assert set(anthropic_wif_litellm_params) <= set(_SERVER_OWNED_WIF_UNCONDITIONAL_BANNED) + assert set(_SERVER_OWNED_WIF_UNCONDITIONAL_BANNED) == set(anthropic_wif_litellm_params) | set( + openai_wif_litellm_params + ) class TestWifServerOwnedParamsAreUnconditional: 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 a152a70bcec..1f010ef19ee 100644 --- a/tests/test_litellm/llms/openai/test_openai_workload_identity.py +++ b/tests/test_litellm/llms/openai/test_openai_workload_identity.py @@ -552,6 +552,17 @@ class TestDiscoverModels: assert not exchange_route.called + @respx.mock + def test_empty_static_key_never_borrows_the_env_key(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("OPENAI_API_KEY", "sk-env-key-that-must-stay-home") + foreign_models: Final = respx.get("https://third-party.example/v1/models").mock( + return_value=httpx.Response(200, json={"data": [{"id": "other-model"}]}) + ) + + assert OpenAIGPTConfig().get_models(api_key="", api_base="https://third-party.example") == ["other-model"] + assert foreign_models.calls.last.request.headers["Authorization"] == "Bearer " + + class TestClientsideBaseOverride: def test_client_api_base_override_clears_deployment_wif(self, deployment_wif: dict[str, str]) -> None: from litellm.router_utils.clientside_credential_handler import get_dynamic_litellm_params From fe6a5351eeec18ca40aa06a9097ecabfd25b40d1 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 3 Sep 2026 13:40:18 -0700 Subject: [PATCH 3/3] fix(openai): let workload identity federation target the regional api.openai.com hosts --- litellm/llms/openai/workload_identity.py | 5 ++++- .../openai/test_openai_workload_identity.py | 17 +++++++++++++++++ 2 files changed, 21 insertions(+), 1 deletion(-) diff --git a/litellm/llms/openai/workload_identity.py b/litellm/llms/openai/workload_identity.py index 48e9b779a34..aef5d16807e 100644 --- a/litellm/llms/openai/workload_identity.py +++ b/litellm/llms/openai/workload_identity.py @@ -18,6 +18,7 @@ if TYPE_CHECKING: OPENAI_WIF_CLIENT_ID: Final = "litellm" _OPENAI_API_HOST: Final = "api.openai.com" +_OPENAI_REGIONAL_HOST_SUFFIX: Final = f".{_OPENAI_API_HOST}" _SDK_UPGRADE_MESSAGE: Final = ( "OpenAI workload identity federation requires openai>=2.32.0. " "Upgrade the installed openai package to use OPENAI_IDENTITY_PROVIDER_ID / " @@ -86,7 +87,9 @@ 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 + if parsed.scheme != "https" or parsed.hostname is None: + return False + return parsed.hostname == _OPENAI_API_HOST or parsed.hostname.endswith(_OPENAI_REGIONAL_HOST_SUFFIX) @lru_cache(maxsize=16) 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 1f010ef19ee..7965ca2df0a 100644 --- a/tests/test_litellm/llms/openai/test_openai_workload_identity.py +++ b/tests/test_litellm/llms/openai/test_openai_workload_identity.py @@ -97,6 +97,23 @@ 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("regional_host", ("eu.api.openai.com", "us.api.openai.com")) + def test_regional_openai_api_base_allows( + self, wif_env: OpenAIWorkloadIdentityConfig, regional_host: str + ) -> None: + assert ( + resolve_openai_workload_identity_config(api_key=None, api_base=f"https://{regional_host}/v1") == wif_env + ) + + @pytest.mark.parametrize( + "lookalike_base", + ("https://api.openai.com.evil.example/v1", "https://openai.com/v1", "https://euapi.openai.com/v1"), + ) + def test_openai_lookalike_api_base_disables( + self, wif_env: OpenAIWorkloadIdentityConfig, lookalike_base: str + ) -> None: + assert resolve_openai_workload_identity_config(api_key=None, api_base=lookalike_base) is None + def test_foreign_env_base_url_disables( self, wif_env: OpenAIWorkloadIdentityConfig, monkeypatch: pytest.MonkeyPatch ) -> None: