From cb3f29d91230019720c99679b7812f496bdaf695 Mon Sep 17 00:00:00 2001 From: derhornspieler <15236687+derhornspieler@users.noreply.github.com> Date: Sun, 23 Aug 2026 04:53:49 -0400 Subject: [PATCH] feat(proxy): provider-level setup for Anthropic workload identity federation Federation was configurable but only per model, and only by hand. This adds the provider-level flow the feature was missing, reusing what already exists rather than introducing parallel concepts Provider metadata gains optional credential variants: a selector plus per-variant field lists, so a provider whose credential shape branches describes that in the same metadata every provider already publishes, and the existing generic renderer drives it with no provider-specific branch. Anthropic publishes five, one per way of authenticating. Fields holding a secret reference render as text, not as password inputs that would invite pasting the secret itself Model discovery could not see a credential configured through litellm_params, only through the environment, so a federated provider had nothing to discover. Providers now expose discovery that receives the deployment params, and an admin-only endpoint discovers a provider's models from a stored credential by name. The credential fields are never accepted in that request, since choosing which server-side secret is read is not a caller's decision Enable and disable reuse the existing blocked flag rather than a second notion of the same thing, which meant letting model creation persist it so a model can be created already disabled instead of being created live and then paused. Alternate names reuse the existing public name and group alias. One credential now feeds many models, so switching a credential's auth variant has to be able to remove the previous variant's fields, which credential updates could not express before The Admin UI gains an Add Provider flow over all of it: pick the provider and auth method, fill that method's fields, export the JWKS when the identity is self-signed, discover the provider's models, then enable, rename, and alias them before creating. Two pre-existing UI defects surfaced while building it: a variant-capable provider could render through the stale flat field list before its metadata loaded and leave a validation rule behind that blocked submission, and the credential editor left the previous variant's fields in place when the variant changed --- litellm/llms/anthropic/common_utils.py | 128 ++- .../llms/base_llm/auth/client_credentials.py | 4 +- litellm/llms/base_llm/base_utils.py | 17 + litellm/models/credentials.py | 2 + .../common_utils/credential_hydration.py | 38 + .../proxy/credential_endpoints/endpoints.py | 94 ++- .../model_management_endpoints.py | 130 ++- .../provider_create_fields.json | 260 +++++- .../model_management_endpoints.py | 21 +- .../public_endpoints/public_endpoints.py | 60 +- litellm/types/router.py | 3 + litellm/utils.py | 5 +- .../anthropic/test_anthropic_common_utils.py | 164 ++++ .../credential_endpoints/test_endpoints.py | 317 ++++++- .../test_model_management_endpoints.py | 780 +++++++++++------- .../public_endpoints/test_public_endpoints.py | 138 +++- .../(dashboard)/models-and-endpoints/page.tsx | 6 + .../AddProviderPanel.integration.test.tsx | 203 +++++ .../panels/add-provider/AddProviderPanel.tsx | 364 ++++++++ .../panels/add-provider/ReviewModelsStep.tsx | 129 +++ .../panels/add-provider/WizardSteps.tsx | 207 +++++ .../panels/add-provider/wizardLogic.test.ts | 149 ++++ .../panels/add-provider/wizardLogic.ts | 128 +++ .../provider_credential_variants.test.ts | 152 ++++ .../add_model/provider_credential_variants.ts | 49 ++ .../provider_specific_fields.test.tsx | 168 ++++ .../add_model/provider_specific_fields.tsx | 220 +++-- .../model_add/CredentialModal.test.tsx | 74 +- .../components/model_add/CredentialModal.tsx | 9 +- .../components/model_add/CredentialsPanel.tsx | 7 +- .../model_add/credential_form_helpers.test.ts | 30 +- .../model_add/credential_form_helpers.ts | 19 + .../src/components/networking.tsx | 74 ++ 33 files changed, 3743 insertions(+), 406 deletions(-) create mode 100644 litellm/proxy/common_utils/credential_hydration.py create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/AddProviderPanel.integration.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/AddProviderPanel.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/ReviewModelsStep.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/WizardSteps.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/wizardLogic.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/wizardLogic.ts create mode 100644 ui/litellm-dashboard/src/components/add_model/provider_credential_variants.test.ts create mode 100644 ui/litellm-dashboard/src/components/add_model/provider_credential_variants.ts diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 23bbe51907b..1f0d7a9355d 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -10,7 +10,7 @@ from types import MappingProxyType from typing import Any, ClassVar, Final, Literal import httpx -from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError import litellm from litellm.constants import DEFAULT_MODEL_CREATED_AT_TIME @@ -128,6 +128,72 @@ class AnthropicError(BaseLLMException): super().__init__(status_code=status_code, message=message, headers=headers) +_MODEL_LIST_PAGE_CAP: Final = 20 + + +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 _AnthropicModelListEntry(BaseModel): + id: str + + +class _AnthropicModelsPage(BaseModel): + data: Sequence[_AnthropicModelListEntry] = Field(default_factory=tuple) + has_more: bool = False + last_id: str | None = None + + +def _sanitized_anthropic_error(response: httpx.Response, detail: str | None = None) -> str: + """A provider error detail built only from structured fields, never ``response.text`` + verbatim: the raw body is untrusted content the caller of ``/v1/models`` did not ask for + and should not have echoed back to it wholesale.""" + if detail is not None: + return f"HTTP {response.status_code}: {detail}" + try: + body: Final = response.json() + except ValueError: + return f"HTTP {response.status_code}" + error: Final = body.get("error") if isinstance(body, dict) else None + message: Final = error.get("message") if isinstance(error, dict) else None + return f"HTTP {response.status_code}: {message}" if isinstance(message, str) else f"HTTP {response.status_code}" + + +def _fetch_anthropic_models_page( + api_base: str, headers: Mapping[str, str], after_id: str | None +) -> _AnthropicModelsPage: + response: Final = litellm.module_level_client.get( + url=f"{api_base}/v1/models", + headers=headers, + params=MappingProxyType({"after_id": after_id}) if after_id else MappingProxyType({}), + follow_redirects=False, + ) + try: + response.raise_for_status() + except httpx.HTTPStatusError: + raise Exception(f"Failed to fetch models from Anthropic. {_sanitized_anthropic_error(response)}") from None + try: + return _AnthropicModelsPage.model_validate(response.json()) + except ValueError as e: + raise Exception( + f"Failed to fetch models from Anthropic. {_sanitized_anthropic_error(response, detail=str(e))}" + ) from None + + +def _fetch_anthropic_model_ids( + api_base: str, headers: Mapping[str, str], after_id: str | None, pages_left: int +) -> tuple[str, ...]: + if pages_left <= 0: + raise Exception(f"Anthropic /v1/models did not terminate within {_MODEL_LIST_PAGE_CAP} pages.") + page: Final = _fetch_anthropic_models_page(api_base, headers, after_id) + page_ids: Final = tuple(entry.id for entry in page.data) + if not page.has_more or page.last_id is None: + return page_ids + return page_ids + _fetch_anthropic_model_ids(api_base, headers, page.last_id, pages_left - 1) + + class AnthropicModelInfo(BaseLLMModelInfo): _workload_identity_eligible: ClassVar[bool] = True @@ -899,39 +965,49 @@ class AnthropicModelInfo(BaseLLMModelInfo): return model.replace("anthropic/", "") if model else None def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: - api_base = AnthropicModelInfo.get_api_base(api_base) - auth_header: Final = AnthropicModelInfo.get_auth_header( - api_key, api_base, allow_workload_identity=config_allows_workload_identity(self) + return self._list_models(api_key=api_key, api_base=api_base, litellm_params=None) + + 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 + """Live discovery for a configured deployment: unlike ``get_models``, this threads the + full ``litellm_params`` into ``get_auth_header`` so a workload-identity-federation source + configured on the deployment (rather than the environment) is honored, gated the same way + every other Anthropic auth surface is via ``config_allows_workload_identity``.""" + return self._list_models( + api_key=_litellm_params_str(litellm_params, "api_key"), + api_base=_litellm_params_str(litellm_params, "api_base"), + litellm_params=litellm_params, ) - if api_base is None or auth_header is None: + + def _list_models( + self, + *, + api_key: str | None, + api_base: str | None, + litellm_params: Mapping[str, object] | None, + ) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override + resolved_api_base: Final = AnthropicModelInfo.get_api_base(api_base) + auth_header: Final = AnthropicModelInfo.get_auth_header( + api_key, + resolved_api_base, + litellm_params=litellm_params, + allow_workload_identity=config_allows_workload_identity(self), + ) + if resolved_api_base is None or auth_header is None: raise ValueError( "ANTHROPIC_API_BASE/ANTHROPIC_BASE_URL or ANTHROPIC_API_KEY/ANTHROPIC_AUTH_TOKEN (or workload " "identity federation via ANTHROPIC_FEDERATION_RULE_ID/ANTHROPIC_ORGANIZATION_ID/" "ANTHROPIC_IDENTITY_TOKEN_FILE) is not set. Please set the environment variable, to query " "Anthropic's `/models` endpoint." ) - headers: Final = {"anthropic-version": "2023-06-01"} - headers.update(auth_header) - response: Final = litellm.module_level_client.get( - url=f"{api_base}/v1/models", - headers=headers, + headers: Final = MappingProxyType({"anthropic-version": "2023-06-01", **auth_header}) + model_ids: Final = _fetch_anthropic_model_ids( + resolved_api_base, headers, after_id=None, pages_left=_MODEL_LIST_PAGE_CAP ) - - try: - response.raise_for_status() - except httpx.HTTPStatusError: - raise Exception( - f"Failed to fetch models from Anthropic. Status code: {response.status_code}, Response: {response.text}" - ) - - models: Final = response.json()["data"] - - litellm_model_names: Final = [] - for model in models: - stripped_model_name = model["id"] - litellm_model_name = "anthropic/" + stripped_model_name - litellm_model_names.append(litellm_model_name) - return litellm_model_names + return [ # mutable-ok: matches get_models' list[str] contract shared by every provider override + "anthropic/" + model_id for model_id in model_ids + ] def get_token_counter(self) -> BaseTokenCounter | None: """ diff --git a/litellm/llms/base_llm/auth/client_credentials.py b/litellm/llms/base_llm/auth/client_credentials.py index 1b70a0a1253..f59412c54ff 100644 --- a/litellm/llms/base_llm/auth/client_credentials.py +++ b/litellm/llms/base_llm/auth/client_credentials.py @@ -159,9 +159,7 @@ def fetch_keycloak_assertion( workload assertion; the caller must not cache the result -- see the module docstring.""" match validate_token_endpoint_url(config.token_url): case InsecureTokenUrl(host=host): - raise ValueError( - f"keycloak token_url must use https; refusing to send the client secret to host {host!r}" - ) + raise ValueError(f"keycloak token_url must use https; refusing to send the client secret to host {host!r}") case _: pass client_secret: Final = _resolve_client_secret(config, secret_reader) diff --git a/litellm/llms/base_llm/base_utils.py b/litellm/llms/base_llm/base_utils.py index c5290b41f7b..bf93308974c 100644 --- a/litellm/llms/base_llm/base_utils.py +++ b/litellm/llms/base_llm/base_utils.py @@ -5,6 +5,7 @@ Utility functions for base LLM classes. import copy import json from abc import ABC, abstractmethod +from collections.abc import Mapping from typing import Any, Final from openai.lib import _parsing, _pydantic @@ -57,6 +58,22 @@ class BaseLLMModelInfo(ABC): """ return [] + 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 + """ + Live model discovery for a configured deployment. Defaults to the api_key/api_base + facade every provider already implements via ``get_models``; a provider whose + discovery needs more of ``litellm_params`` (e.g. Anthropic's workload identity + federation) overrides this instead of widening ``get_models`` for every provider. + """ + api_key: Final = litellm_params.get("api_key") if litellm_params is not None else None + api_base: Final = litellm_params.get("api_base") if litellm_params is not None else None + return self.get_models( + api_key=api_key if isinstance(api_key, str) else None, + api_base=api_base if isinstance(api_base, str) else None, + ) + @staticmethod @abstractmethod def get_api_key(api_key: str | None = None) -> str | None: diff --git a/litellm/models/credentials.py b/litellm/models/credentials.py index 56836234898..9ad85328d93 100644 --- a/litellm/models/credentials.py +++ b/litellm/models/credentials.py @@ -15,6 +15,8 @@ class CredentialBase(BaseModel): class CredentialItem(CredentialBase): credential_values: dict + # PATCH-only: keys to drop from the stored credential_values. + credential_values_to_delete: tuple[str, ...] | None = None class CreateCredentialItem(CredentialBase): diff --git a/litellm/proxy/common_utils/credential_hydration.py b/litellm/proxy/common_utils/credential_hydration.py new file mode 100644 index 00000000000..eb811a9292c --- /dev/null +++ b/litellm/proxy/common_utils/credential_hydration.py @@ -0,0 +1,38 @@ +"""Shared helper for resolving a named Credential's values server-side. + +Memory first (``litellm.credential_list``, already decrypted -- matching +``CredentialAccessor.get_credential_values``), then a DB decrypt fallback for a pod whose +in-memory list has not yet picked up a credential another pod just wrote or updated. +""" + +from types import MappingProxyType +from typing import Final + +import litellm +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper +from litellm.proxy.utils import PrismaClient +from litellm.repositories.credentials_repository import CredentialsRepository +from litellm.types.utils import CredentialItem + + +async def hydrate_named_credential( + credential_name: str, + prisma_client: PrismaClient | None, +) -> CredentialItem | None: + for credential in litellm.credential_list: + if credential.credential_name == credential_name: + return credential + if prisma_client is None: + return None + db_credential: Final = await CredentialsRepository(prisma_client).find_by_name(credential_name) + if db_credential is None: + return None + decrypted_values: Final = MappingProxyType({ + key: decrypt_value_helper(value=value, key=key) or value + for key, value in db_credential.credential_values.items() + }) + return CredentialItem( + credential_name=db_credential.credential_name, + credential_values=decrypted_values, + credential_info=db_credential.credential_info, + ) diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index 3b3e9692eda..a5331ab0c21 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -10,8 +10,19 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.litellm_logging import _get_masked_values -from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth +from litellm.llms.anthropic.wif import ( + _IDENTITY_SOURCE_PARAM, # pyright: ignore[reportPrivateUsage] # one canonical param name, shared with the litellm_params identity-source resolver + _INTERNAL_ISSUER_FIELD_MAP, # pyright: ignore[reportPrivateUsage] # one canonical field map, shared with the litellm_params identity-source resolver + _build_variant, # pyright: ignore[reportPrivateUsage] # one canonical builder, shared with the litellm_params identity-source resolver +) +from litellm.llms.base_llm.auth.identity_source import ( + AnthropicIdentitySourceKind, + InternalIssuerSource, +) +from litellm.llms.base_llm.auth.internal_issuer import internal_issuer_jwks_document +from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.credential_hydration import hydrate_named_credential 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 @@ -87,7 +98,7 @@ async def create_credential( credential_info=credential.credential_info, ) encrypted_credential: Final = CredentialHelperUtils.encrypt_credential_values(processed_credential) - credentials_dict: Final = encrypted_credential.model_dump() + credentials_dict: Final = encrypted_credential.model_dump(exclude_none=True) credentials_dict_jsonified: Final = jsonify_object(credentials_dict) await CredentialsRepository(prisma_client).create( data={ @@ -170,6 +181,70 @@ async def get_credential_by_name( raise handle_exception_on_proxy(e) +@router.get( + "/credentials/{credential_name:path}/jwks", + dependencies=(Depends(user_api_key_auth),), + tags=["credential management"], # mutable-ok: FastAPI's include_router does self.tags.copy(), needs a real list +) +async def get_credential_internal_issuer_jwks( + credential_name: str = Path(..., description="The credential name, percent-decoded; may contain slashes"), + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI resolves the dependency from the default +): + """ + Export the public JWKS for an anthropic ``internal_issuer`` credential, so the operator can + register it on the Anthropic federation issuer from the UI. Never touches the private signing + key: only its derived public JWKS leaves this process. 404s for any other credential shape. + """ + from litellm.proxy.proxy_server import prisma_client + + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail={ # mutable-ok: starlette json.dumps()s HTTPException.detail raw, needs a real dict + "error": "Only proxy admins can export a credential's JWKS." + }, + ) + + try: + credential: Final = await hydrate_named_credential(credential_name, prisma_client) + if credential is None or credential.credential_info.get("custom_llm_provider") != "anthropic": + raise HTTPException( + status_code=404, + detail={ # mutable-ok: starlette json.dumps()s HTTPException.detail raw, needs a real dict + "error": f"No anthropic credential named {credential_name!r}." + }, + ) + configured_source: Final = credential.credential_values.get(_IDENTITY_SOURCE_PARAM) + if configured_source != AnthropicIdentitySourceKind.internal_issuer.value: + raise HTTPException( + status_code=404, + detail={ # mutable-ok: starlette json.dumps()s HTTPException.detail raw, needs a real dict + "error": ( + f"Credential {credential_name!r} is not configured with " + f"{_IDENTITY_SOURCE_PARAM}={AnthropicIdentitySourceKind.internal_issuer.value!r}." + ) + }, + ) + try: + issuer_source: Final = _build_variant( + InternalIssuerSource, credential.credential_values, _INTERNAL_ISSUER_FIELD_MAP + ) + jwks_document: Final = internal_issuer_jwks_document(issuer_source) + except (litellm.AuthenticationError, ValueError) as e: + raise HTTPException( + status_code=400, + detail={ # mutable-ok: starlette json.dumps()s HTTPException.detail raw, needs a real dict + "error": str(e) + }, + ) from e + return Response(content=jwks_document, media_type="application/json") + except HTTPException: + raise + except Exception as e: # noqa: BLE001 # endpoint boundary: every failure becomes the proxy's error contract + verbose_proxy_logger.exception(e) + raise handle_exception_on_proxy(e) + + @router.get( "/credentials/by_model/{model_id}", dependencies=[Depends(user_api_key_auth)], @@ -272,6 +347,9 @@ def update_db_credential( merged_credential.credential_values.update(encrypted_params) + for key in updated_patch.credential_values_to_delete or (): + merged_credential.credential_values.pop(key, None) + # update model info if encrypted_credential.credential_info: """Update credential info""" @@ -300,6 +378,14 @@ async def update_credential( from litellm.proxy.proxy_server import prisma_client try: + overlap: Final = frozenset(credential.credential_values) & frozenset( + credential.credential_values_to_delete or () + ) + if overlap: + raise HTTPException( + status_code=400, + detail=f"credential_values_to_delete overlaps credential_values for key(s): {sorted(overlap)}", + ) if prisma_client is None: raise HTTPException( status_code=500, @@ -310,7 +396,7 @@ async def update_credential( if db_credential is None: raise HTTPException(status_code=404, detail="Credential not found in DB.") merged_credential: Final = update_db_credential(db_credential, credential) - credential_object_jsonified: Final = jsonify_object(merged_credential.model_dump()) + credential_object_jsonified: Final = jsonify_object(merged_credential.model_dump(exclude_none=True)) await credentials_repository.update_by_name( credential_name, data={ @@ -331,6 +417,8 @@ async def update_credential( in_memory_values: Final = dict(existing_in_memory.credential_values or {}) if credential.credential_values: in_memory_values.update(credential.credential_values) + for key in credential.credential_values_to_delete or (): + in_memory_values.pop(key, None) in_memory_info: Final = dict(existing_in_memory.credential_info or {}) if credential.credential_info: in_memory_info.update(credential.credential_info) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 217fc61a56c..fe2f61a910f 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -49,11 +49,13 @@ from litellm.proxy._types import ( TeamModelDeleteRequest, UserAPIKeyAuth, ) +from litellm.proxy.auth.auth_utils import reject_server_owned_wif_params from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.config_sync_pubsub import ( coordination_redis_cache, publish_config_change, ) +from litellm.proxy.common_utils.credential_hydration import hydrate_named_credential from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.management_endpoints.common_utils import _is_user_team_admin @@ -89,6 +91,8 @@ from litellm.router_utils.auto_router_model_naming import ( ) from litellm.types.proxy.management_endpoints.model_management_endpoints import ( AutoRouterClassifierDefaultPromptResponse, + ProviderModelDiscoveryRequest, + ProviderModelDiscoveryResponse, UpdateUsefulLinksRequest, ) from litellm.types.router import ( @@ -98,7 +102,8 @@ from litellm.types.router import ( ModelInfo, updateDeployment, ) -from litellm.utils import get_utc_datetime +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager, get_utc_datetime router: Final = APIRouter() @@ -914,6 +919,8 @@ async def _add_model_to_db( } if model_params.model_info.id is not None: _data["model_id"] = model_params.model_info.id + if model_params.blocked is not None: + _data["blocked"] = model_params.blocked if should_create_model_in_db: model_response = await ModelRepository(prisma_client).table.create(data=_data) else: @@ -1657,6 +1664,116 @@ async def delete_team_model_alias( return removed_model_aliases +def _reject_inline_secret_reference(value: str | None, field: str) -> None: + if value is not None and value.startswith(("os.environ/", "oidc/")): + raise HTTPException( + status_code=400, + detail={ # mutable-ok: starlette json.dumps()s HTTPException.detail raw, needs a real dict + "error": f"{field} may not be an os.environ/ or oidc/ reference in a request body." + }, + ) + + +async def _resolve_discovery_litellm_params( + data: ProviderModelDiscoveryRequest, + prisma_client: PrismaClient | None, +) -> Mapping[str, object]: + if data.litellm_credential_name is None: + return MappingProxyType({k: v for k, v in (("api_key", data.api_key), ("api_base", data.api_base)) if v}) + + credential: Final = await hydrate_named_credential(data.litellm_credential_name, prisma_client) + if credential is None: + raise HTTPException( + status_code=404, + detail={ # mutable-ok: starlette json.dumps()s HTTPException.detail raw, needs a real dict + "error": f"Credential {data.litellm_credential_name!r} not found." + }, + ) + credential_provider: Final = credential.credential_info.get("custom_llm_provider") + if credential_provider is not None and credential_provider != data.custom_llm_provider: + raise HTTPException( + status_code=400, + detail={ # mutable-ok: starlette json.dumps()s HTTPException.detail raw, needs a real dict + "error": ( + f"Credential {data.litellm_credential_name!r} is configured for provider " + f"{credential_provider!r}, not {data.custom_llm_provider!r}." + ) + }, + ) + if data.api_key is None: + return MappingProxyType(dict(credential.credential_values)) + return MappingProxyType({**credential.credential_values, "api_key": data.api_key}) + + +@router.post( + "/provider/models/discover", + description="Live model discovery for a configured provider credential. Proxy-admin only.", + tags=["model management"], # mutable-ok: fastapi's decorator signature types tags as a list + dependencies=[Depends(user_api_key_auth)], # mutable-ok: fastapi's decorator signature types dependencies as a list + response_model=ProviderModelDiscoveryResponse, +) +async def discover_provider_models( + data: ProviderModelDiscoveryRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI resolves the dependency from the default +) -> ProviderModelDiscoveryResponse: + from litellm.proxy.proxy_server import prisma_client + + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise ProxyException( + message="Only proxy admins can discover provider models.", + type=ProxyErrorTypes.auth_error.value, + code=status.HTTP_403_FORBIDDEN, + param=None, + ) + + reject_server_owned_wif_params(data.model_dump(exclude_none=True)) + _reject_inline_secret_reference(data.api_key, "api_key") + _reject_inline_secret_reference(data.api_base, "api_base") + + if data.litellm_credential_name is not None and data.api_base is not None: + raise HTTPException( + status_code=400, + detail={ # mutable-ok: starlette json.dumps()s HTTPException.detail raw, needs a real dict + "error": "api_base cannot be combined with litellm_credential_name." + }, + ) + + try: + provider_enum: Final = LlmProviders(data.custom_llm_provider) + except ValueError: + raise HTTPException( + status_code=400, + detail={ # mutable-ok: starlette json.dumps()s HTTPException.detail raw, needs a real dict + "error": f"Unknown provider: {data.custom_llm_provider!r}" + }, + ) from None + + provider_config: Final = ProviderConfigManager.get_provider_model_info(model=None, provider=provider_enum) + if provider_config is None: + raise HTTPException( + status_code=400, + detail={ # mutable-ok: starlette json.dumps()s HTTPException.detail raw, needs a real dict + "error": f"Provider {data.custom_llm_provider!r} does not support model discovery." + }, + ) + + litellm_params: Final = await _resolve_discovery_litellm_params(data, prisma_client) + + try: + models: Final = await asyncio.to_thread(provider_config.discover_models, litellm_params) + except HTTPException: + raise + except Exception as e: + raise HTTPException( + status_code=502, + detail={ # mutable-ok: starlette json.dumps()s HTTPException.detail raw, needs a real dict + "error": f"Model discovery failed: {e}" + }, + ) from e + + return ProviderModelDiscoveryResponse(models=models) + + #### [BETA] - This is a beta endpoint, format might change based on user feedback. - https://github.com/BerriAI/litellm/issues/964 @router.post( "/model/new", @@ -1728,6 +1845,17 @@ async def add_new_model( premium_user=premium_user, ) + # Same proxy-admin-only rule patch_model applies to the blocked flag: a team admin + # passed the check above for a team-scoped model, but must not be able to create it + # already paused (or explicitly unpaused) out from under the proxy admin. + if model_params.blocked is not None and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise ProxyException( + message="Only proxy admins can set a model's blocked flag.", + type=ProxyErrorTypes.auth_error.value, + code=status.HTTP_403_FORBIDDEN, + param="blocked", + ) + _raise_on_strategy_router_write_violation( incoming_params=model_params.litellm_params, existing_params=None, diff --git a/litellm/proxy/public_endpoints/provider_create_fields.json b/litellm/proxy/public_endpoints/provider_create_fields.json index 4652719a23b..2a640bfc6e1 100644 --- a/litellm/proxy/public_endpoints/provider_create_fields.json +++ b/litellm/proxy/public_endpoints/provider_create_fields.json @@ -343,7 +343,265 @@ "default_value": null } ], - "default_model_placeholder": "claude-3-opus" + "default_model_placeholder": "claude-3-opus", + "credential_variants": { + "selector_label": "Authentication method", + "default_variant": "api_key", + "field_definitions": [ + { + "key": "api_base", + "label": "Upstream API Base", + "placeholder": "https://api.anthropic.com", + "tooltip": "Optional. Where the proxy forwards requests upstream. Leave blank to use Anthropic's public API. Set this only for private Anthropic deployments or reverse proxies. Do NOT set this to your LiteLLM proxy URL — that causes a recursive loop.", + "required": false, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "api_key", + "label": "API Key", + "placeholder": "sk-", + "tooltip": "Leave empty for BYOK (bring-your-own-key) flows, where clients forward their own Anthropic key via the x-api-key header. Requires the 'Forward LLM provider auth headers' UI setting to be enabled.", + "required": false, + "field_type": "password", + "options": null, + "default_value": null + }, + { + "key": "anthropic_federation_rule_id", + "label": "Federation Rule ID", + "placeholder": null, + "tooltip": "The workload identity federation rule id, created in the Anthropic Console.", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "anthropic_organization_id", + "label": "Organization ID", + "placeholder": null, + "tooltip": "The Anthropic organization id the federation rule belongs to.", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "anthropic_service_account_id", + "label": "Service Account ID", + "placeholder": null, + "tooltip": "Optional. Required only if the federation rule is scoped to more than one service account.", + "required": false, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "anthropic_workspace_id", + "label": "Workspace ID", + "placeholder": null, + "tooltip": "Optional. Set this if the federation rule is scoped to a workspace.", + "required": false, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "anthropic_identity_token", + "label": "Identity Token Reference", + "placeholder": "oidc/env/MY_TOKEN_VAR", + "tooltip": "A secret reference to the workload's OIDC identity token, e.g. oidc/env/VAR_NAME, oidc/github/, or oidc/google/. This is a REFERENCE, never the token itself.", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "anthropic_identity_token_file", + "label": "Identity Token File Path", + "placeholder": "/var/run/secrets/tokens/oidc-token", + "tooltip": "Absolute path to a mounted file containing the workload's OIDC identity token.", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "anthropic_issuer_url", + "label": "Issuer URL", + "placeholder": "https://issuer.example.com", + "tooltip": "The 'iss' claim LiteLLM signs into the self-issued workload assertion. Must match the issuer registered on the Anthropic federation rule.", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "anthropic_issuer_subject", + "label": "Subject", + "placeholder": null, + "tooltip": "The 'sub' claim LiteLLM signs into the self-issued workload assertion.", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "anthropic_issuer_audience", + "label": "Audience", + "placeholder": null, + "tooltip": "Optional. The 'aud' claim LiteLLM signs into the self-issued workload assertion.", + "required": false, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "anthropic_issuer_ttl_seconds", + "label": "Assertion TTL (seconds)", + "placeholder": "300", + "tooltip": "Optional. How long each freshly minted assertion is valid for, up to 3600 seconds.", + "required": false, + "field_type": "text", + "options": null, + "default_value": "300" + }, + { + "key": "anthropic_issuer_signing_key_ref", + "label": "Signing Key Reference", + "placeholder": "os.environ/ISSUER_SIGNING_KEY_PEM", + "tooltip": "A secret reference to the ES256 (P-256) private key PEM LiteLLM signs assertions with, e.g. os.environ/VAR_NAME. This is a REFERENCE, never the key itself.", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "anthropic_keycloak_token_url", + "label": "Keycloak Token URL", + "placeholder": "https://keycloak.example.com/realms/my-realm/protocol/openid-connect/token", + "tooltip": "The Keycloak client_credentials token endpoint LiteLLM fetches the workload assertion from.", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "anthropic_keycloak_client_id", + "label": "Keycloak Client ID", + "placeholder": null, + "tooltip": "The confidential client id LiteLLM authenticates to Keycloak as.", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "anthropic_keycloak_auth_method", + "label": "Keycloak Client Auth Method", + "placeholder": null, + "tooltip": "How the client secret is presented to Keycloak's token endpoint.", + "required": false, + "field_type": "select", + "options": ["client_secret_basic", "client_secret_post"], + "default_value": "client_secret_basic" + }, + { + "key": "anthropic_keycloak_client_secret_ref", + "label": "Keycloak Client Secret Reference", + "placeholder": "os.environ/KEYCLOAK_CLIENT_SECRET", + "tooltip": "A secret reference to the Keycloak client secret, e.g. os.environ/VAR_NAME. This is a REFERENCE, never the secret itself.", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "anthropic_keycloak_scope", + "label": "Keycloak Scope", + "placeholder": null, + "tooltip": "Optional. The OAuth scope requested from Keycloak's token endpoint.", + "required": false, + "field_type": "text", + "options": null, + "default_value": null + } + ], + "variants": [ + { + "id": "api_key", + "label": "API Key", + "field_keys": ["api_base", "api_key"], + "fixed_values": {} + }, + { + "id": "wif_token", + "label": "Workload Identity Federation (external token)", + "field_keys": [ + "api_base", + "anthropic_federation_rule_id", + "anthropic_organization_id", + "anthropic_service_account_id", + "anthropic_workspace_id", + "anthropic_identity_token" + ], + "fixed_values": {} + }, + { + "id": "wif_token_file", + "label": "Workload Identity Federation (token file)", + "field_keys": [ + "api_base", + "anthropic_federation_rule_id", + "anthropic_organization_id", + "anthropic_service_account_id", + "anthropic_workspace_id", + "anthropic_identity_token_file" + ], + "fixed_values": {} + }, + { + "id": "wif_internal_issuer", + "label": "Workload Identity Federation (LiteLLM-signed)", + "field_keys": [ + "api_base", + "anthropic_federation_rule_id", + "anthropic_organization_id", + "anthropic_service_account_id", + "anthropic_workspace_id", + "anthropic_issuer_url", + "anthropic_issuer_subject", + "anthropic_issuer_audience", + "anthropic_issuer_ttl_seconds", + "anthropic_issuer_signing_key_ref" + ], + "fixed_values": { + "anthropic_identity_source": "internal_issuer" + } + }, + { + "id": "wif_keycloak", + "label": "Workload Identity Federation (Keycloak)", + "field_keys": [ + "api_base", + "anthropic_federation_rule_id", + "anthropic_organization_id", + "anthropic_service_account_id", + "anthropic_workspace_id", + "anthropic_keycloak_token_url", + "anthropic_keycloak_client_id", + "anthropic_keycloak_auth_method", + "anthropic_keycloak_client_secret_ref", + "anthropic_keycloak_scope" + ], + "fixed_values": { + "anthropic_identity_source": "keycloak" + } + } + ] + } }, { "provider": "ANTHROPIC_TEXT", diff --git a/litellm/types/proxy/management_endpoints/model_management_endpoints.py b/litellm/types/proxy/management_endpoints/model_management_endpoints.py index 6e18787a224..a16e03e199a 100644 --- a/litellm/types/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/types/proxy/management_endpoints/model_management_endpoints.py @@ -1,6 +1,7 @@ +from collections.abc import Sequence from typing import Any -from pydantic import BaseModel, Field +from pydantic import BaseModel, ConfigDict, Field from ...router import ModelGroupInfo @@ -61,3 +62,21 @@ class AccessGroupInfo(BaseModel): class ListAccessGroupsResponse(BaseModel): access_groups: list[AccessGroupInfo] + + +class ProviderModelDiscoveryRequest(BaseModel): + """Body for POST /provider/models/discover. Exactly one of ``litellm_credential_name`` or + inline ``api_key``/``api_base`` names the credential to probe; extra fields are rejected + outright rather than silently ignored, since this is a security boundary -- see + ``reject_server_owned_wif_params`` in the handler for why.""" + + custom_llm_provider: str + litellm_credential_name: str | None = None + api_key: str | None = None + api_base: str | None = None + + model_config = ConfigDict(extra="forbid") + + +class ProviderModelDiscoveryResponse(BaseModel): + models: Sequence[str] diff --git a/litellm/types/proxy/public_endpoints/public_endpoints.py b/litellm/types/proxy/public_endpoints/public_endpoints.py index f6ee054ceaa..cabc047816c 100644 --- a/litellm/types/proxy/public_endpoints/public_endpoints.py +++ b/litellm/types/proxy/public_endpoints/public_endpoints.py @@ -1,7 +1,8 @@ from collections.abc import Mapping -from typing import Any, Literal +from typing import Any, Final, Literal -from pydantic import BaseModel +from pydantic import BaseModel, Field, model_validator +from typing_extensions import Self class PublicModelHubInfo(BaseModel): @@ -25,12 +26,67 @@ class ProviderCredentialField(BaseModel): default_value: str | None = None +class ProviderCredentialVariant(BaseModel): + """One selectable auth shape for a provider, e.g. 'api_key' or 'workload identity + federation, Keycloak'. ``field_keys`` names entries in the parent + ``ProviderCredentialVariants.field_definitions`` to mount when this variant is active; + ``fixed_values`` are litellm_params values the variant implies (e.g. a discriminator like + ``anthropic_identity_source: keycloak``) and are submitted without a form field for them.""" + + id: str + label: str + field_keys: tuple[str, ...] + fixed_values: Mapping[str, str] = Field(default_factory=dict) + + +def _validate_variant(variant: "ProviderCredentialVariant", defined_keys: frozenset[str]) -> None: + """Each variant may only reference declared fields, and a fixed value may not also be a field the + form would mount, or the form and the payload would disagree about who owns that key.""" + unresolved: Final = tuple(key for key in variant.field_keys if key not in defined_keys) + if unresolved: + raise ValueError(f"variant {variant.id!r} references undefined field_keys: {unresolved}") + overlap: Final = sorted(frozenset(variant.field_keys) & frozenset(variant.fixed_values)) + if overlap: + raise ValueError(f"variant {variant.id!r} has fixed_values overlapping field_keys: {overlap}") + + +class ProviderCredentialVariants(BaseModel): + """Declares selectable auth variants for a provider whose credential shape branches (e.g. + API key vs. one of several workload-identity-federation identity sources), so the field + list itself depends on a choice the form makes, not just which fields are shown. + + ``field_definitions`` is the full pool of fields any variant may reference by key; + a UI mounts only the active variant's ``field_keys``, keeping every other field both + unmounted and unsubmitted. ``default_variant`` seeds the selector on a fresh form. + """ + + selector_label: str + default_variant: str + field_definitions: tuple[ProviderCredentialField, ...] + variants: tuple[ProviderCredentialVariant, ...] + + @model_validator(mode="after") + def _validate_structure(self) -> Self: + variant_ids: Final = tuple(variant.id for variant in self.variants) + if len(variant_ids) != len(frozenset(variant_ids)): + raise ValueError(f"credential_variants.variants ids must be unique, got {variant_ids}") + if self.default_variant not in variant_ids: + raise ValueError( + f"credential_variants.default_variant {self.default_variant!r} is not one of {variant_ids}" + ) + defined_keys: Final = frozenset(field.key for field in self.field_definitions) + for variant in self.variants: + _validate_variant(variant, defined_keys) + return self + + class ProviderCreateInfo(BaseModel): provider: str provider_display_name: str litellm_provider: str credential_fields: list[ProviderCredentialField] default_model_placeholder: str | None = None + credential_variants: ProviderCredentialVariants | None = None class AgentCredentialField(BaseModel): diff --git a/litellm/types/router.py b/litellm/types/router.py index 5f331009754..7d007b98498 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -536,6 +536,9 @@ class Deployment(BaseModel): model_name: str litellm_params: LiteLLM_Params model_info: ModelInfo + # admin-toggled pause flag; mirrors LiteLLM_ProxyModelTable.blocked. None means "don't set it + # on create" -- the Prisma column defaults to False -- rather than "explicitly unblocked". + blocked: bool | None = None model_config = ConfigDict(extra="allow", protected_namespaces=()) diff --git a/litellm/utils.py b/litellm/utils.py index e5ce7157e77..0025d97deee 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7207,9 +7207,8 @@ def _get_valid_models_from_provider_api( if cached_result is not None: return cached_result - models: Final = provider_config.get_models( - api_key=litellm_params.api_key if litellm_params is not None else None, - api_base=litellm_params.api_base if litellm_params is not None else None, + models: Final = provider_config.discover_models( + litellm_params=litellm_params.model_dump(exclude_none=True) if litellm_params is not None else None ) _model_cache.set_cached_model_info(custom_llm_provider, litellm_params, 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 01907901f75..24f0dea8c05 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -2954,6 +2954,170 @@ class TestWifProviderAllowlist: assert poster.requests == [] +def _models_page_response(page: dict, status_code: int = 200): + import httpx + + return httpx.Response(status_code, json=page, request=httpx.Request("GET", "https://api.anthropic.com/v1/models")) + + +class RecordingModelsClient: + """Records every call and answers with the queued responses in order, cycling the last one + once exhausted so a runaway pagination loop degrades to a repeated page rather than an + IndexError, letting the page-cap test observe the cap firing instead of a test bug.""" + + def __init__(self, pages: list[dict] | None = None, responses=None): + self.calls = [] + self._responses = responses if responses is not None else [_models_page_response(page) for page in pages] + + def get(self, url, headers=None, params=None, follow_redirects=None, timeout=None): + self.calls.append(SimpleNamespace(url=url, headers=headers, params=params, follow_redirects=follow_redirects)) + index = min(len(self.calls) - 1, len(self._responses) - 1) + return self._responses[index] + + +class TestModelDiscovery: + """AnthropicModelInfo.get_models / discover_models: pagination, redirect refusal, and + sanitized errors on the upstream Anthropic /v1/models call itself (issue #28607 gap: a + WIF source configured in litellm_params, rather than the environment, could not + discover).""" + + def test_get_models_paginates_via_has_more_and_last_id(self, monkeypatch, clean_anthropic_env): + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + client = RecordingModelsClient( + [ + {"data": [{"id": "claude-a"}, {"id": "claude-b"}], "has_more": True, "last_id": "claude-b"}, + {"data": [{"id": "claude-c"}], "has_more": False, "last_id": "claude-c"}, + ] + ) + monkeypatch.setattr("litellm.module_level_client", client) + + models = AnthropicModelInfo().get_models(api_base="https://api.anthropic.com") + + assert models == ["anthropic/claude-a", "anthropic/claude-b", "anthropic/claude-c"] + assert len(client.calls) == 2 + assert client.calls[0].params == {} + assert client.calls[1].params == {"after_id": "claude-b"} + + def test_get_models_refuses_to_follow_redirects(self, monkeypatch, clean_anthropic_env): + """Only the configured api_base is validated, so a redirected /v1/models must not be + allowed to replay the credential to an unvalidated origin -- same rule already applied + to the WIF token exchange itself.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + client = RecordingModelsClient([{"data": [], "has_more": False, "last_id": None}]) + monkeypatch.setattr("litellm.module_level_client", client) + + AnthropicModelInfo().get_models(api_base="https://api.anthropic.com") + + assert client.calls[0].follow_redirects is False + + def test_get_models_page_cap_stops_a_runaway_has_more(self, monkeypatch, clean_anthropic_env): + from litellm.llms.anthropic.common_utils import ( + _MODEL_LIST_PAGE_CAP, + AnthropicModelInfo, + ) + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + client = RecordingModelsClient([{"data": [{"id": "claude-loop"}], "has_more": True, "last_id": "claude-loop"}]) + monkeypatch.setattr("litellm.module_level_client", client) + + with pytest.raises(Exception, match="did not terminate"): + AnthropicModelInfo().get_models(api_base="https://api.anthropic.com") + + assert len(client.calls) == _MODEL_LIST_PAGE_CAP + + def test_get_models_error_is_sanitized_not_raw_response_text(self, monkeypatch, clean_anthropic_env): + """A failed discovery call must never echo the raw response body verbatim -- only the + structured error message, so an unrelated/oversized/reflected body is not surfaced.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + reflected_payload = "" * 50 + client = RecordingModelsClient( + responses=[ + _models_page_response( + { + "type": "error", + "error": { + "type": "authentication_error", + "message": "invalid x-api-key", + "reflected": reflected_payload, + }, + }, + status_code=401, + ) + ] + ) + monkeypatch.setattr("litellm.module_level_client", client) + + with pytest.raises(Exception, match="invalid x-api-key") as exc_info: # noqa: B017, PT011 # the callee raises a bare Exception; match pins the sanitized text + AnthropicModelInfo().get_models(api_base="https://api.anthropic.com") + + assert "invalid x-api-key" in str(exc_info.value) + assert reflected_payload not in str(exc_info.value) + + def test_discover_models_threads_litellm_params_into_wif(self, monkeypatch, wif_engine): + """The gap this phase fixes: get_models only ever saw api_key/api_base, so a WIF source + configured in litellm_params (rather than ANTHROPIC_* env vars) could not discover.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + poster, calls = wif_engine + client = RecordingModelsClient([{"data": [{"id": "claude-wif"}], "has_more": False, "last_id": None}]) + monkeypatch.setattr("litellm.module_level_client", client) + monkeypatch.setenv("DISC_JWT", "jwt-assertion-value") + + models = AnthropicModelInfo().discover_models( + litellm_params={ + "anthropic_federation_rule_id": "fdrl_disc", + "anthropic_organization_id": "org-disc", + "anthropic_identity_token": "oidc/env/DISC_JWT", + } + ) + + assert models == ["anthropic/claude-wif"] + assert len(poster.requests) == 1 + assert client.calls[0].headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + + def test_discover_models_without_litellm_params_behaves_like_get_models(self, monkeypatch, clean_anthropic_env): + """No litellm_params (the wildcard-discovery call shape) must fall back to the + env-only resolution get_models has always used -- zero behavior change for that path.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + client = RecordingModelsClient([{"data": [{"id": "claude-env"}], "has_more": False, "last_id": None}]) + monkeypatch.setattr("litellm.module_level_client", client) + + models = AnthropicModelInfo().discover_models(litellm_params=None) + + assert models == ["anthropic/claude-env"] + assert client.calls[0].headers["x-api-key"] == FAKE_REGULAR_KEY + + def test_discover_models_explicit_api_key_beats_wif(self, monkeypatch, wif_engine): + """Same precedence discover_models must honor as every other Anthropic auth surface: + WIF is the lowest tier.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + poster, calls = wif_engine + client = RecordingModelsClient([{"data": [], "has_more": False, "last_id": None}]) + monkeypatch.setattr("litellm.module_level_client", client) + + AnthropicModelInfo().discover_models( + litellm_params={ + "api_key": FAKE_REGULAR_KEY, + "anthropic_federation_rule_id": "fdrl_disc", + "anthropic_organization_id": "org-disc", + "anthropic_identity_token": "oidc/env/DISC_JWT", + } + ) + + assert client.calls[0].headers["x-api-key"] == FAKE_REGULAR_KEY + assert calls == [] + assert poster.requests == [] + + class TestWifExchangeTransportHardening: def test_token_exchange_client_does_not_follow_redirects(self): """Only the initial token URL is validated, so a 3xx must not be allowed to replay the diff --git a/tests/test_litellm/proxy/credential_endpoints/test_endpoints.py b/tests/test_litellm/proxy/credential_endpoints/test_endpoints.py index dcd8e6881bd..5810400fb77 100644 --- a/tests/test_litellm/proxy/credential_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/credential_endpoints/test_endpoints.py @@ -1,11 +1,14 @@ """Tests for the credential management endpoints.""" +import json from unittest.mock import AsyncMock, MagicMock, patch import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import ec from fastapi.testclient import TestClient - +import litellm from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.proxy_server import app @@ -35,19 +38,63 @@ def _patch_credential(name: str, body: dict): app.dependency_overrides[user_api_key_auth] = previous_override +def _post_credential(body: dict): + missing = object() + previous_override = app.dependency_overrides.get(user_api_key_auth, missing) + app.dependency_overrides[user_api_key_auth] = _as_admin + try: + return client.post("/credentials", json=body, headers={"Authorization": "Bearer test-key"}) + finally: + if previous_override is missing: + app.dependency_overrides.pop(user_api_key_auth, None) + else: + app.dependency_overrides[user_api_key_auth] = previous_override + + +def test_create_credential_write_omits_the_patch_only_deletion_field(): + """Regression: CredentialItem.credential_values_to_delete is a PATCH-only field that + defaults to None on every other construction path. A bare .model_dump() (without + exclude_none) on the create path put a `credential_values_to_delete: null` key into the + Prisma write, which litellm_credentialstable has no column for.""" + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.master_key", "sk-test-master"), + patch("litellm.proxy.credential_endpoints.endpoints.CredentialsRepository") as repository, + ): + create_mock = AsyncMock(return_value=None) + repository.return_value.create = create_mock + + response = _post_credential( + { + "credential_name": "new-cred", + "credential_values": {"api_key": "sk-new"}, + "credential_info": {"custom_llm_provider": "openai"}, + } + ) + + assert response.status_code == 200, response.text + written_data = create_mock.await_args.kwargs["data"] + assert "credential_values_to_delete" not in written_data + + def test_update_credential_answers_404_when_the_credential_does_not_exist(): """Regression: the handler used to ``return handle_exception_on_proxy(e)``, which makes the exception the response body and lets FastAPI answer 200, so a write the handler rejected read as a success to every caller that checks the status. The dashboard's API client branches on the status, so it reported a failed edit as applied.""" - with patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch( - "litellm.proxy.credential_endpoints.endpoints.CredentialsRepository" - ) as repository: + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.credential_endpoints.endpoints.CredentialsRepository") as repository, + ): repository.return_value.find_by_name = AsyncMock(return_value=None) response = _patch_credential( "definitely-not-there", - {"credential_name": "definitely-not-there", "credential_values": {"api_key": "sk-x"}, "credential_info": {}}, + { + "credential_name": "definitely-not-there", + "credential_values": {"api_key": "sk-x"}, + "credential_info": {}, + }, ) assert response.status_code == 404, f"rejected write answered {response.status_code}: {response.text}" @@ -73,9 +120,11 @@ def test_update_credential_still_answers_200_on_a_successful_write(): credential_values={"api_key": "sk-old"}, credential_info={"custom_llm_provider": "openai"}, ) - with patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch( - "litellm.proxy.proxy_server.master_key", "sk-test-master" - ), patch("litellm.proxy.credential_endpoints.endpoints.CredentialsRepository") as repository: + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.master_key", "sk-test-master"), + patch("litellm.proxy.credential_endpoints.endpoints.CredentialsRepository") as repository, + ): repository.return_value.find_by_name = AsyncMock(return_value=stored) repository.return_value.update_by_name = AsyncMock(return_value=None) @@ -86,3 +135,255 @@ def test_update_credential_still_answers_200_on_a_successful_write(): assert response.status_code == 200, response.text assert response.json()["success"] is True + + +def _get_jwks(name: str): + missing = object() + previous_override = app.dependency_overrides.get(user_api_key_auth, missing) + app.dependency_overrides[user_api_key_auth] = _as_admin + try: + return client.get(f"/credentials/{name}/jwks", headers={"Authorization": "Bearer test-key"}) + finally: + if previous_override is missing: + app.dependency_overrides.pop(user_api_key_auth, None) + else: + app.dependency_overrides[user_api_key_auth] = previous_override + + +@pytest.fixture +def restore_credential_list(monkeypatch): + monkeypatch.setattr(litellm, "credential_list", []) + + +def test_update_credential_rejects_overlap_between_update_and_delete(): + """A key in both sets is ambiguous (set to what value, before or after the delete?), so the + endpoint must reject it outright rather than picking a resolution order silently.""" + response = _patch_credential( + "any-name", + { + "credential_name": "any-name", + "credential_values": {"api_key": "sk-new"}, + "credential_values_to_delete": ["api_key"], + "credential_info": {}, + }, + ) + + assert response.status_code == 400, response.text + assert "api_key" in response.json()["error"]["message"] + + +def test_update_credential_deletion_removes_the_key_from_the_db_write(restore_credential_list): + """The bug this closes: switching WIF identity sources (or WIF -> api_key) left the old + variant's fields behind in the DB row, which wif.py then rejects by presence.""" + stored = CredentialItem( + credential_name="wif-cred", + credential_values={"anthropic_identity_source": "keycloak", "anthropic_keycloak_client_id": "old-client"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.credential_endpoints.endpoints.CredentialsRepository") as repository, + ): + repository.return_value.find_by_name = AsyncMock(return_value=stored) + update_mock = AsyncMock(return_value=None) + repository.return_value.update_by_name = update_mock + + response = _patch_credential( + "wif-cred", + { + "credential_name": "wif-cred", + "credential_values": {}, + "credential_values_to_delete": ["anthropic_keycloak_client_id"], + "credential_info": {}, + }, + ) + + assert response.status_code == 200, response.text + written_values = json.loads(update_mock.await_args.kwargs["data"]["credential_values"]) + assert "anthropic_keycloak_client_id" not in written_values + assert written_values["anthropic_identity_source"] == "keycloak" + + +def test_update_credential_deletion_updates_in_memory_credential_list(restore_credential_list, monkeypatch): + """The in-memory list is what the request-time auth resolvers read; a deletion that only + landed in the DB would leave the stale field servable until the next process restart.""" + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="wif-cred", + credential_values={ + "anthropic_identity_source": "keycloak", + "anthropic_keycloak_client_id": "old-client", + }, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ], + ) + stored = CredentialItem( + credential_name="wif-cred", + credential_values={"anthropic_identity_source": "keycloak", "anthropic_keycloak_client_id": "old-client"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.credential_endpoints.endpoints.CredentialsRepository") as repository, + ): + repository.return_value.find_by_name = AsyncMock(return_value=stored) + repository.return_value.update_by_name = AsyncMock(return_value=None) + + response = _patch_credential( + "wif-cred", + { + "credential_name": "wif-cred", + "credential_values": {}, + "credential_values_to_delete": ["anthropic_keycloak_client_id"], + "credential_info": {}, + }, + ) + + assert response.status_code == 200, response.text + in_memory = next(c for c in litellm.credential_list if c.credential_name == "wif-cred") + assert "anthropic_keycloak_client_id" not in in_memory.credential_values + assert in_memory.credential_values["anthropic_identity_source"] == "keycloak" + + +def test_update_credential_leaves_untouched_fields_alone(): + """Regression for the masked-value hazard: GET /credentials masks values, so a PATCH that + only names the field being changed must not let an untouched field be nulled or overwritten + by anything a round-tripped (masked) form value could contain.""" + stored = CredentialItem( + credential_name="existing", + credential_values={"api_key": "sk-real-value", "api_base": "https://api.anthropic.com"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.master_key", "sk-test-master"), + patch("litellm.proxy.credential_endpoints.endpoints.CredentialsRepository") as repository, + ): + repository.return_value.find_by_name = AsyncMock(return_value=stored) + update_mock = AsyncMock(return_value=None) + repository.return_value.update_by_name = update_mock + + response = _patch_credential( + "existing", + {"credential_name": "existing", "credential_values": {"api_key": "sk-rotated"}, "credential_info": {}}, + ) + + assert response.status_code == 200, response.text + written_values = json.loads(update_mock.await_args.kwargs["data"]["credential_values"]) + assert written_values["api_base"] == "https://api.anthropic.com" + + +def _generate_es256_pem() -> str: + key = ec.generate_private_key(ec.SECP256R1()) + return key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ).decode() + + +class TestCredentialJwksExport: + def test_jwks_export_succeeds_for_an_internal_issuer_credential(self, restore_credential_list, monkeypatch): + monkeypatch.setenv("JWKS_TEST_SIGNING_KEY", _generate_es256_pem()) + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="anthropic-issuer", + credential_values={ + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.example.com", + "anthropic_issuer_subject": "my-workload", + "anthropic_issuer_signing_key_ref": "os.environ/JWKS_TEST_SIGNING_KEY", + }, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ], + ) + + response = _get_jwks("anthropic-issuer") + + assert response.status_code == 200, response.text + body = response.json() + assert body["keys"][0]["kty"] == "EC" + assert body["keys"][0]["crv"] == "P-256" + # The private key material must never leave the process via this endpoint. + assert "JWKS_TEST_SIGNING_KEY" not in response.text + assert "PRIVATE KEY" not in response.text + + def test_jwks_export_404s_for_a_non_anthropic_credential(self, restore_credential_list, monkeypatch): + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="openai-key", + credential_values={"api_key": "sk-x"}, + credential_info={"custom_llm_provider": "openai"}, + ) + ], + ) + + response = _get_jwks("openai-key") + + assert response.status_code == 404, response.text + + def test_jwks_export_404s_for_an_anthropic_credential_without_internal_issuer( + self, restore_credential_list, monkeypatch + ): + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="anthropic-apikey", + credential_values={"api_key": "sk-ant"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ], + ) + + response = _get_jwks("anthropic-apikey") + + assert response.status_code == 404, response.text + + def test_jwks_export_404s_for_an_unknown_credential(self, restore_credential_list): + with patch("litellm.proxy.proxy_server.prisma_client", None): + response = _get_jwks("does-not-exist") + + assert response.status_code == 404, response.text + + def test_jwks_export_requires_proxy_admin(self, restore_credential_list, monkeypatch): + monkeypatch.setenv("JWKS_TEST_SIGNING_KEY", _generate_es256_pem()) + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="anthropic-issuer", + credential_values={ + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.example.com", + "anthropic_issuer_subject": "my-workload", + "anthropic_issuer_signing_key_ref": "os.environ/JWKS_TEST_SIGNING_KEY", + }, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ], + ) + + def _as_internal_user(): + return UserAPIKeyAuth(api_key="test-key", user_role="internal_user") + + app.dependency_overrides[user_api_key_auth] = _as_internal_user + try: + response = client.get("/credentials/anthropic-issuer/jwks", headers={"Authorization": "Bearer test-key"}) + finally: + app.dependency_overrides.pop(user_api_key_auth, None) + + assert response.status_code == 403, response.text 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 097230108d4..f5201e85c28 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 @@ -4,6 +4,7 @@ from typing import Dict, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest +from fastapi import HTTPException from fastapi.testclient import TestClient from litellm._uuid import uuid @@ -44,11 +45,7 @@ class MockPrismaClient: return LiteLLM_TeamTable( team_id=where["team_id"], team_alias="test_team", - members_with_roles=[ - Member( - user_id="test_user", role="admin" if self.user_admin else "user" - ) - ], + members_with_roles=[Member(user_id="test_user", role="admin" if self.user_admin else "user")], ) return None @@ -62,10 +59,7 @@ class MockPrismaClient: # Support model_name startswith filter (used by _get_team_deployments) if where and "model_name" in where: model_name_filter = where["model_name"] - if ( - isinstance(model_name_filter, dict) - and "startswith" in model_name_filter - ): + if isinstance(model_name_filter, dict) and "startswith" in model_name_filter: prefix = model_name_filter["startswith"] results = [d for d in results if d.model_name.startswith(prefix)] @@ -110,13 +104,9 @@ class MockProxyConfig: class TestModelManagementAuthChecks: def setup_method(self): """Setup test cases""" - self.admin_user = UserAPIKeyAuth( - user_id="test_admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + self.admin_user = UserAPIKeyAuth(user_id="test_admin", user_role=LitellmUserRoles.PROXY_ADMIN) - self.normal_user = UserAPIKeyAuth( - user_id="test_user", user_role=LitellmUserRoles.INTERNAL_USER - ) + self.normal_user = UserAPIKeyAuth(user_id="test_user", user_role=LitellmUserRoles.INTERNAL_USER) self.team_admin_user = UserAPIKeyAuth( user_id="test_user", @@ -135,7 +125,7 @@ class TestModelManagementAuthChecks: @pytest.mark.asyncio async def test_can_user_make_team_model_call_non_premium_fails(self): """Test that non-premium users cannot make team model calls""" - with pytest.raises(Exception, match='You must be a LiteLLM Enterprise user to use this feature\\.') as exc_info: + with pytest.raises(Exception, match="You must be a LiteLLM Enterprise user to use this feature\\.") as exc_info: ModelManagementAuthChecks.can_user_make_team_model_call( team_id="test_team", user_api_key_dict=self.admin_user, @@ -149,9 +139,7 @@ class TestModelManagementAuthChecks: team_obj = LiteLLM_TeamTable( team_id="test_team", team_alias="test_team", - members_with_roles=[ - Member(user_id=self.team_admin_user.user_id, role="admin") - ], + members_with_roles=[Member(user_id=self.team_admin_user.user_id, role="admin")], ) result = ModelManagementAuthChecks.can_user_make_team_model_call( @@ -190,7 +178,7 @@ class TestModelManagementAuthChecks: ) prisma_client = MockPrismaClient(team_exists=True) - with pytest.raises(Exception, match='You must be a LiteLLM Enterprise user to use this feature\\.') as exc_info: + with pytest.raises(Exception, match="You must be a LiteLLM Enterprise user to use this feature\\.") as exc_info: await ModelManagementAuthChecks.allow_team_model_action( model_params=model_params, user_api_key_dict=self.admin_user, @@ -325,29 +313,21 @@ class TestDeleteTeamModelAlias: mock_prisma.db = MockPrismaWrapper(model_aliases_list) # Call the function - await delete_team_model_alias( - public_model_name="public_model_1", prisma_client=mock_prisma - ) + await delete_team_model_alias(public_model_name="public_model_1", prisma_client=mock_prisma) # Verify results mock_db = mock_prisma.db.litellm_modeltable - assert ( - len(mock_db.update_calls) == 2 - ) # Should have 2 update calls since public_model_1 appears twice + assert len(mock_db.update_calls) == 2 # Should have 2 update calls since public_model_1 appears twice # Verify first update first_update = mock_db.update_calls[0] assert first_update["where"] == {"id": 1} - assert json.loads(first_update["data"]["model_aliases"]) == { - "alias2": "public_model_2" - } + assert json.loads(first_update["data"]["model_aliases"]) == {"alias2": "public_model_2"} # Verify second update second_update = mock_db.update_calls[1] assert second_update["where"] == {"id": 2} - assert json.loads(second_update["data"]["model_aliases"]) == { - "alias3": "public_model_3" - } + assert json.loads(second_update["data"]["model_aliases"]) == {"alias3": "public_model_3"} @pytest.mark.asyncio async def test_delete_team_model_alias_no_matches(self): @@ -383,9 +363,7 @@ class TestDeleteTeamModelAlias: mock_prisma.db = MockPrismaWrapper(model_aliases_list) # Call the function with non-existent model - await delete_team_model_alias( - public_model_name="non_existent_model", prisma_client=mock_prisma - ) + await delete_team_model_alias(public_model_name="non_existent_model", prisma_client=mock_prisma) # Verify no updates were made mock_db = mock_prisma.db.litellm_modeltable @@ -882,18 +860,12 @@ class TestUpdateModel: updated_row.model_dump_json.return_value = "{}" mock_prisma = MagicMock() - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=existing_row - ) - mock_prisma.db.litellm_proxymodeltable.update = AsyncMock( - return_value=updated_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row) mock_router = MagicMock() mock_router.get_model_ids.return_value = [model_id] - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -910,9 +882,7 @@ class TestUpdateModel: ), patch( "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", - new=AsyncMock( - return_value=ReconcileOutcome(still_desired=None, live_after=None) - ), + new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)), ) as mock_clear_cache, ): await update_model( @@ -957,9 +927,7 @@ class TestUpdatePublicModelGroups: mock_proxy_config.get_config = mock_get_config mock_proxy_config.save_config = AsyncMock() - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) request = UpdatePublicModelGroupsRequest(model_groups=new_models) @@ -1015,9 +983,7 @@ class TestUpdatePublicModelGroups: mock_proxy_config.get_config = mock_get_config mock_proxy_config.save_config = AsyncMock() - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) request = UpdateUsefulLinksRequest(useful_links=new_links) @@ -1182,9 +1148,7 @@ class TestTeamModelSiblingRouting: ) # Global deployment should be accessible when team_id is provided - deployments = router._get_all_deployments( - model_name="global-gpt-4o", team_id="teamA" - ) + deployments = router._get_all_deployments(model_name="global-gpt-4o", team_id="teamA") assert len(deployments) == 1 assert deployments[0]["model_name"] == "global-gpt-4o" @@ -1233,9 +1197,7 @@ class TestTeamModelUpdate: patch( "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" ) as mock_team_model_add, - patch( - "litellm.proxy.management_endpoints.model_management_endpoints.update_team" - ) as mock_update_team, + patch("litellm.proxy.management_endpoints.model_management_endpoints.update_team") as mock_update_team, ): result = await _update_team_model_in_db( db_model=db_model, @@ -1265,9 +1227,7 @@ class TestTeamModelUpdate: db_model = Deployment( model_name="model_name_team_123_uuid1", litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"), - model_info=ModelInfo( - team_id="team_123", team_public_model_name="old-public-name" - ), + model_info=ModelInfo(team_id="team_123", team_public_model_name="old-public-name"), ) # Create a sibling deployment that still uses the old public name @@ -1278,9 +1238,7 @@ class TestTeamModelUpdate: "team_public_model_name": "old-public-name", } - prisma_client = MockPrismaClient( - team_exists=True, sibling_deployments=[sibling_deployment] - ) + prisma_client = MockPrismaClient(team_exists=True, sibling_deployments=[sibling_deployment]) patch_data = updateDeployment( model_name="new-public-name", @@ -1293,12 +1251,8 @@ class TestTeamModelUpdate: ) with ( - patch( - "litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete" - ) as mock_delete, - patch( - "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" - ) as mock_add, + patch("litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete") as mock_delete, + patch("litellm.proxy.management_endpoints.model_management_endpoints.team_model_add") as mock_add, ): await _update_existing_team_model_assignment( team_id="team_123", @@ -1339,12 +1293,8 @@ class TestTeamModelUpdate: ) with ( - patch( - "litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete" - ) as mock_delete, - patch( - "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" - ) as mock_add, + patch("litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete") as mock_delete, + patch("litellm.proxy.management_endpoints.model_management_endpoints.team_model_add") as mock_add, ): await _update_existing_team_model_assignment( team_id="team_123", @@ -1369,9 +1319,7 @@ class TestTeamModelUpdate: db_model = Deployment( model_name="model_name_team_123_uuid1", litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"), - model_info=ModelInfo( - team_id="team_123", team_public_model_name="old-public-name" - ), + model_info=ModelInfo(team_id="team_123", team_public_model_name="old-public-name"), ) patch_data = updateDeployment( model_name="new-public-name", @@ -1406,20 +1354,14 @@ class TestTeamModelUpdate: db_model = Deployment( model_name="model_name_team_123_uuid1", litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"), - model_info=ModelInfo( - team_id="team_123", team_public_model_name="old-public-name" - ), + model_info=ModelInfo(team_id="team_123", team_public_model_name="old-public-name"), ) sibling_deployment = MagicMock() sibling_deployment.model_name = "model_name_team_123_uuid2" - sibling_deployment.model_info = ( - '{"team_id":"team_123","team_public_model_name":"old-public-name"}' - ) + sibling_deployment.model_info = '{"team_id":"team_123","team_public_model_name":"old-public-name"}' - prisma_client = MockPrismaClient( - team_exists=True, sibling_deployments=[sibling_deployment] - ) + prisma_client = MockPrismaClient(team_exists=True, sibling_deployments=[sibling_deployment]) patch_data = updateDeployment( model_name="new-public-name", @@ -1432,12 +1374,8 @@ class TestTeamModelUpdate: ) with ( - patch( - "litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete" - ) as mock_delete, - patch( - "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" - ) as mock_add, + patch("litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete") as mock_delete, + patch("litellm.proxy.management_endpoints.model_management_endpoints.team_model_add") as mock_add, ): await _update_existing_team_model_assignment( team_id="team_123", @@ -1514,10 +1452,7 @@ class TestTeamModelUpdate: ), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "gpt-5.2-low-rpm-testing" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "gpt-5.2-low-rpm-testing" def test_get_public_model_name_preserves_db_public_name_when_internal_name_unchanged( self, @@ -1544,10 +1479,7 @@ class TestTeamModelUpdate: model_info=ModelInfo(team_id="test-team"), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "gpt-5.2-low-rpm-testing" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "gpt-5.2-low-rpm-testing" def test_get_public_model_name_allows_top_level_rename(self): """A genuine rename via the top-level model_name field (no @@ -1572,10 +1504,7 @@ class TestTeamModelUpdate: model_info=ModelInfo(team_id="test-team"), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "new-public-name" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "new-public-name" def test_get_public_model_name_top_level_rename_wins_over_stale_model_info(self): """Regression (codex review): on a dashboard rename the UI sends the new @@ -1592,9 +1521,7 @@ class TestTeamModelUpdate: db_model = Deployment( model_name="model_name_team-a_abc123", litellm_params=LiteLLM_Params(model="azure/gpt-4.1"), - model_info=ModelInfo( - team_id="team-a", team_public_model_name="old-public-name" - ), + model_info=ModelInfo(team_id="team-a", team_public_model_name="old-public-name"), ) patch_data = updateDeployment( model_name="new-public-name", @@ -1604,10 +1531,7 @@ class TestTeamModelUpdate: ), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "new-public-name" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "new-public-name" def test_get_public_model_name_falls_back_to_db_public_name(self): """When patch_data carries no name hints at all (neither model_name @@ -1630,10 +1554,7 @@ class TestTeamModelUpdate: model_info=ModelInfo(team_id="test-team"), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "gpt-5.2-low-rpm-testing" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "gpt-5.2-low-rpm-testing" def test_get_public_model_name_last_resort_returns_db_model_name(self): """Legacy rows may have no team_public_model_name anywhere; the @@ -1653,10 +1574,7 @@ class TestTeamModelUpdate: model_info=ModelInfo(team_id="test-team"), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "legacy-model" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "legacy-model" def test_get_public_model_name_ignores_different_internal_shape_name(self): """A stale client may PATCH an internal-shaped model_name that does not @@ -1680,10 +1598,7 @@ class TestTeamModelUpdate: model_info=ModelInfo(team_id="test-team"), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "gpt-5.2-low-rpm-testing" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "gpt-5.2-low-rpm-testing" def test_get_public_model_name_ignores_internal_shape_patch_public(self): """If a corrupted row round-trips an internal-shaped value in @@ -1709,10 +1624,7 @@ class TestTeamModelUpdate: ), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "gpt-5.2-low-rpm-testing" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "gpt-5.2-low-rpm-testing" @pytest.mark.asyncio async def test_dashboard_edit_preserves_public_name_and_acl(self): @@ -1779,9 +1691,7 @@ class TestTeamModelUpdate: # the merged model_info written to the DB must keep the public name model_info_json = result.get("model_info", "") parsed_model_info = json.loads(model_info_json) - assert ( - parsed_model_info.get("team_public_model_name") == "gpt-5.2-low-rpm-testing" - ) + assert parsed_model_info.get("team_public_model_name") == "gpt-5.2-low-rpm-testing" # the internal model_name must not have been overwritten (caller # intentionally clears patch_data.model_name so the DB row's name @@ -1823,9 +1733,7 @@ class TestModelInfoEndpoint: model_info=ModelInfo(id="gpt-4"), ) - result = await model_info( - model_id="gpt-4", user_api_key_dict=user_api_key_dict - ) + result = await model_info(model_id="gpt-4", user_api_key_dict=user_api_key_dict) assert result["id"] == "gpt-4" assert result["object"] == "model" @@ -1900,9 +1808,7 @@ class TestModelInfoEndpoint: model_info=ModelInfo(id="team-model-1"), ) - result = await model_info( - model_id="team-model-1", user_api_key_dict=user_api_key_dict - ) + result = await model_info(model_id="team-model-1", user_api_key_dict=user_api_key_dict) assert result["id"] == "team-model-1" assert result["object"] == "model" @@ -1934,9 +1840,7 @@ class TestAddAndDeleteModelLifecycle: ) model_id = "lifecycle-test-model-123" - admin_user = UserAPIKeyAuth( - user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN) # Build a real LiteLLM_ProxyModelTable for the DB mock to return db_row = LiteLLM_ProxyModelTable( @@ -1952,9 +1856,7 @@ class TestAddAndDeleteModelLifecycle: mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=db_row) - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_proxy_config = MagicMock() @@ -1978,14 +1880,11 @@ class TestAddAndDeleteModelLifecycle: patch(f"{_PS}.llm_router", mock_router), patch(_ENCRYPT, side_effect=lambda value, **kwargs: value), ): - # --- ADD --- add_result = await add_new_model( model_params=Deployment( model_name="lifecycle-model", - litellm_params=LiteLLM_Params( - model="openai/gpt-4.1-nano", api_key="fake-key" - ), + litellm_params=LiteLLM_Params(model="openai/gpt-4.1-nano", api_key="fake-key"), model_info={"id": model_id}, ), user_api_key_dict=admin_user, @@ -2000,9 +1899,7 @@ class TestAddAndDeleteModelLifecycle: assert "deleted successfully" in delete_result["message"] # --- DELETE again should fail (model not found) --- - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=None) from litellm.proxy.proxy_server import ProxyException with pytest.raises(ProxyException) as exc_info: @@ -2063,24 +1960,18 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma = MagicMock() mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) # After the row delete no team deployment remains -> nothing backs the public name. mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_teamtable = AsyncMock() mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=updated_team_row - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=updated_team_row) # Team BYOK models have no alias row; delete_team_model_alias finds nothing. mock_prisma.db.litellm_modeltable = AsyncMock() mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[]) - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -2145,9 +2036,7 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma = MagicMock() mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_teamtable = AsyncMock() @@ -2157,9 +2046,7 @@ class TestDeleteTeamBYOKModelGhost: # No alias row matches -> delete_team_model_alias returns nothing, but it still ran. mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[]) - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -2221,25 +2108,17 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma = MagicMock() mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=deleted_row - ) - mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock( - return_value=deleted_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=deleted_row) + mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=deleted_row) # After the deleted replica's row is gone, the sibling still backs the public name. - mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock( - return_value=[sibling_row] - ) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[sibling_row]) mock_prisma.db.litellm_teamtable = AsyncMock() mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=team_row) mock_prisma.db.litellm_modeltable = AsyncMock() mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[]) - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -2297,35 +2176,27 @@ class TestDeleteTeamBYOKModelGhost: members_with_roles=[Member(user_id="admin", role="admin")], models=[public_name], ) - alias_row = MagicMock( - id="alias-row-1", model_aliases={public_name: internal_name} - ) + alias_row = MagicMock(id="alias-row-1", model_aliases={public_name: internal_name}) alias_row.team = MagicMock() alias_row.team.team_id = team_id mock_prisma = MagicMock() mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_teamtable = AsyncMock() mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=team_row) mock_prisma.db.litellm_modeltable = AsyncMock() - mock_prisma.db.litellm_modeltable.find_many = AsyncMock( - return_value=[alias_row] - ) + mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[alias_row]) mock_prisma.db.litellm_modeltable.update = AsyncMock() mock_router = MagicMock() mock_router.model_name_to_deployment_indices = {public_name: [0]} - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -2387,9 +2258,7 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma = MagicMock() mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_teamtable = AsyncMock() @@ -2402,9 +2271,7 @@ class TestDeleteTeamBYOKModelGhost: mock_router = MagicMock() mock_router.model_name_to_deployment_indices = {internal_name: [0]} - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -2456,9 +2323,7 @@ class TestDeleteModelTeamAuth: mock_prisma = MagicMock() mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) # The team is gone -> every team lookup returns None. @@ -2480,9 +2345,7 @@ class TestDeleteModelTeamAuth: model_id = "orphaned-byok-1" mock_prisma = self._orphaned_model_mocks(team_id, model_id) - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -2518,9 +2381,7 @@ class TestDeleteModelTeamAuth: model_id = "orphaned-byok-2" mock_prisma = self._orphaned_model_mocks(team_id, model_id) - non_admin = UserAPIKeyAuth( - user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER - ) + non_admin = UserAPIKeyAuth(user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -2574,9 +2435,7 @@ class TestDeleteModelTeamAuth: mock_prisma = MagicMock() mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_teamtable = AsyncMock() @@ -2586,9 +2445,7 @@ class TestDeleteModelTeamAuth: # A team member who is not the team admin: rejected before the delete runs, # so the only team lookup is the single one inside the auth check. - non_admin = UserAPIKeyAuth( - user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER - ) + non_admin = UserAPIKeyAuth(user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -2782,15 +2639,11 @@ class TestDeleteTeamModels: prisma = _TxPrismaClient(rows) router = _RecordingRouter(prisma.events) - await delete_team_models( - team_ids=["team_a", "team_b"], prisma_client=prisma, llm_router=router - ) + await delete_team_models(team_ids=["team_a", "team_b"], prisma_client=prisma, llm_router=router) commit_idx = prisma.events.index(("commit",)) router_indices = [i for i, e in enumerate(prisma.events) if e[0] == "router"] - delete_indices = [ - i for i, e in enumerate(prisma.events) if e[0] == "delete_many" - ] + delete_indices = [i for i, e in enumerate(prisma.events) if e[0] == "delete_many"] assert router_indices, "router was never synced" assert all(i > commit_idx for i in router_indices) assert all(i < commit_idx for i in delete_indices) @@ -2806,9 +2659,7 @@ class TestDeleteTeamModels: prisma = _TxPrismaClient([mine, intruder]) router = _RecordingRouter(prisma.events) - deleted = await delete_team_models( - team_ids=["team_a"], prisma_client=prisma, llm_router=router - ) + deleted = await delete_team_models(team_ids=["team_a"], prisma_client=prisma, llm_router=router) assert deleted == ["a1"] assert router.deleted == ["a1"] @@ -2818,9 +2669,7 @@ class TestDeleteTeamModels: prisma = _TxPrismaClient([]) router = _RecordingRouter(prisma.events) - deleted = await delete_team_models( - team_ids=["team_a"], prisma_client=prisma, llm_router=router - ) + deleted = await delete_team_models(team_ids=["team_a"], prisma_client=prisma, llm_router=router) assert deleted == [] assert router.deleted == [] @@ -2831,9 +2680,7 @@ class TestDeleteTeamModels: rows = [_model_row("a1", "team_a")] prisma = _TxPrismaClient(rows) - deleted = await delete_team_models( - team_ids=["team_a"], prisma_client=prisma, llm_router=None - ) + deleted = await delete_team_models(team_ids=["team_a"], prisma_client=prisma, llm_router=None) assert deleted == ["a1"] assert any(e[0] == "delete_many" for e in prisma.events) @@ -2921,9 +2768,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=_build_db_model_with_pricing(), - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(input_cost_per_token=None) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(input_cost_per_token=None)), ) params = json.loads(result["litellm_params"]) @@ -2942,9 +2787,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=_build_db_model_with_pricing(), - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(output_cost_per_token=None) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(output_cost_per_token=None)), ) params = json.loads(result["litellm_params"]) @@ -2960,9 +2803,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=_build_db_model_with_pricing(), - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(input_cost_per_token=0.000005) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(input_cost_per_token=0.000005)), ) params = json.loads(result["litellm_params"]) @@ -2977,9 +2818,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=_build_db_model_with_pricing(), - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(output_cost_per_token=0.000007) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(output_cost_per_token=0.000007)), ) params = json.loads(result["litellm_params"]) @@ -3014,9 +2853,7 @@ class TestUpdateDBModelClearPricing: # or any other non-pricing field from the merged dict. result = update_db_model( db_model=db_model, - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(api_base=None) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(api_base=None)), ) info = json.loads(result["model_info"]) @@ -3051,9 +2888,7 @@ class TestUpdateDBModelClearPricing: params = json.loads(result["litellm_params"]) info = json.loads(result["model_info"]) assert "input_cost_per_token" not in params - assert ( - "input_cost_per_token" not in info - ), "model_info passthrough must not resurrect the cleared override" + assert "input_cost_per_token" not in info, "model_info passthrough must not resurrect the cleared override" def test_clear_via_model_info_clears_both_blobs(self): """The mirror works in the reverse direction too: nulling a pricing field @@ -3065,9 +2900,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=_build_db_model_with_pricing(), - updated_patch=updateDeployment( - model_info=ModelInfo(id="dep-pricing-0", input_cost_per_token=None) - ), + updated_patch=updateDeployment(model_info=ModelInfo(id="dep-pricing-0", input_cost_per_token=None)), ) params = json.loads(result["litellm_params"]) @@ -3099,9 +2932,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=db_model, - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(cache_read_input_token_cost=None) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(cache_read_input_token_cost=None)), ) params = json.loads(result["litellm_params"]) @@ -3133,9 +2964,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=db_model, - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(cache_creation_input_token_cost=None) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(cache_creation_input_token_cost=None)), ) params = json.loads(result["litellm_params"]) @@ -3169,9 +2998,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=db_model, - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(cache_read_input_token_cost=None) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(cache_read_input_token_cost=None)), ) params = json.loads(result["litellm_params"]) @@ -3237,9 +3064,7 @@ class TestPatchModelBlockedAuthGate: existing_row.model_dump_json.return_value = "{}" mock_prisma = MagicMock() - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=existing_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -3280,12 +3105,8 @@ class TestPatchModelBlockedAuthGate: updated_row.model_dump_json.return_value = "{}" mock_prisma = MagicMock() - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=existing_row - ) - mock_prisma.db.litellm_proxymodeltable.update = AsyncMock( - return_value=updated_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -3298,9 +3119,7 @@ class TestPatchModelBlockedAuthGate: ), patch( "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", - new=AsyncMock( - return_value=ReconcileOutcome(still_desired=None, live_after=None) - ), + new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)), ), ): result = await patch_model( @@ -3382,9 +3201,7 @@ class TestWriteSurfacesReloadDrop: ) with pytest.raises(ProxyException, match="m-gone"): - raise_if_reload_degraded_serving( - before=frozenset(), written_models=[("m-gone", None)], action="update" - ) + raise_if_reload_degraded_serving(before=frozenset(), written_models=[("m-gone", None)], action="update") with pytest.raises(ProxyException, match="m-collateral"): raise_if_reload_degraded_serving( @@ -3499,10 +3316,7 @@ class TestConcurrentModelWritesDoNotEvictEachOther: config = ProxyConfig() await asyncio.gather( - *[ - config.add_deployment(prisma_client=MagicMock(), proxy_logging_obj=MagicMock()) - for _ in range(5) - ] + *[config.add_deployment(prisma_client=MagicMock(), proxy_logging_obj=MagicMock()) for _ in range(5)] ) assert observed_max == 1 @@ -3725,9 +3539,7 @@ class TestDeleteEvictionsHoldTheReconcileLock: ) async def call() -> None: - await delete_team_models( - team_ids=["team-1"], prisma_client=prisma, llm_router=router - ) + await delete_team_models(team_ids=["team-1"], prisma_client=prisma, llm_router=router) await self._assert_evicts_under_lock(monkeypatch, call, model_id) router.delete_deployment.assert_called_once_with(id=model_id) @@ -4122,13 +3934,17 @@ class TestAutoRouterClassifierDefaultPrompt: from litellm.router_strategy.complexity_router import ClassificationRubric, classification_system_prompt for preset in ClassificationRubric: - response = await get_auto_router_classifier_default_prompt(context_window_size=5, classification_rubric=preset) + response = await get_auto_router_classifier_default_prompt( + context_window_size=5, classification_rubric=preset + ) assert response.system_prompt == classification_system_prompt(5, classification_rubric=preset) agentic = await get_auto_router_classifier_default_prompt( context_window_size=5, classification_rubric=ClassificationRubric.AGENTIC ) - chat = await get_auto_router_classifier_default_prompt(context_window_size=5, classification_rubric=ClassificationRubric.CHAT) + chat = await get_auto_router_classifier_default_prompt( + context_window_size=5, classification_rubric=ClassificationRubric.CHAT + ) unset = await get_auto_router_classifier_default_prompt(context_window_size=5) assert "Calibration on engineering tasks" in agentic.system_prompt assert "Calibration on engineering tasks" not in chat.system_prompt @@ -4201,3 +4017,405 @@ class TestAutoRouterClassifierDefaultPrompt: for empty in (None, "", "{}"): response = await get_auto_router_classifier_default_prompt(context_window_size=5, tier_labels=empty) assert response.system_prompt == classification_system_prompt(5) + + +class TestAddModelToDbBlocked: + """`_add_model_to_db` must thread `blocked` into the initial insert, so the wizard can + create a discovered-but-unchecked model already paused instead of active-then-patched.""" + + @staticmethod + def _deployment(blocked): + from litellm.types.router import ModelInfo + + return Deployment( + model_name="anthropic/claude-discovered", + litellm_params=LiteLLM_Params(model="anthropic/claude-discovered"), + model_info=ModelInfo(id="dep-blocked-create-0"), + blocked=blocked, + ) + + @pytest.mark.asyncio + async def test_add_model_to_db_writes_blocked_true(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _add_model_to_db, + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=MagicMock()) + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with patch("litellm.proxy.proxy_server.master_key", "sk-test-master"): + await _add_model_to_db( + model_params=self._deployment(True), user_api_key_dict=admin, prisma_client=mock_prisma + ) + + _, kwargs = mock_prisma.db.litellm_proxymodeltable.create.call_args + assert kwargs["data"]["blocked"] is True + + @pytest.mark.asyncio + async def test_add_model_to_db_writes_blocked_false(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _add_model_to_db, + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=MagicMock()) + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with patch("litellm.proxy.proxy_server.master_key", "sk-test-master"): + await _add_model_to_db( + model_params=self._deployment(False), user_api_key_dict=admin, prisma_client=mock_prisma + ) + + _, kwargs = mock_prisma.db.litellm_proxymodeltable.create.call_args + assert kwargs["data"]["blocked"] is False + + @pytest.mark.asyncio + async def test_add_model_to_db_omits_blocked_when_not_set(self): + """None means "don't set it" -- the Prisma column defaults to False -- not "explicitly + unblocked", so the key must be absent from the write entirely.""" + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _add_model_to_db, + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=MagicMock()) + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with patch("litellm.proxy.proxy_server.master_key", "sk-test-master"): + await _add_model_to_db( + model_params=self._deployment(None), user_api_key_dict=admin, prisma_client=mock_prisma + ) + + _, kwargs = mock_prisma.db.litellm_proxymodeltable.create.call_args + assert "blocked" not in kwargs["data"] + + +class TestAddNewModelBlockedAuthGate: + """Same proxy-admin-only rule patch_model applies to `blocked` must hold at create time + too: a team admin authorized for a team-scoped model must not be able to create it already + paused (or explicitly unpaused) out from under the proxy admin.""" + + @pytest.mark.asyncio + async def test_non_admin_cannot_set_blocked_on_create(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + add_new_model, + ) + + non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER) + mock_prisma = MagicMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + ): + with pytest.raises(ProxyException) as exc_info: + await add_new_model( + model_params=Deployment( + model_name="my-model", + litellm_params=LiteLLM_Params(model="openai/gpt-4o-mini"), + model_info={"id": "blocked-gate-create-0"}, + blocked=True, + ), + user_api_key_dict=non_admin, + ) + assert "proxy admin" in str(exc_info.value.message).lower() + mock_prisma.db.litellm_proxymodeltable.create.assert_not_called() + + @pytest.mark.asyncio + async def test_proxy_admin_can_create_a_blocked_model(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + add_new_model, + ) + + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + mock_prisma = MagicMock() + created_row = MagicMock() + created_row.model_id = "blocked-gate-create-1" + created_row.model_dump_json.return_value = "{}" + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=created_row) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.master_key", "sk-test-master"), + patch( + "litellm.proxy.proxy_server.llm_router", + MagicMock(**{"get_model_ids.return_value": ["blocked-gate-create-1"]}), + ), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + patch( + "litellm.proxy.proxy_server.proxy_config", + MagicMock(add_deployment=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None))), + ), + ): + result = await add_new_model( + model_params=Deployment( + model_name="my-model", + litellm_params=LiteLLM_Params(model="openai/gpt-4o-mini"), + model_info={"id": "blocked-gate-create-1"}, + blocked=True, + ), + user_api_key_dict=admin, + ) + assert result is created_row + _, kwargs = mock_prisma.db.litellm_proxymodeltable.create.call_args + assert kwargs["data"]["blocked"] is True + + +class TestDiscoverProviderModels: + """POST /provider/models/discover: proxy-admin-only, credential-name-only contract for + server-owned auth (WIF), never a silent [] on failure.""" + + @staticmethod + def _admin(): + return UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + @pytest.mark.asyncio + async def test_non_admin_is_rejected(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + discover_provider_models, + ) + from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + ProviderModelDiscoveryRequest, + ) + + non_admin = UserAPIKeyAuth(user_id="not-admin", user_role=LitellmUserRoles.INTERNAL_USER) + + with pytest.raises(ProxyException) as exc_info: + await discover_provider_models( + data=ProviderModelDiscoveryRequest(custom_llm_provider="anthropic"), + user_api_key_dict=non_admin, + ) + assert "proxy admin" in str(exc_info.value.message).lower() + + @pytest.mark.asyncio + async def test_inline_api_base_with_credential_name_is_rejected(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + discover_provider_models, + ) + from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + ProviderModelDiscoveryRequest, + ) + + with pytest.raises(HTTPException) as exc_info: + await discover_provider_models( + data=ProviderModelDiscoveryRequest( + custom_llm_provider="anthropic", + litellm_credential_name="anthropic-wif", + api_base="https://attacker.example.com", + ), + user_api_key_dict=self._admin(), + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_inline_os_environ_reference_is_rejected(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + discover_provider_models, + ) + from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + ProviderModelDiscoveryRequest, + ) + + with pytest.raises(HTTPException) as exc_info: + await discover_provider_models( + data=ProviderModelDiscoveryRequest( + custom_llm_provider="anthropic", api_key="os.environ/ANTHROPIC_API_KEY" + ), + user_api_key_dict=self._admin(), + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_unknown_provider_is_rejected(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + discover_provider_models, + ) + from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + ProviderModelDiscoveryRequest, + ) + + with pytest.raises(HTTPException) as exc_info: + await discover_provider_models( + data=ProviderModelDiscoveryRequest(custom_llm_provider="not-a-real-provider"), + user_api_key_dict=self._admin(), + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_credential_provider_mismatch_is_rejected(self, monkeypatch): + 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-key", + credential_values={"api_key": "sk-x"}, + credential_info={"custom_llm_provider": "openai"}, + ) + ], + ) + with patch("litellm.proxy.proxy_server.prisma_client", MagicMock()): + with pytest.raises(HTTPException) as exc_info: + await discover_provider_models( + data=ProviderModelDiscoveryRequest( + custom_llm_provider="anthropic", litellm_credential_name="openai-key" + ), + user_api_key_dict=self._admin(), + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_unknown_credential_name_is_404(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + discover_provider_models, + ) + from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + ProviderModelDiscoveryRequest, + ) + + with patch("litellm.proxy.proxy_server.prisma_client", None): + with pytest.raises(HTTPException) as exc_info: + await discover_provider_models( + data=ProviderModelDiscoveryRequest( + custom_llm_provider="anthropic", litellm_credential_name="does-not-exist" + ), + user_api_key_dict=self._admin(), + ) + assert exc_info.value.status_code == 404 + + @pytest.mark.asyncio + async def test_discovery_success_via_named_wif_credential(self, monkeypatch): + """The end-to-end contract: a WIF credential's fields never appear in the request + body, only its name does, and discovery still succeeds by hydrating them server-side.""" + 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="anthropic-wif", + credential_values={ + "anthropic_federation_rule_id": "rule-1", + "anthropic_organization_id": "org-1", + "anthropic_identity_token": "oidc/env/TOK", + }, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ], + ) + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch( + "litellm.llms.anthropic.common_utils.AnthropicModelInfo.discover_models", + return_value=["anthropic/claude-disc"], + ) as discover_mock, + ): + result = await discover_provider_models( + data=ProviderModelDiscoveryRequest( + custom_llm_provider="anthropic", litellm_credential_name="anthropic-wif" + ), + user_api_key_dict=self._admin(), + ) + assert result.models == ["anthropic/claude-disc"] + 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["anthropic_federation_rule_id"] == "rule-1" + assert called_params["anthropic_identity_token"] == "oidc/env/TOK" + + @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 ( + discover_provider_models, + ) + from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + ProviderModelDiscoveryRequest, + ) + + with patch( + "litellm.llms.anthropic.common_utils.AnthropicModelInfo.discover_models", + side_effect=Exception("Failed to fetch models from Anthropic. HTTP 401: invalid x-api-key"), + ): + with pytest.raises(HTTPException) as exc_info: + await discover_provider_models( + data=ProviderModelDiscoveryRequest(custom_llm_provider="anthropic", api_key="sk-ant-bad"), + user_api_key_dict=self._admin(), + ) + assert exc_info.value.status_code == 502 + assert "invalid x-api-key" in str(exc_info.value.detail) + + +class TestOneCredentialFeedsManyModelsNoWifCopy: + """Regression: one named WIF credential feeds multiple model rows, and no WIF field is + ever copied onto a model row -- litellm_params carries only `model` and + `litellm_credential_name`, the same shape the wizard's per-row /model/new call produces.""" + + @pytest.mark.asyncio + async def test_two_discovered_models_share_the_credential_reference_only(self): + from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _add_model_to_db, + ) + from litellm.types.router import ModelInfo + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=MagicMock()) + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch("litellm.proxy.proxy_server.master_key", "sk-test-master"), + patch("litellm.proxy.common_utils.encrypt_decrypt_utils._get_salt_key", return_value="sk-test-master"), + ): + for i, discovered_id in enumerate(["claude-a", "claude-b"]): + model_params = Deployment( + model_name=discovered_id, + litellm_params=LiteLLM_Params( + model=f"anthropic/{discovered_id}", litellm_credential_name="anthropic-wif" + ), + model_info=ModelInfo(id=f"dep-shared-{i}"), + blocked=False, + ) + await _add_model_to_db(model_params=model_params, user_api_key_dict=admin, prisma_client=mock_prisma) + + assert mock_prisma.db.litellm_proxymodeltable.create.await_count == 2 + for call in mock_prisma.db.litellm_proxymodeltable.create.await_args_list: + written_litellm_params = json.loads(call.kwargs["data"]["litellm_params"]) + decrypted_credential_name = decrypt_value_helper( + value=written_litellm_params["litellm_credential_name"], key="litellm_credential_name" + ) + assert decrypted_credential_name == "anthropic-wif" + assert "anthropic_federation_rule_id" not in written_litellm_params + assert "anthropic_identity_token" not in written_litellm_params + assert call.kwargs["data"]["blocked"] is False 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 31430da71e8..e6ffa8ef926 100644 --- a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py @@ -2,16 +2,20 @@ from datetime import datetime, timezone from unittest.mock import AsyncMock, MagicMock, patch import pytest - - from fastapi import FastAPI from fastapi.testclient import TestClient +from pydantic import ValidationError from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.public_endpoints import router from litellm.types.proxy.management_endpoints.model_management_endpoints import ( ModelGroupInfoProxy, ) +from litellm.types.proxy.public_endpoints.public_endpoints import ( + ProviderCredentialField, + ProviderCredentialVariant, + ProviderCredentialVariants, +) from litellm.types.utils import LlmProviders @@ -1037,3 +1041,133 @@ def test_public_mcp_hub_does_not_expose_upstream_url(): assert all("url" not in item for item in data) assert secret_url not in response.text app.dependency_overrides.clear() + + +def _field(key: str) -> ProviderCredentialField: + return ProviderCredentialField(key=key, label=key) + + +def test_credential_variants_rejects_duplicate_variant_ids(): + with pytest.raises(ValidationError, match="unique"): + ProviderCredentialVariants( + selector_label="Auth method", + default_variant="a", + field_definitions=[_field("x")], + variants=[ + ProviderCredentialVariant(id="a", label="A", field_keys=["x"]), + ProviderCredentialVariant(id="a", label="A again", field_keys=["x"]), + ], + ) + + +def test_credential_variants_rejects_unknown_default_variant(): + with pytest.raises(ValidationError, match="default_variant"): + ProviderCredentialVariants( + selector_label="Auth method", + default_variant="does-not-exist", + field_definitions=[_field("x")], + variants=[ProviderCredentialVariant(id="a", label="A", field_keys=["x"])], + ) + + +def test_credential_variants_rejects_unresolved_field_keys(): + with pytest.raises(ValidationError, match="undefined field_keys"): + ProviderCredentialVariants( + selector_label="Auth method", + default_variant="a", + field_definitions=[_field("x")], + variants=[ProviderCredentialVariant(id="a", label="A", field_keys=["x", "does-not-exist"])], + ) + + +def test_credential_variants_rejects_fixed_values_overlapping_field_keys(): + with pytest.raises(ValidationError, match="overlapping"): + ProviderCredentialVariants( + selector_label="Auth method", + default_variant="a", + field_definitions=[_field("x")], + variants=[ + ProviderCredentialVariant(id="a", label="A", field_keys=["x"], fixed_values={"x": "fixed"}), + ], + ) + + +def test_credential_variants_accepts_a_well_formed_schema(): + variants = ProviderCredentialVariants( + selector_label="Auth method", + default_variant="a", + field_definitions=[_field("x"), _field("y")], + variants=[ + ProviderCredentialVariant(id="a", label="A", field_keys=["x"]), + ProviderCredentialVariant(id="b", label="B", field_keys=["y"], fixed_values={"mode": "b"}), + ], + ) + assert [v.id for v in variants.variants] == ["a", "b"] + + +def test_anthropic_provider_fields_expose_credential_variants(): + """The Anthropic provider must publish a variant selector covering API-key auth + and every workload-identity-federation identity source, while the legacy + credential_fields stays exactly api_base + api_key for old dashboards.""" + 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() + + anthropic = next((p for p in providers if p["provider"] == "Anthropic"), None) + assert anthropic is not None + + # Legacy fallback is untouched: still exactly api_base + api_key. + assert {f["key"] for f in anthropic["credential_fields"]} == {"api_base", "api_key"} + + variants_block = anthropic["credential_variants"] + assert variants_block is not None + assert variants_block["default_variant"] == "api_key" + + variant_ids = {v["id"] for v in variants_block["variants"]} + assert variant_ids == { + "api_key", + "wif_token", + "wif_token_file", + "wif_internal_issuer", + "wif_keycloak", + } + + field_defs_by_key = {f["key"]: f for f in variants_block["field_definitions"]} + + # Every *_ref secret-pointer field is a plain text input, never password: pasting a + # secret into it would be exactly the exfiltration primitive the WIF hardening closed. + for ref_key in ("anthropic_issuer_signing_key_ref", "anthropic_keycloak_client_secret_ref", "anthropic_identity_token"): + assert field_defs_by_key[ref_key]["field_type"] == "text", f"{ref_key} must not render as a password field" + + variants_by_id = {v["id"]: v for v in variants_block["variants"]} + assert variants_by_id["wif_internal_issuer"]["fixed_values"] == {"anthropic_identity_source": "internal_issuer"} + assert variants_by_id["wif_keycloak"]["fixed_values"] == {"anthropic_identity_source": "keycloak"} + assert variants_by_id["api_key"]["fixed_values"] == {} + + # Every field_key referenced by a variant must resolve in field_definitions -- if this + # ever drifted the endpoint's own response_model validation would already 500, but assert + # it explicitly so a future edit fails fast in this test instead. + 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 + see no change at all.""" + 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 + assert openai.get("credential_variants") is None diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx index 94737d88d0f..d57941ede43 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx @@ -17,6 +17,7 @@ import { useModelDashboardData } from "@/app/(dashboard)/models-and-endpoints/us import AllModelsPanel from "@/app/(dashboard)/models-and-endpoints/panels/AllModelsPanel"; import AutoRoutersTabPanel from "@/app/(dashboard)/models-and-endpoints/panels/AutoRoutersTabPanel"; import AddModelPanel from "@/app/(dashboard)/models-and-endpoints/panels/AddModelPanel"; +import AddProviderPanel from "@/app/(dashboard)/models-and-endpoints/panels/add-provider/AddProviderPanel"; import LlmCredentialsPanel from "@/app/(dashboard)/models-and-endpoints/panels/LlmCredentialsPanel"; import PassThroughPanel from "@/app/(dashboard)/models-and-endpoints/panels/PassThroughPanel"; import HealthStatusPanel from "@/app/(dashboard)/models-and-endpoints/panels/HealthStatusPanel"; @@ -28,6 +29,7 @@ import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; type ModelTabSlug = | "add" + | "add-provider" | "auto-routers" | "llm-credentials" | "pass-through" @@ -40,6 +42,7 @@ const BASE_TAB_KEY = "all-models"; const TAB_LABELS: Record = { add: "Add Model", + "add-provider": "Add Provider", "auto-routers": "Auto-Routers", "llm-credentials": "LLM Credentials", "pass-through": "Pass-Through Endpoints", @@ -57,6 +60,8 @@ const renderPanel = (key: string) => { return ; case "add": return ; + case "add-provider": + return ; case "llm-credentials": return ; case "pass-through": @@ -100,6 +105,7 @@ export default function ModelsAndEndpointsPage() { () => [ "", ...(canCreate ? (["add"] as const) : []), + ...(isAdmin ? (["add-provider"] as const) : []), ...(isAdmin || canCreate ? (["auto-routers"] as const) : []), ...(isAdmin ? (["llm-credentials", "pass-through", "health", "retry-settings", "model-group-alias", "price-data"] as const) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/AddProviderPanel.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/AddProviderPanel.integration.test.tsx new file mode 100644 index 00000000000..cae7cba2cf0 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/AddProviderPanel.integration.test.tsx @@ -0,0 +1,203 @@ +import { renderWithProviders, screen, waitFor, within } from "../../../../../../tests/test-utils"; +import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import AddProviderPanel from "./AddProviderPanel"; + +const discoverProviderModelsCall = vi.fn(); +const credentialCreateCall = vi.fn(); +const credentialUpdateCall = vi.fn(); +const createProviderModelCall = vi.fn(); +const listAllModelsCall = vi.fn(); +const getCallbacksCall = vi.fn(); +const setCallbacksCall = vi.fn(); +const mockAuthorized = vi.fn(); + +vi.mock("@/components/networking", async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + discoverProviderModelsCall: (...args: unknown[]) => discoverProviderModelsCall(...args), + credentialCreateCall: (...args: unknown[]) => credentialCreateCall(...args), + credentialUpdateCall: (...args: unknown[]) => credentialUpdateCall(...args), + createProviderModelCall: (...args: unknown[]) => createProviderModelCall(...args), + listAllModelsCall: (...args: unknown[]) => listAllModelsCall(...args), + getCallbacksCall: (...args: unknown[]) => getCallbacksCall(...args), + setCallbacksCall: (...args: unknown[]) => setCallbacksCall(...args), + }; +}); + +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: () => mockAuthorized() })); + +vi.mock("@/app/(dashboard)/hooks/credentials/useCredentials", () => ({ + useCredentials: () => ({ data: { credentials: [] } }), +})); + +vi.mock("@/app/(dashboard)/hooks/providers/useProviderFields", () => ({ + useProviderFields: () => ({ + data: [ + { + provider: "Anthropic", + provider_display_name: "Anthropic", + litellm_provider: "anthropic", + default_model_placeholder: "claude-3-opus", + credential_fields: [ + { key: "api_base", label: "Upstream API Base", field_type: "text" }, + { key: "api_key", label: "API Key", field_type: "password" }, + ], + credential_variants: { + selector_label: "Authentication method", + default_variant: "api_key", + field_definitions: [ + { key: "api_base", label: "Upstream API Base", field_type: "text" }, + { key: "api_key", label: "API Key", field_type: "password" }, + ], + variants: [{ id: "api_key", label: "API Key", field_keys: ["api_base", "api_key"], fixed_values: {} }], + }, + }, + ], + isLoading: false, + error: null, + }), +})); + +const PROXY_ADMIN = { accessToken: "test-access-token" }; + +const setup = async () => { + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + renderWithProviders(); + await screen.findByLabelText("Provider"); + return { user }; +}; + +const chooseProvider = async (user: ReturnType, name: string) => { + await user.click(screen.getByLabelText("Provider")); + await user.click(await screen.findByText(name)); +}; + +const rowFor = (upstreamId: string) => within(screen.getByText(upstreamId).closest("tr") as HTMLElement); + +describe("AddProviderPanel", () => { + beforeEach(() => { + vi.clearAllMocks(); + mockAuthorized.mockReturnValue(PROXY_ADMIN); + credentialCreateCall.mockResolvedValue({}); + credentialUpdateCall.mockResolvedValue({}); + listAllModelsCall.mockResolvedValue({ data: [] }); + createProviderModelCall.mockResolvedValue({ model_id: "new-id" }); + getCallbacksCall.mockResolvedValue({ router_settings: {} }); + setCallbacksCall.mockResolvedValue({}); + }); + + it("walks provider -> credential -> discover -> review -> create, with blocked and aliases wired correctly", async () => { + discoverProviderModelsCall.mockResolvedValue({ models: ["claude-3-opus", "claude-3-haiku"] }); + const { user } = await setup(); + + await chooseProvider(user, "Anthropic"); + await user.type(screen.getByLabelText("Credential name"), "anthropic-prod"); + await user.click(screen.getByRole("button", { name: /Next/ })); + + await user.type(await screen.findByLabelText("API Key"), "sk-ant-test"); + await user.click(screen.getByRole("button", { name: "Save credential" })); + + expect(await screen.findByText("claude-3-opus")).toBeInTheDocument(); + expect(credentialCreateCall).toHaveBeenCalledWith("test-access-token", { + credential_name: "anthropic-prod", + credential_values: { api_key: "sk-ant-test" }, + credential_info: { custom_llm_provider: "anthropic" }, + }); + expect(discoverProviderModelsCall).toHaveBeenCalledWith("test-access-token", { + custom_llm_provider: "anthropic", + litellm_credential_name: "anthropic-prod", + }); + + // Disable the first discovered row. + await user.click(rowFor("claude-3-opus").getByRole("switch")); + + // Add an alternate name to the second discovered row. + const haikuAltNames = rowFor("claude-3-haiku").getByRole("combobox"); + await user.type(haikuAltNames, "gpt-4o-mini"); + await user.click(await screen.findByText('Create "gpt-4o-mini"')); + + // Add a manual (hidden) model. + await user.type(screen.getByPlaceholderText("upstream model id"), "claude-hidden"); + await user.click(screen.getByRole("button", { name: "Add" })); + + await user.click(screen.getByRole("button", { name: "Create 3 models" })); + + await waitFor(() => expect(createProviderModelCall).toHaveBeenCalledTimes(3)); + expect(createProviderModelCall).toHaveBeenCalledWith("test-access-token", { + model_name: "claude-3-opus", + litellm_params: { model: "anthropic/claude-3-opus", litellm_credential_name: "anthropic-prod" }, + model_info: {}, + blocked: true, + }); + expect(createProviderModelCall).toHaveBeenCalledWith("test-access-token", { + model_name: "claude-3-haiku", + litellm_params: { model: "anthropic/claude-3-haiku", litellm_credential_name: "anthropic-prod" }, + model_info: {}, + blocked: false, + }); + expect(createProviderModelCall).toHaveBeenCalledWith("test-access-token", { + model_name: "claude-hidden", + litellm_params: { model: "anthropic/claude-hidden", litellm_credential_name: "anthropic-prod" }, + model_info: {}, + blocked: false, + }); + + await waitFor(() => + expect(setCallbacksCall).toHaveBeenCalledWith("test-access-token", { + router_settings: { model_group_alias: { "gpt-4o-mini": "claude-3-haiku" } }, + }), + ); + + expect(await screen.findByText(/claude-3-opus: created/)).toBeInTheDocument(); + expect(screen.getByText(/claude-3-haiku: created/)).toBeInTheDocument(); + expect(screen.getByText(/claude-hidden: created/)).toBeInTheDocument(); + }); + + it("shows a sanitized discovery error with a working retry", async () => { + discoverProviderModelsCall.mockRejectedValueOnce(new Error("upstream auth failed")); + discoverProviderModelsCall.mockResolvedValueOnce({ models: ["claude-3-opus"] }); + const { user } = await setup(); + + await chooseProvider(user, "Anthropic"); + await user.type(screen.getByLabelText("Credential name"), "anthropic-prod"); + await user.click(screen.getByRole("button", { name: /Next/ })); + await user.type(await screen.findByLabelText("API Key"), "sk-ant-test"); + await user.click(screen.getByRole("button", { name: "Save credential" })); + + expect(await screen.findByText("Discovery failed")).toBeInTheDocument(); + expect(screen.getByText("upstream auth failed")).toBeInTheDocument(); + + await user.click(screen.getByRole("button", { name: "Retry" })); + + expect(await screen.findByText("claude-3-opus")).toBeInTheDocument(); + expect(discoverProviderModelsCall).toHaveBeenCalledTimes(2); + }); + + it("skips a row already created under this credential on a re-run", async () => { + discoverProviderModelsCall.mockResolvedValue({ models: ["claude-3-opus"] }); + listAllModelsCall.mockResolvedValue({ + data: [ + { + model_name: "claude-3-opus", + litellm_params: { model: "anthropic/claude-3-opus", litellm_credential_name: "anthropic-prod" }, + model_info: { id: "existing-id" }, + }, + ], + }); + const { user } = await setup(); + + await chooseProvider(user, "Anthropic"); + await user.type(screen.getByLabelText("Credential name"), "anthropic-prod"); + await user.click(screen.getByRole("button", { name: /Next/ })); + await user.type(await screen.findByLabelText("API Key"), "sk-ant-test"); + await user.click(screen.getByRole("button", { name: "Save credential" })); + + await screen.findByText("claude-3-opus"); + await user.click(screen.getByRole("button", { name: "Create 1 model" })); + + expect(await screen.findByText(/claude-3-opus: skipped/)).toBeInTheDocument(); + expect(createProviderModelCall).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/AddProviderPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/AddProviderPanel.tsx new file mode 100644 index 00000000000..eae5b5bb0bb --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/AddProviderPanel.tsx @@ -0,0 +1,364 @@ +"use client"; + +import React from "react"; +import { useForm, FormProvider } from "react-hook-form"; +import { useQueryClient } from "@tanstack/react-query"; +import { Button } from "@/components/ui/button"; +import { Card, CardContent } from "@/components/ui/card"; +import { toast } from "@/lib/toast"; +import { useProviderFields } from "@/app/(dashboard)/hooks/providers/useProviderFields"; +import { useCredentials } from "@/app/(dashboard)/hooks/credentials/useCredentials"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { + MountedFormProvider, + projectMountedValues, + useMountRegistry, + type MountedFormValues, +} from "@/components/common_components/MountedFormField"; +import { computeCredentialValuesToDelete } from "@/components/model_add/credential_form_helpers"; +import ProviderSpecificFields from "@/components/add_model/provider_specific_fields"; +import { ProviderLogo } from "@/components/molecules/models/ProviderLogo"; +import { Providers } from "@/components/provider_info_helpers"; +import { + credentialCreateCall, + credentialUpdateCall, + discoverProviderModelsCall, + getCredentialJwksCall, + listAllModelsCall, + getCallbacksCall, + setCallbacksCall, + createProviderModelCall, + type AnthropicJwks, + type DeploymentInfoRow, + type ProviderCreateInfo, +} from "@/components/networking"; +import type { SearchSelectOption } from "@/components/shared/SearchSelect"; +import { ArrowLeft } from "lucide-react"; +import { extractProxyErrorMessage } from "@/lib/http/client"; +import { + aliasAdditionsFromRows, + buildDiscoveredRows, + buildModelCreationPayload, + mergeModelGroupAliases, + rowsPendingCreation, + type CreationResult, + type DiscoveredModelRow, + type ModelGroupAliasMap, +} from "./wizardLogic"; +import ReviewModelsStep from "./ReviewModelsStep"; +import { DiscoverStep, JwksStep, ProviderStep, ResultsStep } from "./WizardSteps"; + +type WizardStep = "provider" | "credential" | "jwks" | "discover" | "review" | "creating" | "done"; + +const STEP_ORDER: readonly WizardStep[] = ["provider", "credential", "jwks", "discover", "review", "creating", "done"]; + +const STEP_LABELS: Record = { + provider: "Provider", + credential: "Authentication", + jwks: "Register issuer", + discover: "Discover models", + review: "Review models", + creating: "Creating", + done: "Done", +}; + +const ANTHROPIC_INTERNAL_ISSUER_DISCRIMINATOR = "internal_issuer"; + +const StepIndicator: React.FC<{ step: WizardStep; skipJwks: boolean }> = ({ step, skipJwks }) => { + const visibleSteps = STEP_ORDER.filter((s) => s !== "creating" && (!skipJwks || s !== "jwks")); + const currentIndex = visibleSteps.indexOf(step === "creating" ? "done" : step); + return ( +
+ {visibleSteps.map((s, index) => ( + + {index > 0 && {"->"}} + + {STEP_LABELS[s]} + + + ))} +
+ ); +}; + +export default function AddProviderPanel() { + const { accessToken } = useAuthorized(); + const queryClient = useQueryClient(); + const { data: providerMetadata } = useProviderFields(); + const { data: credentialsResponse } = useCredentials(); + + const [step, setStep] = React.useState("provider"); + const [selectedProvider, setSelectedProvider] = React.useState(null); + const [credentialName, setCredentialName] = React.useState(""); + const [credentialSaved, setCredentialSaved] = React.useState(false); + const [savedValues, setSavedValues] = React.useState>({}); + const [federationRuleId, setFederationRuleId] = React.useState(""); + const [jwks, setJwks] = React.useState(null); + const [jwksError, setJwksError] = React.useState(null); + const [discoveryError, setDiscoveryError] = React.useState(null); + const [isDiscovering, setIsDiscovering] = React.useState(false); + const [rows, setRows] = React.useState([]); + const [isCreating, setIsCreating] = React.useState(false); + const [creationResults, setCreationResults] = React.useState([]); + const [aliasCollisions, setAliasCollisions] = React.useState([]); + + const form = useForm({ mode: "onChange" }); + const registry = useMountRegistry(); + + const providerOptions: SearchSelectOption[] = React.useMemo( + () => + (providerMetadata ?? []) + .slice() + .sort((a, b) => a.provider_display_name.localeCompare(b.provider_display_name)) + .map((p) => ({ + label: p.provider_display_name, + value: p.provider_display_name, + icon: , + })), + [providerMetadata], + ); + + const selectedProviderInfo: ProviderCreateInfo | undefined = React.useMemo( + () => providerMetadata?.find((p) => p.provider_display_name === selectedProvider), + [providerMetadata, selectedProvider], + ); + const litellmProvider = selectedProviderInfo?.litellm_provider ?? ""; + + const nameCollision = + credentialName.length > 0 && + (credentialsResponse?.credentials ?? []).some((c) => c.credential_name === credentialName); + + const goTo = (next: WizardStep) => setStep(next); + + const saveCredential = async () => { + if (!accessToken || !selectedProvider) { + return; + } + const isValid = await form.trigger(registry.mountedNames() as string[]); + if (!isValid) { + return; + } + const values = projectMountedValues(registry, form.getValues) as Record; + const nonEmptyValues = Object.fromEntries( + Object.entries(values).filter(([, v]) => v !== "" && v !== undefined && v !== null), + ); + try { + if (!credentialSaved) { + await credentialCreateCall(accessToken, { + credential_name: credentialName, + credential_values: nonEmptyValues, + credential_info: { custom_llm_provider: litellmProvider }, + }); + } else { + const credentialValuesToDelete = computeCredentialValuesToDelete(savedValues, values); + const updatePayload = { + credential_name: credentialName, + credential_values: nonEmptyValues, + credential_info: { custom_llm_provider: litellmProvider }, + ...(credentialValuesToDelete.length > 0 ? { credential_values_to_delete: credentialValuesToDelete } : {}), + }; + await credentialUpdateCall(accessToken, credentialName, updatePayload); + } + setSavedValues(values); + setCredentialSaved(true); + setFederationRuleId(typeof values.anthropic_federation_rule_id === "string" ? values.anthropic_federation_rule_id : ""); + queryClient.invalidateQueries({ queryKey: ["credentials"] }); + toast.success(`Credential "${credentialName}" saved`); + if (values.anthropic_identity_source === ANTHROPIC_INTERNAL_ISSUER_DISCRIMINATOR) { + goTo("jwks"); + void loadJwks(); + } else { + goTo("discover"); + void runDiscovery(); + } + } catch (error) { + toast.fromError(`Failed to save credential: ${extractProxyErrorMessage(error)}`); + } + }; + + const loadJwks = async () => { + if (!accessToken) return; + setJwksError(null); + try { + const result = await getCredentialJwksCall(accessToken, credentialName); + setJwks(result); + } catch (error) { + setJwksError(extractProxyErrorMessage(error)); + } + }; + + const confirmFederationRuleId = async () => { + if (!accessToken) return; + if (federationRuleId !== savedValues.anthropic_federation_rule_id) { + try { + await credentialUpdateCall(accessToken, credentialName, { + credential_name: credentialName, + credential_values: { anthropic_federation_rule_id: federationRuleId }, + credential_info: { custom_llm_provider: litellmProvider }, + }); + setSavedValues((prev) => ({ ...prev, anthropic_federation_rule_id: federationRuleId })); + } catch (error) { + toast.fromError(`Failed to save the federation rule id: ${extractProxyErrorMessage(error)}`); + return; + } + } + goTo("discover"); + void runDiscovery(); + }; + + const runDiscovery = async () => { + if (!accessToken) return; + setIsDiscovering(true); + setDiscoveryError(null); + try { + const result = await discoverProviderModelsCall(accessToken, { + custom_llm_provider: litellmProvider, + litellm_credential_name: credentialName, + }); + setRows(buildDiscoveredRows(result.models)); + goTo("review"); + } catch (error) { + setDiscoveryError(extractProxyErrorMessage(error)); + } finally { + setIsDiscovering(false); + } + }; + + const createModels = async () => { + if (!accessToken) return; + setIsCreating(true); + setCreationResults([]); + setAliasCollisions([]); + goTo("creating"); + + let existing: DeploymentInfoRow[] = []; + try { + existing = (await listAllModelsCall(accessToken)).data; + } catch { + existing = []; + } + const pending = rowsPendingCreation(rows, litellmProvider, credentialName, existing); + const pendingIds = new Set(pending.map((r) => r.id)); + + const results: CreationResult[] = []; + for (const row of rows) { + if (!pendingIds.has(row.id)) { + results.push({ row, status: "skipped", detail: "already created" }); + continue; + } + try { + await createProviderModelCall(accessToken, buildModelCreationPayload(litellmProvider, credentialName, row)); + results.push({ row, status: "created" }); + } catch (error) { + results.push({ row, status: "failed", detail: extractProxyErrorMessage(error) }); + } + } + setCreationResults(results); + + const additions = aliasAdditionsFromRows(rows); + if (additions.length > 0) { + try { + const config = await getCallbacksCall(accessToken, "", ""); + const existingAliasMap: ModelGroupAliasMap = config?.router_settings?.model_group_alias ?? {}; + const { merged, collisions } = mergeModelGroupAliases(existingAliasMap, additions); + if (collisions.length > 0) { + setAliasCollisions([...collisions]); + } + await setCallbacksCall(accessToken, { router_settings: { model_group_alias: merged } }); + } catch (error) { + toast.fromError(`Failed to save alternate names: ${extractProxyErrorMessage(error)}`); + } + } + + queryClient.invalidateQueries({ queryKey: ["models", "list"] }); + setIsCreating(false); + goTo("done"); + }; + + const isInternalIssuer = savedValues.anthropic_identity_source === ANTHROPIC_INTERNAL_ISSUER_DISCRIMINATOR; + + return ( +
+

Add Provider

+ + + {step === "provider" && ( + goTo("credential")} + /> + )} + + {step === "credential" && selectedProvider && ( + + + + +
{ + e.preventDefault(); + void saveCredential(); + }} + > + +
+ + +
+ +
+
+
+
+ )} + + {step === "jwks" && ( + goTo("credential")} + onNext={() => void confirmFederationRuleId()} + /> + )} + + {step === "discover" && ( + goTo(isInternalIssuer ? "jwks" : "credential")} + onRetry={() => void runDiscovery()} + /> + )} + + {step === "review" && ( + { + goTo("discover"); + void runDiscovery(); + }} + onCreateModels={() => void createModels()} + /> + )} + + {(step === "creating" || step === "done") && ( + + )} +
+ ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/ReviewModelsStep.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/ReviewModelsStep.tsx new file mode 100644 index 00000000000..a7a9b8525d4 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/ReviewModelsStep.tsx @@ -0,0 +1,129 @@ +"use client"; + +import React from "react"; +import { Button } from "@/components/ui/button"; +import { Card, CardContent } from "@/components/ui/card"; +import { Input } from "@/components/ui/input"; +import { Switch } from "@/components/ui/switch"; +import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import { ArrowLeft, Plus, X } from "lucide-react"; +import { buildManualRow, type DiscoveredModelRow } from "./wizardLogic"; + +interface ReviewModelsStepProps { + rows: DiscoveredModelRow[]; + setRows: React.Dispatch>; + onBack: () => void; + onCreateModels: () => void; +} + +const ReviewModelsStep: React.FC = ({ rows, setRows, onBack, onCreateModels }) => { + const [manualId, setManualId] = React.useState(""); + + const updateRow = (id: string, patch: Partial) => + setRows((current) => current.map((row) => (row.id === id ? { ...row, ...patch } : row))); + + const removeRow = (id: string) => setRows((current) => current.filter((row) => row.id !== id)); + + const addManualRow = () => { + const trimmed = manualId.trim(); + if (!trimmed) return; + setRows((current) => [...current, buildManualRow(trimmed)]); + setManualId(""); + }; + + return ( + + + {rows.length === 0 ? ( +

No models discovered. Add one manually below.

+ ) : ( +
+ + + + Upstream ID + Enabled + Model name + Alternate names + + + + + {rows.map((row) => ( + + {row.upstreamId} + + updateRow(row.id, { enabled: checked })} + aria-label={`Enable ${row.upstreamId}`} + /> + + + updateRow(row.id, { modelName: e.target.value })} + aria-label={`Model name for ${row.upstreamId}`} + /> + + + updateRow(row.id, { alternateNames: value })} + options={[]} + allowCustomValues + placeholder="Add alternate names" + emptyText="Type to add an alternate name" + /> + + + {row.manual && ( + + )} + + + ))} + +
+
+ )} + +
+
+ + setManualId(e.target.value)} + placeholder="upstream model id" + /> +
+ +
+ +
+ + +
+
+
+ ); +}; + +export default ReviewModelsStep; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/WizardSteps.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/WizardSteps.tsx new file mode 100644 index 00000000000..6875fecfb59 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/WizardSteps.tsx @@ -0,0 +1,207 @@ +"use client"; + +import React from "react"; +import { Button } from "@/components/ui/button"; +import { Card, CardContent } from "@/components/ui/card"; +import { Input } from "@/components/ui/input"; +import { Field, FieldLabel } from "@/components/shared/form/field"; +import { SearchSelect, type SearchSelectOption } from "@/components/shared/SearchSelect"; +import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert"; +import CopyButton from "@/components/shared/CopyButton"; +import { AlertTriangle, ArrowLeft, ArrowRight, Loader2 } from "lucide-react"; +import { Providers } from "@/components/provider_info_helpers"; +import type { AnthropicJwks } from "@/components/networking"; +import type { CreationResult } from "./wizardLogic"; + +const CREATION_RESULT_CLASS_NAME: Record = { + failed: "text-destructive", + skipped: "text-muted-foreground", + created: "text-success", +}; + +interface ProviderStepProps { + providerOptions: SearchSelectOption[]; + selectedProvider: Providers | null; + onSelectProvider: (provider: Providers) => void; + credentialName: string; + onCredentialNameChange: (name: string) => void; + nameCollision: boolean; + onNext: () => void; +} + +export const ProviderStep: React.FC = ({ + providerOptions, + selectedProvider, + onSelectProvider, + credentialName, + onCredentialNameChange, + nameCollision, + onNext, +}) => ( + + + + Provider + onSelectProvider(value as Providers)} + /> + + + Credential name + onCredentialNameChange(e.target.value)} + placeholder="e.g. anthropic-prod" + /> + {nameCollision &&

A credential with this name already exists.

} +
+
+ +
+
+
+); + +interface JwksStepProps { + jwks: AnthropicJwks | null; + jwksError: string | null; + federationRuleId: string; + onFederationRuleIdChange: (value: string) => void; + onBack: () => void; + onNext: () => void; +} + +export const JwksStep: React.FC = ({ + jwks, + jwksError, + federationRuleId, + onFederationRuleIdChange, + onBack, + onNext, +}) => ( + + + + Register this JWKS with Anthropic + + Register this public JWKS as the inline issuer for your federation rule in the Anthropic Console, then + paste the resulting Federation Rule ID below. + + + {jwksError && ( + + + Could not load JWKS + {jwksError} + + )} + {jwks && ( +
+ +
{JSON.stringify(jwks, null, 2)}
+
+ )} + + Federation Rule ID + onFederationRuleIdChange(e.target.value)} + /> + +
+ + +
+
+
+); + +interface DiscoverStepProps { + isDiscovering: boolean; + discoveryError: string | null; + onBack: () => void; + onRetry: () => void; +} + +export const DiscoverStep: React.FC = ({ isDiscovering, discoveryError, onBack, onRetry }) => ( + + + {isDiscovering && ( +

+ Discovering models... +

+ )} + {discoveryError && ( + + + Discovery failed + {discoveryError} + + )} +
+ + {discoveryError && ( + + )} +
+
+
+); + +interface ResultsStepProps { + isCreating: boolean; + isDone: boolean; + creationResults: CreationResult[]; + aliasCollisions: string[]; +} + +export const ResultsStep: React.FC = ({ isCreating, isDone, creationResults, aliasCollisions }) => ( + + + {isCreating && ( +

+ Creating models... +

+ )} + {isDone && ( + <> +
    + {creationResults.map((result) => ( +
  • + + {result.row.modelName}: {result.status} + {result.detail ? ` (${result.detail})` : ""} + +
  • + ))} +
+ {aliasCollisions.length > 0 && ( + + + Some alternate names were not saved + + These alias names already exist and were left unchanged: {aliasCollisions.join(", ")} + + + )} + + )} +
+
+); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/wizardLogic.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/wizardLogic.test.ts new file mode 100644 index 00000000000..fc945d82697 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/wizardLogic.test.ts @@ -0,0 +1,149 @@ +import { describe, expect, it } from "vitest"; +import type { DeploymentInfoRow } from "@/components/networking"; +import { + aliasAdditionsFromRows, + buildDiscoveredRows, + buildManualRow, + buildModelCreationPayload, + litellmModelForUpstreamId, + mergeModelGroupAliases, + rowsPendingCreation, + type DiscoveredModelRow, +} from "./wizardLogic"; + +describe("litellmModelForUpstreamId", () => { + it("prefixes a bare upstream id with the provider", () => { + expect(litellmModelForUpstreamId("anthropic", "claude-3-opus")).toBe("anthropic/claude-3-opus"); + }); + + it("does not double-prefix an id that already carries the provider prefix", () => { + expect(litellmModelForUpstreamId("anthropic", "anthropic/claude-3-opus")).toBe("anthropic/claude-3-opus"); + }); +}); + +describe("buildDiscoveredRows", () => { + it("defaults every row to enabled, with model_name equal to the discovered id", () => { + const rows = buildDiscoveredRows(["claude-3-opus", "claude-3-haiku"]); + expect(rows).toHaveLength(2); + expect(rows[0]).toMatchObject({ + upstreamId: "claude-3-opus", + modelName: "claude-3-opus", + enabled: true, + alternateNames: [], + manual: false, + }); + }); + + it("gives every row a unique id even for duplicate upstream ids", () => { + const rows = buildDiscoveredRows(["same-id", "same-id"]); + expect(rows[0].id).not.toBe(rows[1].id); + }); +}); + +describe("buildManualRow", () => { + it("marks the row manual and enabled by default", () => { + const row = buildManualRow("hidden-model"); + expect(row).toMatchObject({ upstreamId: "hidden-model", modelName: "hidden-model", enabled: true, manual: true }); + }); +}); + +describe("buildModelCreationPayload", () => { + const baseRow: DiscoveredModelRow = { + id: "row-1", + upstreamId: "claude-3-opus", + modelName: "my-claude", + enabled: true, + alternateNames: [], + manual: false, + }; + + it("maps enabled=true to blocked=false", () => { + const payload = buildModelCreationPayload("anthropic", "my-cred", baseRow); + expect(payload).toEqual({ + model_name: "my-claude", + litellm_params: { model: "anthropic/claude-3-opus", litellm_credential_name: "my-cred" }, + model_info: {}, + blocked: false, + }); + }); + + it("maps enabled=false to blocked=true", () => { + const payload = buildModelCreationPayload("anthropic", "my-cred", { ...baseRow, enabled: false }); + expect(payload.blocked).toBe(true); + }); +}); + +describe("rowsPendingCreation", () => { + const rows: DiscoveredModelRow[] = [ + { id: "1", upstreamId: "claude-3-opus", modelName: "opus", enabled: true, alternateNames: [], manual: false }, + { id: "2", upstreamId: "claude-3-haiku", modelName: "haiku", enabled: true, alternateNames: [], manual: false }, + ]; + + it("returns every row when nothing exists yet", () => { + expect(rowsPendingCreation(rows, "anthropic", "my-cred", [])).toHaveLength(2); + }); + + it("skips a row already created under the same credential", () => { + const existing: DeploymentInfoRow[] = [ + { + model_name: "opus", + litellm_params: { model: "anthropic/claude-3-opus", litellm_credential_name: "my-cred" }, + model_info: { id: "abc" }, + }, + ]; + const pending = rowsPendingCreation(rows, "anthropic", "my-cred", existing); + expect(pending.map((r) => r.upstreamId)).toEqual(["claude-3-haiku"]); + }); + + it("does not skip a same-model row that belongs to a different credential", () => { + const existing: DeploymentInfoRow[] = [ + { + model_name: "opus", + litellm_params: { model: "anthropic/claude-3-opus", litellm_credential_name: "someone-elses-cred" }, + model_info: { id: "abc" }, + }, + ]; + expect(rowsPendingCreation(rows, "anthropic", "my-cred", existing)).toHaveLength(2); + }); +}); + +describe("mergeModelGroupAliases", () => { + it("adds new aliases to an empty map", () => { + const { merged, collisions } = mergeModelGroupAliases({}, [{ alias: "gpt-4o", targetModelGroup: "opus" }]); + expect(merged).toEqual({ "gpt-4o": "opus" }); + expect(collisions).toEqual([]); + }); + + it("preserves an existing object-valued entry untouched", () => { + const existing = { "hidden-alias": { model: "some-model", hidden: true } }; + const { merged } = mergeModelGroupAliases(existing, [{ alias: "new-alias", targetModelGroup: "opus" }]); + expect(merged["hidden-alias"]).toEqual({ model: "some-model", hidden: true }); + expect(merged["new-alias"]).toBe("opus"); + }); + + it("rejects a collision with an existing alias rather than overwriting it", () => { + const existing = { "gpt-4o": "some-other-model" }; + const { merged, collisions } = mergeModelGroupAliases(existing, [{ alias: "gpt-4o", targetModelGroup: "opus" }]); + expect(merged["gpt-4o"]).toBe("some-other-model"); + expect(collisions).toEqual(["gpt-4o"]); + }); + + it("rejects a collision between two additions in the same batch", () => { + const { merged, collisions } = mergeModelGroupAliases({}, [ + { alias: "dup", targetModelGroup: "opus" }, + { alias: "dup", targetModelGroup: "haiku" }, + ]); + expect(merged.dup).toBe("opus"); + expect(collisions).toEqual(["dup"]); + }); +}); + +describe("aliasAdditionsFromRows", () => { + it("flattens each row's alternate names against its model_name", () => { + const rows: DiscoveredModelRow[] = [ + { id: "1", upstreamId: "opus", modelName: "my-opus", enabled: true, alternateNames: ["gpt-4o"], manual: false }, + { id: "2", upstreamId: "haiku", modelName: "my-haiku", enabled: true, alternateNames: [], manual: false }, + ]; + expect(aliasAdditionsFromRows(rows)).toEqual([{ alias: "gpt-4o", targetModelGroup: "my-opus" }]); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/wizardLogic.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/wizardLogic.ts new file mode 100644 index 00000000000..2637168545a --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/add-provider/wizardLogic.ts @@ -0,0 +1,128 @@ +import type { DeploymentInfoRow } from "@/components/networking"; + +export interface DiscoveredModelRow { + id: string; + upstreamId: string; + modelName: string; + enabled: boolean; + alternateNames: string[]; + manual: boolean; +} + +/** + * The litellm_params.model this upstream id resolves to: provider-prefixed, but never + * double-prefixed if the discovered/typed id already carries the provider's own prefix. + */ +export const litellmModelForUpstreamId = (litellmProvider: string, upstreamId: string): string => + upstreamId.startsWith(`${litellmProvider}/`) ? upstreamId : `${litellmProvider}/${upstreamId}`; + +export const buildDiscoveredRows = (upstreamIds: readonly string[]): DiscoveredModelRow[] => + upstreamIds.map((upstreamId, index) => ({ + id: `discovered-${index}-${upstreamId}`, + upstreamId, + modelName: upstreamId, + enabled: true, + alternateNames: [], + manual: false, + })); + +export const buildManualRow = (upstreamId: string): DiscoveredModelRow => ({ + id: `manual-${upstreamId}-${Math.random().toString(36).slice(2)}`, + upstreamId, + modelName: upstreamId, + enabled: true, + alternateNames: [], + manual: true, +}); + +export interface CreationResult { + row: DiscoveredModelRow; + status: "created" | "skipped" | "failed"; + detail?: string; +} + +export interface ModelCreationPayload { + model_name: string; + litellm_params: { model: string; litellm_credential_name: string }; + model_info: Record; + blocked: boolean; +} + +/** blocked is the inverse of the review table's enabled toggle -- a disabled row is still + * created, just paused, so a later enable is a one-click unblock instead of re-discovery. */ +export const buildModelCreationPayload = ( + litellmProvider: string, + credentialName: string, + row: DiscoveredModelRow, +): ModelCreationPayload => ({ + model_name: row.modelName, + litellm_params: { + model: litellmModelForUpstreamId(litellmProvider, row.upstreamId), + litellm_credential_name: credentialName, + }, + model_info: {}, + blocked: !row.enabled, +}); + +const deploymentKey = (credentialName: string, litellmModel: string): string => `${credentialName}::${litellmModel}`; + +/** + * Rows not yet created for this credential, matched against every existing deployment by + * (litellm_credential_name, litellm_params.model) -- the partial-failure-recovery key: re-running + * the wizard after some rows already landed must not attempt to recreate them. + */ +export const rowsPendingCreation = ( + rows: readonly DiscoveredModelRow[], + litellmProvider: string, + credentialName: string, + existingDeployments: readonly DeploymentInfoRow[], +): DiscoveredModelRow[] => { + const existingKeys = new Set( + existingDeployments + .filter((deployment) => deployment.litellm_params.litellm_credential_name === credentialName) + .map((deployment) => deploymentKey(credentialName, deployment.litellm_params.model ?? "")), + ); + return rows.filter( + (row) => !existingKeys.has(deploymentKey(credentialName, litellmModelForUpstreamId(litellmProvider, row.upstreamId))), + ); +}; + +export type ModelGroupAliasValue = string | { model: string; hidden?: boolean }; +export type ModelGroupAliasMap = Record; + +export interface AliasAddition { + alias: string; + targetModelGroup: string; +} + +export interface AliasMergeResult { + merged: ModelGroupAliasMap; + collisions: readonly string[]; +} + +/** + * Merges new alias entries into the existing model_group_alias map for a single + * /config/update write. Never touches an existing key's value (object-valued {model, hidden} + * entries survive untouched) and never overwrites a name collision -- a colliding alias is + * reported back rather than silently dropped or silently replacing what was there. + */ +export const mergeModelGroupAliases = ( + existing: ModelGroupAliasMap, + additions: readonly AliasAddition[], +): AliasMergeResult => { + const collisions: string[] = []; + const added: ModelGroupAliasMap = {}; + for (const { alias, targetModelGroup } of additions) { + if (alias in existing || alias in added) { + collisions.push(alias); + continue; + } + added[alias] = targetModelGroup; + } + return { merged: { ...existing, ...added }, collisions }; +}; + +/** Flattens the review table's per-row alternate names into the alias additions + * mergeModelGroupAliases expects, skipping rows the operator removed/left disabled-but-empty. */ +export const aliasAdditionsFromRows = (rows: readonly DiscoveredModelRow[]): AliasAddition[] => + rows.flatMap((row) => row.alternateNames.map((alias) => ({ alias, targetModelGroup: row.modelName }))); diff --git a/ui/litellm-dashboard/src/components/add_model/provider_credential_variants.test.ts b/ui/litellm-dashboard/src/components/add_model/provider_credential_variants.test.ts new file mode 100644 index 00000000000..64ef7f956bb --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/provider_credential_variants.test.ts @@ -0,0 +1,152 @@ +import { describe, expect, it } from "vitest"; +import { getVariant, inferActiveVariant, resolveVariantFieldDefs } from "./provider_credential_variants"; +import type { ProviderCredentialVariants } from "../networking"; + +const field = (key: string, required = false): ProviderCredentialVariants["field_definitions"][number] => ({ + key, + label: key, + required, + field_type: "text", +}); + +const anthropicVariants: ProviderCredentialVariants = { + selector_label: "Authentication method", + default_variant: "api_key", + field_definitions: [ + field("api_base"), + field("api_key"), + field("anthropic_federation_rule_id", true), + field("anthropic_organization_id", true), + field("anthropic_identity_token", true), + field("anthropic_identity_token_file", true), + field("anthropic_issuer_url", true), + field("anthropic_issuer_signing_key_ref", true), + field("anthropic_keycloak_token_url", true), + field("anthropic_keycloak_client_id", true), + field("anthropic_keycloak_client_secret_ref", true), + ], + variants: [ + { id: "api_key", label: "API Key", field_keys: ["api_base", "api_key"], fixed_values: {} }, + { + id: "wif_token", + label: "WIF (token)", + field_keys: ["anthropic_federation_rule_id", "anthropic_organization_id", "anthropic_identity_token"], + fixed_values: {}, + }, + { + id: "wif_token_file", + label: "WIF (token file)", + field_keys: ["anthropic_federation_rule_id", "anthropic_organization_id", "anthropic_identity_token_file"], + fixed_values: {}, + }, + { + id: "wif_internal_issuer", + label: "WIF (internal issuer)", + field_keys: ["anthropic_federation_rule_id", "anthropic_organization_id", "anthropic_issuer_url", "anthropic_issuer_signing_key_ref"], + fixed_values: { anthropic_identity_source: "internal_issuer" }, + }, + { + id: "wif_keycloak", + label: "WIF (keycloak)", + field_keys: [ + "anthropic_federation_rule_id", + "anthropic_organization_id", + "anthropic_keycloak_token_url", + "anthropic_keycloak_client_id", + "anthropic_keycloak_client_secret_ref", + ], + fixed_values: { anthropic_identity_source: "keycloak" }, + }, + ], +}; + +describe("getVariant", () => { + it("finds a variant by id", () => { + expect(getVariant(anthropicVariants, "wif_keycloak")?.label).toBe("WIF (keycloak)"); + }); + + it("returns undefined for an unknown id", () => { + expect(getVariant(anthropicVariants, "nope")).toBeUndefined(); + }); +}); + +describe("resolveVariantFieldDefs", () => { + it("resolves field_keys to their field_definitions, in order", () => { + const fields = resolveVariantFieldDefs(anthropicVariants, "wif_token_file"); + expect(fields.map((f) => f.key)).toEqual([ + "anthropic_federation_rule_id", + "anthropic_organization_id", + "anthropic_identity_token_file", + ]); + }); + + it("returns an empty list for an unknown variant id", () => { + expect(resolveVariantFieldDefs(anthropicVariants, "nope")).toEqual([]); + }); + + it("drops a field_key that has no matching field_definitions entry, rather than crashing", () => { + const variants: ProviderCredentialVariants = { + ...anthropicVariants, + variants: [{ id: "broken", label: "Broken", field_keys: ["api_key", "missing_field"], fixed_values: {} }], + }; + expect(resolveVariantFieldDefs(variants, "broken").map((f) => f.key)).toEqual(["api_key"]); + }); +}); + +describe("inferActiveVariant", () => { + it("defaults to default_variant on a blank form", () => { + expect(inferActiveVariant(anthropicVariants, {})).toBe("api_key"); + }); + + it("picks wif_token when only the inline identity token is set", () => { + const values = { + anthropic_federation_rule_id: "rule-1", + anthropic_organization_id: "org-1", + anthropic_identity_token: "oidc/env/TOKEN", + }; + expect(inferActiveVariant(anthropicVariants, values)).toBe("wif_token"); + }); + + it("picks wif_token_file over wif_token when the file field is set instead", () => { + const values = { + anthropic_federation_rule_id: "rule-1", + anthropic_organization_id: "org-1", + anthropic_identity_token_file: "/var/run/secrets/tokens/oidc-token", + }; + expect(inferActiveVariant(anthropicVariants, values)).toBe("wif_token_file"); + }); + + it("picks wif_internal_issuer from its fixed discriminator, not just field presence", () => { + const values = { + anthropic_identity_source: "internal_issuer", + anthropic_federation_rule_id: "rule-1", + anthropic_organization_id: "org-1", + anthropic_issuer_url: "https://issuer.example.com", + anthropic_issuer_signing_key_ref: "os.environ/SIGNING_KEY", + }; + expect(inferActiveVariant(anthropicVariants, values)).toBe("wif_internal_issuer"); + }); + + it("picks wif_keycloak from its fixed discriminator", () => { + const values = { + anthropic_identity_source: "keycloak", + anthropic_federation_rule_id: "rule-1", + anthropic_organization_id: "org-1", + anthropic_keycloak_token_url: "https://keycloak.example.com/token", + anthropic_keycloak_client_id: "client-1", + anthropic_keycloak_client_secret_ref: "os.environ/SECRET", + }; + expect(inferActiveVariant(anthropicVariants, values)).toBe("wif_keycloak"); + }); + + it("does not match a fixed-values variant whose required fields are still incomplete", () => { + // The discriminator alone (e.g. left over from a variant switch) is not enough. + const values = { anthropic_identity_source: "keycloak" }; + expect(inferActiveVariant(anthropicVariants, values)).toBe("api_key"); + }); + + it("ignores non-required optional fields when deciding a match", () => { + const values = { api_base: "https://api.anthropic.com" }; + expect(inferActiveVariant(anthropicVariants, values)).toBe("api_key"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/provider_credential_variants.ts b/ui/litellm-dashboard/src/components/add_model/provider_credential_variants.ts new file mode 100644 index 00000000000..80b7622a685 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/provider_credential_variants.ts @@ -0,0 +1,49 @@ +import type { ProviderCredentialFieldMetadata, ProviderCredentialVariant, ProviderCredentialVariants } from "../networking"; + +export const getVariant = (variants: ProviderCredentialVariants, variantId: string): ProviderCredentialVariant | undefined => + variants.variants.find((variant) => variant.id === variantId); + +export const resolveVariantFieldDefs = ( + variants: ProviderCredentialVariants, + variantId: string, +): ProviderCredentialFieldMetadata[] => { + const variant = getVariant(variants, variantId); + if (!variant) { + return []; + } + const byKey = new Map(variants.field_definitions.map((field) => [field.key, field])); + return variant.field_keys + .map((key) => byKey.get(key)) + .filter((field): field is ProviderCredentialFieldMetadata => field !== undefined); +}; + +const hasValue = (value: unknown): boolean => (typeof value === "string" ? value.trim() !== "" : value != null); + +const isFullySatisfied = ( + variant: ProviderCredentialVariant, + fieldsByKey: Map, + values: Record, +): boolean => { + const fixedValuesMatch = Object.entries(variant.fixed_values).every(([key, value]) => values[key] === value); + if (!fixedValuesMatch) { + return false; + } + return variant.field_keys.every((key) => !fieldsByKey.get(key)?.required || hasValue(values[key])); +}; + +/** + * Infers which variant a set of existing field values belongs to, for pre-selecting the + * selector when editing a saved credential. Checks variants in declaration order and keeps + * the last one whose fixed_values match and whose required fields are all present, so a more + * specific variant (one with a matching discriminator, or more required fields set) wins over + * a broader one earlier in the list (e.g. a bare api_key variant with no required fields, + * which always matches trivially). Falls back to default_variant when nothing matches, which + * is also what a blank/new form resolves to. + */ +export const inferActiveVariant = (variants: ProviderCredentialVariants, values: Record): string => { + const fieldsByKey = new Map(variants.field_definitions.map((field) => [field.key, field])); + return variants.variants.reduce( + (best, variant) => (isFullySatisfied(variant, fieldsByKey, values) ? variant.id : best), + variants.default_variant, + ); +}; diff --git a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.test.tsx b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.test.tsx index 1ea9505201c..eada05d4f04 100644 --- a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.test.tsx @@ -1,5 +1,6 @@ import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { fireEvent, render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { beforeAll, describe, expect, it, vi } from "vitest"; import { useFormContext } from "react-hook-form"; import { Providers } from "../provider_info_helpers"; @@ -109,6 +110,65 @@ vi.mock("../networking", async () => { }, ], }, + { + provider: "Anthropic", + provider_display_name: Providers.Anthropic, + litellm_provider: "anthropic", + default_model_placeholder: "claude-3-opus", + credential_fields: [ + { key: "api_base", label: "Upstream API Base", field_type: "text" }, + { key: "api_key", label: "API Key", field_type: "password" }, + ], + credential_variants: { + selector_label: "Authentication method", + default_variant: "api_key", + field_definitions: [ + { key: "api_base", label: "Upstream API Base", field_type: "text" }, + { key: "api_key", label: "API Key", field_type: "password" }, + { key: "anthropic_federation_rule_id", label: "Federation Rule ID", field_type: "text", required: true }, + { key: "anthropic_organization_id", label: "Organization ID", field_type: "text", required: true }, + { + key: "anthropic_identity_token", + label: "Identity Token Reference", + field_type: "text", + required: true, + }, + { + key: "anthropic_issuer_url", + label: "Issuer URL", + field_type: "text", + required: true, + }, + { + key: "anthropic_issuer_signing_key_ref", + label: "Signing Key Reference", + field_type: "text", + required: true, + tooltip: "A secret REFERENCE, e.g. os.environ/VAR_NAME. Never the key itself.", + }, + ], + variants: [ + { id: "api_key", label: "API Key", field_keys: ["api_base", "api_key"], fixed_values: {} }, + { + id: "wif_token", + label: "Workload Identity Federation (external token)", + field_keys: ["anthropic_federation_rule_id", "anthropic_organization_id", "anthropic_identity_token"], + fixed_values: {}, + }, + { + id: "wif_internal_issuer", + label: "Workload Identity Federation (LiteLLM-signed)", + field_keys: [ + "anthropic_federation_rule_id", + "anthropic_organization_id", + "anthropic_issuer_url", + "anthropic_issuer_signing_key_ref", + ], + fixed_values: { anthropic_identity_source: "internal_issuer" }, + }, + ], + }, + }, ]), }; }); @@ -419,4 +479,112 @@ describe("ProviderSpecificFields", () => { expect(apiVersionInput).toHaveValue("2025-01-01-preview"); }); }); + + describe("credential_variants", () => { + const IdentitySourceProbe = () => { + const { watch } = useFormContext(); + return {String(watch("anthropic_identity_source") ?? "")}; + }; + + it("defaults to the api_key variant and hides WIF fields", async () => { + const queryClient = createQueryClient(); + render( + + + + + , + ); + + expect(await screen.findByLabelText("API Key")).toBeInTheDocument(); + expect(screen.getByLabelText("Upstream API Base")).toBeInTheDocument(); + expect(screen.queryByLabelText("Federation Rule ID")).not.toBeInTheDocument(); + expect(screen.queryByLabelText("Identity Token Reference")).not.toBeInTheDocument(); + }); + + it("switching to a WIF variant swaps the rendered fields and unmounts the previous variant's", async () => { + const queryClient = createQueryClient(); + render( + + + + + , + ); + + const user = userEvent.setup(); + await screen.findByLabelText("API Key"); + await user.click(await screen.findByRole("combobox", { name: "Authentication method" })); + await user.click(await screen.findByRole("option", { name: "Workload Identity Federation (external token)" })); + + expect(await screen.findByLabelText("Federation Rule ID")).toBeInTheDocument(); + expect(screen.getByLabelText("Identity Token Reference")).toBeInTheDocument(); + // api_key/api_base belong only to the api_key variant here, so they must be gone. + expect(screen.queryByLabelText("API Key")).not.toBeInTheDocument(); + expect(screen.queryByLabelText("Upstream API Base")).not.toBeInTheDocument(); + }); + + it("injects the fixed discriminator for a variant without rendering a field for it", async () => { + const queryClient = createQueryClient(); + render( + + + + + + , + ); + + const user = userEvent.setup(); + await screen.findByLabelText("API Key"); + expect(screen.getByTestId("identity-source")).toBeEmptyDOMElement(); + + await user.click(await screen.findByRole("combobox", { name: "Authentication method" })); + await user.click(await screen.findByRole("option", { name: "Workload Identity Federation (LiteLLM-signed)" })); + + expect(await screen.findByLabelText("Issuer URL")).toBeInTheDocument(); + expect(screen.queryByLabelText("anthropic_identity_source")).not.toBeInTheDocument(); + await waitFor(() => expect(screen.getByTestId("identity-source")).toHaveTextContent("internal_issuer")); + }); + + it("labels a secret-reference field as a reference, not a raw secret input", async () => { + const queryClient = createQueryClient(); + render( + + + + + , + ); + + const user = userEvent.setup(); + await screen.findByLabelText("API Key"); + await user.click(await screen.findByRole("combobox", { name: "Authentication method" })); + await user.click(await screen.findByRole("option", { name: "Workload Identity Federation (LiteLLM-signed)" })); + + const signingKeyInput = await screen.findByLabelText("Signing Key Reference"); + // A *_ref field renders as plain text, never password -- it must never invite a pasted secret. + expect(signingKeyInput).toHaveAttribute("type", "text"); + }); + + it("infers the wif_token variant is already active when editing a credential that has one", async () => { + const queryClient = createQueryClient(); + render( + + + + + , + ); + + expect(await screen.findByLabelText("Identity Token Reference")).toBeInTheDocument(); + expect(screen.queryByLabelText("API Key")).not.toBeInTheDocument(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx index e1a47e9792d..bb6ecd52bac 100644 --- a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx +++ b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx @@ -10,12 +10,15 @@ import { useFormContext } from "react-hook-form"; import { requiredRule } from "../common_components/formRules"; import { MountedFormField, + useMountedName, type MountedFieldControlProps, type MountedFormValues, } from "../common_components/MountedFormField"; -import { CredentialItem, ProviderCredentialFieldMetadata } from "../networking"; +import { CredentialItem, ProviderCredentialFieldMetadata, ProviderCredentialVariants } from "../networking"; import { provider_map, Providers } from "../provider_info_helpers"; import { labelWithHint } from "@/components/shared/form/LabelWithHint"; +import { Field, FieldLabel } from "@/components/shared/form/field"; +import { getVariant, inferActiveVariant, resolveVariantFieldDefs } from "./provider_credential_variants"; interface ProviderSpecificFieldsProps { selectedProvider: Providers; @@ -88,6 +91,13 @@ const mapFieldMetadataToUiField = (field: ProviderCredentialFieldMetadata): Prov // non-React helpers like createCredentialFromModel. const providerFieldsByDisplayName: Record = {}; +// Companion cache for credential_variants, keyed and populated the same way. Without this, +// `allFields`'s cache-first lookup can resolve before providerMetadata (and thus `variants`) +// has loaded, briefly rendering a variant-capable provider through the flat legacy field list +// -- which registers the wrong `required` rule with react-hook-form for a field name a variant +// later reuses (e.g. api_key), and that stale rule outlives the field's unmount. +const providerVariantsByDisplayName: Record = {}; + export const createCredentialFromModel = (provider: string, modelData: any): CredentialItem => { const enumKey = Object.keys(provider_map).find((key) => provider_map[key].toLowerCase() === provider.toLowerCase()); if (!enumKey) { @@ -117,6 +127,20 @@ export const createCredentialFromModel = (provider: string, modelData: any): Cre return credential; }; +// Mounts a litellm_params value a credential_variants variant implies (e.g. the +// anthropic_identity_source discriminator) without a visible form field for it, so it is +// submitted alongside whatever the operator actually typed. Registering via useMountedName +// (rather than rendering a MountedFormField) keeps this a real function component, so the +// effect below follows the ordinary rules of hooks instead of running inside a render-prop. +const FixedValueField: React.FC<{ name: string; value: string }> = ({ name, value }) => { + const form = useFormContext(); + useMountedName(name); + React.useEffect(() => { + form.setValue(name, value); + }, [form, name, value]); + return null; +}; + const ProviderSpecificFields: React.FC = ({ selectedProvider }) => { const selectedProviderEnum = Providers[selectedProvider as keyof typeof Providers] as Providers; const form = useFormContext(); @@ -139,45 +163,50 @@ const ProviderSpecificFields: React.FC = ({ selecte } // Compute cache entries keyed by provider display name and identifiers - const entries: Record = {}; + const fieldEntries: Record = {}; + const variantEntries: Record = {}; providerMetadata.forEach((providerInfo) => { - const displayName = providerInfo.provider_display_name; const mappedFields = providerInfo.credential_fields.map(mapFieldMetadataToUiField); - - // Primary key: human-readable display name - entries[displayName] = mappedFields; - - // Also cache by backend identifiers so lookups by provider slug work - if (providerInfo.provider) { - entries[providerInfo.provider] = mappedFields; - } - if (providerInfo.litellm_provider) { - entries[providerInfo.litellm_provider] = mappedFields; - } + const keys = [providerInfo.provider_display_name, providerInfo.provider, providerInfo.litellm_provider].filter( + (key): key is string => Boolean(key), + ); + keys.forEach((key) => { + fieldEntries[key] = mappedFields; + variantEntries[key] = providerInfo.credential_variants ?? null; + }); }); - return entries; + return { fieldEntries, variantEntries }; }, [providerMetadata]); - // Sync memoized cache entries to module-level cache + // Sync memoized cache entries to the module-level caches together, so a lookup can never see + // one updated and the other still stale. React.useEffect(() => { if (!cacheEntries) { return; } - Object.assign(providerFieldsByDisplayName, cacheEntries); + Object.assign(providerFieldsByDisplayName, cacheEntries.fieldEntries); + Object.assign(providerVariantsByDisplayName, cacheEntries.variantEntries); }, [cacheEntries]); - const allFields = React.useMemo(() => { - // First try to resolve from the in-memory cache. We support both the - // enum/display-name form and the raw provider slug (e.g. "petals"). + // Resolves this provider's fields and variants together -- from the module-level cache when + // an earlier mount (of any provider) has already populated it, else from providerMetadata once + // loaded. Combined into one lookup so the two can never disagree about whether this provider + // has credential_variants on a given render (see the cache comment above for why that matters). + const { allFields, variants } = React.useMemo((): { + allFields: ProviderCredentialField[]; + variants: ProviderCredentialVariants | null; + } => { const cachedFields = providerFieldsByDisplayName[selectedProviderEnum] ?? providerFieldsByDisplayName[selectedProvider]; if (cachedFields) { - return cachedFields; + const cachedVariants = + providerVariantsByDisplayName[selectedProviderEnum] ?? providerVariantsByDisplayName[selectedProvider]; + return { allFields: cachedFields, variants: cachedVariants ?? null }; } if (!providerMetadata) { - return []; + return { allFields: [], variants: null }; } const providerInfo = providerMetadata.find( @@ -187,21 +216,42 @@ const ProviderSpecificFields: React.FC = ({ selecte p.litellm_provider === selectedProvider, ); if (!providerInfo) { - return []; + return { allFields: [], variants: null }; } const mapped = providerInfo.credential_fields.map(mapFieldMetadataToUiField); - providerFieldsByDisplayName[providerInfo.provider_display_name] = mapped; - if (providerInfo.provider) { - providerFieldsByDisplayName[providerInfo.provider] = mapped; - } - if (providerInfo.litellm_provider) { - providerFieldsByDisplayName[providerInfo.litellm_provider] = mapped; - } - return mapped; + const resolvedVariants = providerInfo.credential_variants ?? null; + [providerInfo.provider_display_name, providerInfo.provider, providerInfo.litellm_provider] + .filter((key): key is string => Boolean(key)) + .forEach((key) => { + providerFieldsByDisplayName[key] = mapped; + providerVariantsByDisplayName[key] = resolvedVariants; + }); + return { allFields: mapped, variants: resolvedVariants }; }, [selectedProviderEnum, selectedProvider, providerMetadata]); - const hasApiVersionField = React.useMemo(() => allFields.some((field) => field.key === "api_version"), [allFields]); + // The selector's own choice, once the operator makes one; otherwise undefined so the variant + // stays derived (inferred fresh every render from `variants`/the current field values, so a + // still-loading `variants` or a provider switch resolve correctly with no effect needed to + // "catch up" afterward). A choice left over from a different provider's variant list is + // treated as unset rather than resolving to a variant id that provider doesn't have. + const [userChosenVariantId, setUserChosenVariantId] = React.useState(undefined); + const validUserChoice = variants?.variants.some((variant) => variant.id === userChosenVariantId) + ? userChosenVariantId + : undefined; + const activeVariantId = validUserChoice ?? (variants ? inferActiveVariant(variants, form.getValues()) : ""); + + const activeVariantFields = React.useMemo( + () => (variants ? resolveVariantFieldDefs(variants, activeVariantId).map(mapFieldMetadataToUiField) : []), + [variants, activeVariantId], + ); + + const currentFields = variants ? activeVariantFields : allFields; + + const hasApiVersionField = React.useMemo( + () => currentFields.some((field) => field.key === "api_version"), + [currentFields], + ); const lastInferredApiVersionRef = React.useRef(null); const handleApiBaseChange = React.useCallback( @@ -313,49 +363,83 @@ const ProviderSpecificFields: React.FC = ({ selecte ); }; + const renderFieldEntry = (field: ProviderCredentialField) => ( + + + {(control) => renderFieldControl(field, control)} + + + {/* Special case for Vertex Credentials help text */} + {field.key === "vertex_credentials" && ( +

Give a gcp service account(.json file)

+ )} + + {/* Special case for Azure Base Model help text */} + {field.key === "base_model" && ( +
+

+ The actual model your azure deployment uses. Used for accurate cost tracking. Select name from{" "} + + here + +

+
+ )} +
+ ); + + const activeVariant = variants ? getVariant(variants, activeVariantId) : undefined; + return ( <> - {isLoading && allFields.length === 0 &&

Loading provider fields...

} - {loadError && allFields.length === 0 && ( + {isLoading && currentFields.length === 0 &&

Loading provider fields...

} + {loadError && currentFields.length === 0 && (

{loadError instanceof Error ? loadError.message : "Failed to load provider credential fields"}

)} - {allFields.map((field) => ( - - + {variants.selector_label} + + + )} + {/* Keyed by the active variant so switching variants fully unmounts the previous + variant's fields (deregistering them from submission) and re-applies fixed_values + fresh, rather than leaving a stale value behind under a reused field name. */} + + {currentFields.map(renderFieldEntry)} + {activeVariant && + Object.entries(activeVariant.fixed_values).map(([key, value]) => ( + + ))} + ); }; diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx index 8d94ae0d213..080c5e86d61 100644 --- a/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx +++ b/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx @@ -1,5 +1,6 @@ import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import { render, screen, waitFor } from "@testing-library/react"; +import { fireEvent, render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { describe, expect, it, vi } from "vitest"; import { Providers } from "../provider_info_helpers"; import { CredentialItem } from "../networking"; @@ -43,6 +44,30 @@ vi.mock("../networking", async () => { required: true, }, ], + credential_variants: { + selector_label: "Authentication method", + default_variant: "api_key", + field_definitions: [ + { key: "api_key", label: "Anthropic API Key", field_type: "password" }, + { key: "anthropic_federation_rule_id", label: "Federation Rule ID", field_type: "text", required: true }, + { key: "anthropic_organization_id", label: "Organization ID", field_type: "text", required: true }, + { + key: "anthropic_identity_token", + label: "Identity Token Reference", + field_type: "text", + required: true, + }, + ], + variants: [ + { id: "api_key", label: "API Key", field_keys: ["api_key"], fixed_values: {} }, + { + id: "wif_token", + label: "Workload Identity Federation", + field_keys: ["anthropic_federation_rule_id", "anthropic_organization_id", "anthropic_identity_token"], + fixed_values: {}, + }, + ], + }, }, ]), }; @@ -125,4 +150,51 @@ describe("CredentialModal", () => { expect(screen.getByLabelText("Credential Name:")).toBeDisabled(); }); }); + + describe("credential_values_to_delete", () => { + const anthropicWifCredential: CredentialItem = { + credential_name: "anthropic-wif-cred", + credential_values: { + anthropic_federation_rule_id: "rule-1", + anthropic_organization_id: "org-1", + anthropic_identity_token: "oidc/env/TOKEN", + }, + credential_info: { + custom_llm_provider: Providers.Anthropic, + }, + }; + + it("flags the previous variant's fields for deletion when the operator switches variants", async () => { + const onSubmit = vi.fn(); + const user = userEvent.setup(); + renderModal({ mode: "edit", existingCredential: anthropicWifCredential, onSubmit }); + + await screen.findByLabelText("Federation Rule ID"); + await user.click(await screen.findByRole("combobox", { name: "Authentication method" })); + await user.click(await screen.findByRole("option", { name: "API Key" })); + await screen.findByLabelText("Anthropic API Key"); + + await user.click(screen.getByRole("button", { name: "Update Credential" })); + + await waitFor(() => expect(onSubmit).toHaveBeenCalled()); + const [, deletedKeys] = onSubmit.mock.calls[0]; + expect([...deletedKeys].sort()).toEqual([ + "anthropic_federation_rule_id", + "anthropic_identity_token", + "anthropic_organization_id", + ]); + }); + + it("flags nothing for deletion when the variant and values are untouched", async () => { + const onSubmit = vi.fn(); + renderModal({ mode: "edit", existingCredential: anthropicWifCredential, onSubmit }); + + await screen.findByLabelText("Federation Rule ID"); + fireEvent.click(screen.getByRole("button", { name: "Update Credential" })); + + await waitFor(() => expect(onSubmit).toHaveBeenCalled()); + const [, deletedKeys] = onSubmit.mock.calls[0]; + expect(deletedKeys).toEqual([]); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx index 8108cc1c20e..b9a976c222b 100644 --- a/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx +++ b/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx @@ -17,7 +17,7 @@ import { import { CredentialItem } from "../networking"; import { Providers } from "../provider_info_helpers"; import { Logo } from "@/components/molecules/logo/Logo"; -import { resetCredentialFormOnProviderChange } from "./credential_form_helpers"; +import { computeCredentialValuesToDelete, resetCredentialFormOnProviderChange } from "./credential_form_helpers"; import { Dialog, DialogContent, DialogHeader, DialogTitle } from "@/components/ui/dialog"; const providerOptions: SearchSelectOption[] = Object.entries(Providers).map(([providerEnum, providerDisplayName]) => ({ @@ -29,7 +29,7 @@ const providerOptions: SearchSelectOption[] = Object.entries(Providers).map(([pr interface CredentialModalProps { open: boolean; onCancel: () => void; - onSubmit: (values: any) => void; + onSubmit: (values: any, credentialValuesToDelete: string[]) => void; mode: "add" | "edit"; existingCredential?: CredentialItem | null; } @@ -77,7 +77,10 @@ export default function CredentialModal({ } return acc; }, {} as any); - onSubmit(filteredValues); + const credentialValuesToDelete = isEdit + ? computeCredentialValuesToDelete(existingCredential?.credential_values ?? {}, values) + : []; + onSubmit(filteredValues, credentialValuesToDelete); form.reset(); }; diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialsPanel.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialsPanel.tsx index 80a8216dbaf..cc9e94772ab 100644 --- a/ui/litellm-dashboard/src/components/model_add/CredentialsPanel.tsx +++ b/ui/litellm-dashboard/src/components/model_add/CredentialsPanel.tsx @@ -47,12 +47,15 @@ export default function CredentialsPanel() { const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); const [isCredentialDeleting, setIsCredentialDeleting] = useState(false); - const handleUpdateCredential = async (values: Record) => { + const handleUpdateCredential = async (values: Record, credentialValuesToDelete: string[] = []) => { if (!accessToken) { return; } try { - const newCredential = buildCredential(values, stripMaskedSecrets(withoutRestrictedFields(values))); + const newCredential = { + ...buildCredential(values, stripMaskedSecrets(withoutRestrictedFields(values))), + ...(credentialValuesToDelete.length > 0 ? { credential_values_to_delete: credentialValuesToDelete } : {}), + }; await credentialUpdateCall(accessToken, values.credential_name as string, newCredential); toast.success("Credential updated successfully"); setIsUpdateModalOpen(false); diff --git a/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.test.ts b/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.test.ts index d7218f6f4d3..759e4a43ccb 100644 --- a/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.test.ts +++ b/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.test.ts @@ -1,6 +1,6 @@ import { describe, expect, it, vi } from "vitest"; import { Providers } from "../provider_info_helpers"; -import { resetCredentialFormOnProviderChange } from "./credential_form_helpers"; +import { computeCredentialValuesToDelete, resetCredentialFormOnProviderChange } from "./credential_form_helpers"; /** * Build a minimal FormInstance stub that records calls. We don't depend @@ -81,3 +81,31 @@ describe("resetCredentialFormOnProviderChange", () => { expect(credentialNameCalls).toHaveLength(0); }); }); + +describe("computeCredentialValuesToDelete", () => { + it("flags a field that is no longer mounted at all", () => { + // e.g. switching from the api_key variant to a WIF variant unmounts api_key/api_base. + const original = { api_base: "https://api.anthropic.com", api_key: "sk-***1234" }; + const mounted = { anthropic_federation_rule_id: "rule-1" }; + + expect(computeCredentialValuesToDelete(original, mounted)).toEqual(["api_base", "api_key"]); + }); + + it("keeps a masked-but-untouched field, since it is still mounted", () => { + const original = { api_key: "sk-***1234" }; + const mounted = { api_key: "sk-***1234" }; + + expect(computeCredentialValuesToDelete(original, mounted)).toEqual([]); + }); + + it("keeps a field the operator genuinely changed", () => { + const original = { api_key: "sk-***1234" }; + const mounted = { api_key: "sk-new-real-key" }; + + expect(computeCredentialValuesToDelete(original, mounted)).toEqual([]); + }); + + it("returns nothing when nothing existed before", () => { + expect(computeCredentialValuesToDelete({}, { api_key: "sk-new" })).toEqual([]); + }); +}); diff --git a/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.ts b/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.ts index 7190c539c81..85d0d913db6 100644 --- a/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.ts +++ b/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.ts @@ -36,3 +36,22 @@ export function resetCredentialFormOnProviderChange( setSelectedProvider(newProvider); form.setFieldValue("custom_llm_provider", newProvider); } + +/** + * Keys to drop from a saved credential's `credential_values` on update: whatever the form had + * mounted before that it does not have mounted now. A field stops being mounted either because + * the operator cleared it or because a credential_variants switch (e.g. api_key -> Keycloak) + * unmounted it, and either way the backend must actually delete it rather than merge over it + * -- a leftover field from a different auth variant fails validation on the next request + * (wif.py rejects foreign-variant fields by presence). + * + * `mountedValues` must be the full projected form state (masked-but-untouched fields included), + * not the caller's post-filter payload: a masked value that the operator never touched is still + * mounted and must be preserved, not read as "absent, so delete it". + */ +export function computeCredentialValuesToDelete( + originalValues: Record, + mountedValues: Record, +): string[] { + return Object.keys(originalValues).filter((key) => !(key in mountedValues)); +} diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 5b6d70b4771..979f761c8ce 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -197,6 +197,8 @@ export interface Model { model_name: string; litellm_params: object; model_info: object | null; + // Admin-toggled pause flag, mirrors LiteLLM_ProxyModelTable.blocked. Proxy-admin only. + blocked?: boolean; } interface PromptInfo { @@ -258,6 +260,8 @@ export interface CredentialItem { description?: string; required?: boolean; }; + // PATCH-only: keys to drop from the stored credential_values. + credential_values_to_delete?: string[] | null; } export interface ProviderCredentialFieldMetadata { @@ -271,12 +275,27 @@ export interface ProviderCredentialFieldMetadata { default_value?: string | null; } +export interface ProviderCredentialVariant { + id: string; + label: string; + field_keys: string[]; + fixed_values: Record; +} + +export interface ProviderCredentialVariants { + selector_label: string; + default_variant: string; + field_definitions: ProviderCredentialFieldMetadata[]; + variants: ProviderCredentialVariant[]; +} + export interface ProviderCreateInfo { provider: string; provider_display_name: string; litellm_provider: string; default_model_placeholder?: string | null; credential_fields: ProviderCredentialFieldMetadata[]; + credential_variants?: ProviderCredentialVariants | null; } export interface AgentCredentialFieldMetadata { @@ -372,6 +391,39 @@ export const getProviderCreateMetadata = async (): Promise return jsonData; }; +export interface ProviderModelDiscoveryRequest { + custom_llm_provider: string; + litellm_credential_name?: string; + api_key?: string; + api_base?: string; +} + +export interface ProviderModelDiscoveryResponse { + models: string[]; +} + +export const discoverProviderModelsCall = async ( + accessToken: string, + body: ProviderModelDiscoveryRequest, +): Promise => { + /** + * Live model discovery for a configured provider credential. Proxy-admin only. + */ + return await apiClient.post(`/provider/models/discover`, { accessToken, body }); +}; + +export interface AnthropicJwks { + keys: Record[]; +} + +export const getCredentialJwksCall = async (accessToken: string, credentialName: string): Promise => { + /** + * Export the public JWKS for an anthropic internal_issuer credential, so it can be + * registered on the Anthropic federation issuer. + */ + return await apiClient.get(`/credentials/${encodeURIComponent(credentialName)}/jwks`, { accessToken }); +}; + export interface ComplexityScorerDefaults { tier_boundaries: Record; token_thresholds: Record; @@ -621,6 +673,14 @@ export const modelCreateCall = async (accessToken: string, formValues: Model) => } }; +export const createProviderModelCall = async (accessToken: string, formValues: Model): Promise<{ model_id?: string }> => { + /** + * Same endpoint as modelCreateCall, without the per-call success toast -- used by the Add + * Provider wizard, which creates many rows at once and reports its own bulk-creation summary. + */ + return await apiClient.post(`/model/new`, { accessToken, body: formValues }); +}; + export const modelDeleteCall = async (accessToken: string, model_id: string) => { try { const data = await apiClient.post(`/model/delete`, { @@ -1755,6 +1815,20 @@ export const modelInfoV1Call = async (accessToken: string, modelId: string) => { } }; +export interface DeploymentInfoRow { + model_name: string; + litellm_params: { model?: string; litellm_credential_name?: string }; + model_info: { id: string }; +} + +export const listAllModelsCall = async (accessToken: string): Promise<{ data: DeploymentInfoRow[] }> => { + /** + * Every deployment the caller can see, unpaginated. Used to dedupe model creation against + * what already exists (e.g. when resuming a partially-failed Add Provider wizard run). + */ + return await apiClient.get(`/model/info`, { accessToken }); +}; + export const modelHubPublicModelsCall = async () => { const url = proxyBaseUrl ? `${proxyBaseUrl}/public/model_hub` : `/public/model_hub`; const response = await fetch(url, {