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
This commit is contained in:
derhornspieler 2026-08-23 04:53:49 -04:00
parent 66a37443c3
commit cb3f29d912
33 changed files with 3743 additions and 406 deletions

View file

@ -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:
"""

View file

@ -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)

View file

@ -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:

View file

@ -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):

View file

@ -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,
)

View file

@ -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)

View file

@ -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,

View file

@ -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/<audience>, or oidc/google/<audience>. 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",

View file

@ -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]

View file

@ -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):

View file

@ -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=())

View file

@ -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)

View file

@ -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 = "<script>evil()</script>" * 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

View file

@ -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

View file

@ -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

View file

@ -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<ModelTabSlug, string> = {
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 <AutoRoutersTabPanel />;
case "add":
return <AddModelPanel />;
case "add-provider":
return <AddProviderPanel />;
case "llm-credentials":
return <LlmCredentialsPanel />;
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)

View file

@ -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<typeof import("@/components/networking")>();
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(<AddProviderPanel />);
await screen.findByLabelText("Provider");
return { user };
};
const chooseProvider = async (user: ReturnType<typeof userEvent.setup>, 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();
});
});

View file

@ -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<WizardStep, string> = {
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 (
<div className="mb-6 flex flex-wrap items-center gap-2 text-sm">
{visibleSteps.map((s, index) => (
<React.Fragment key={s}>
{index > 0 && <span className="text-muted-foreground">{"->"}</span>}
<span className={index <= currentIndex ? "font-medium text-foreground" : "text-muted-foreground"}>
{STEP_LABELS[s]}
</span>
</React.Fragment>
))}
</div>
);
};
export default function AddProviderPanel() {
const { accessToken } = useAuthorized();
const queryClient = useQueryClient();
const { data: providerMetadata } = useProviderFields();
const { data: credentialsResponse } = useCredentials();
const [step, setStep] = React.useState<WizardStep>("provider");
const [selectedProvider, setSelectedProvider] = React.useState<Providers | null>(null);
const [credentialName, setCredentialName] = React.useState("");
const [credentialSaved, setCredentialSaved] = React.useState(false);
const [savedValues, setSavedValues] = React.useState<Record<string, unknown>>({});
const [federationRuleId, setFederationRuleId] = React.useState("");
const [jwks, setJwks] = React.useState<AnthropicJwks | null>(null);
const [jwksError, setJwksError] = React.useState<string | null>(null);
const [discoveryError, setDiscoveryError] = React.useState<string | null>(null);
const [isDiscovering, setIsDiscovering] = React.useState(false);
const [rows, setRows] = React.useState<DiscoveredModelRow[]>([]);
const [isCreating, setIsCreating] = React.useState(false);
const [creationResults, setCreationResults] = React.useState<CreationResult[]>([]);
const [aliasCollisions, setAliasCollisions] = React.useState<string[]>([]);
const form = useForm<MountedFormValues>({ 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: <ProviderLogo provider={p.provider_display_name} className="w-5 h-5" />,
})),
[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<string, unknown>;
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 (
<div className="mx-auto max-w-3xl">
<h2 className="mb-4 text-2xl font-semibold text-foreground">Add Provider</h2>
<StepIndicator step={step} skipJwks={!isInternalIssuer} />
{step === "provider" && (
<ProviderStep
providerOptions={providerOptions}
selectedProvider={selectedProvider}
onSelectProvider={setSelectedProvider}
credentialName={credentialName}
onCredentialNameChange={setCredentialName}
nameCollision={Boolean(nameCollision)}
onNext={() => goTo("credential")}
/>
)}
{step === "credential" && selectedProvider && (
<Card>
<CardContent>
<FormProvider {...form}>
<MountedFormProvider value={{ control: form.control, registry }}>
<form
onSubmit={(e) => {
e.preventDefault();
void saveCredential();
}}
>
<ProviderSpecificFields selectedProvider={selectedProvider} />
<div className="flex justify-between">
<Button type="button" variant="outline" onClick={() => goTo("provider")}>
<ArrowLeft className="mr-1 size-4" /> Back
</Button>
<Button type="submit">{credentialSaved ? "Save changes" : "Save credential"}</Button>
</div>
</form>
</MountedFormProvider>
</FormProvider>
</CardContent>
</Card>
)}
{step === "jwks" && (
<JwksStep
jwks={jwks}
jwksError={jwksError}
federationRuleId={federationRuleId}
onFederationRuleIdChange={setFederationRuleId}
onBack={() => goTo("credential")}
onNext={() => void confirmFederationRuleId()}
/>
)}
{step === "discover" && (
<DiscoverStep
isDiscovering={isDiscovering}
discoveryError={discoveryError}
onBack={() => goTo(isInternalIssuer ? "jwks" : "credential")}
onRetry={() => void runDiscovery()}
/>
)}
{step === "review" && (
<ReviewModelsStep
rows={rows}
setRows={setRows}
onBack={() => {
goTo("discover");
void runDiscovery();
}}
onCreateModels={() => void createModels()}
/>
)}
{(step === "creating" || step === "done") && (
<ResultsStep
isCreating={isCreating}
isDone={step === "done"}
creationResults={creationResults}
aliasCollisions={aliasCollisions}
/>
)}
</div>
);
}

View file

@ -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<React.SetStateAction<DiscoveredModelRow[]>>;
onBack: () => void;
onCreateModels: () => void;
}
const ReviewModelsStep: React.FC<ReviewModelsStepProps> = ({ rows, setRows, onBack, onCreateModels }) => {
const [manualId, setManualId] = React.useState("");
const updateRow = (id: string, patch: Partial<DiscoveredModelRow>) =>
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 (
<Card>
<CardContent className="space-y-4">
{rows.length === 0 ? (
<p className="text-sm text-muted-foreground">No models discovered. Add one manually below.</p>
) : (
<div className="overflow-x-auto rounded-md border">
<Table>
<TableHeader>
<TableRow>
<TableHead>Upstream ID</TableHead>
<TableHead>Enabled</TableHead>
<TableHead>Model name</TableHead>
<TableHead>Alternate names</TableHead>
<TableHead />
</TableRow>
</TableHeader>
<TableBody>
{rows.map((row) => (
<TableRow key={row.id}>
<TableCell className="whitespace-nowrap font-mono text-xs">{row.upstreamId}</TableCell>
<TableCell>
<Switch
checked={row.enabled}
onCheckedChange={(checked) => updateRow(row.id, { enabled: checked })}
aria-label={`Enable ${row.upstreamId}`}
/>
</TableCell>
<TableCell className="min-w-[160px]">
<Input
value={row.modelName}
onChange={(e) => updateRow(row.id, { modelName: e.target.value })}
aria-label={`Model name for ${row.upstreamId}`}
/>
</TableCell>
<TableCell className="min-w-[200px]">
<MultiSelect
value={row.alternateNames}
onValueChange={(value) => updateRow(row.id, { alternateNames: value })}
options={[]}
allowCustomValues
placeholder="Add alternate names"
emptyText="Type to add an alternate name"
/>
</TableCell>
<TableCell>
{row.manual && (
<Button
variant="ghost"
size="icon-sm"
aria-label={`Remove ${row.upstreamId}`}
onClick={() => removeRow(row.id)}
>
<X className="size-4" />
</Button>
)}
</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</div>
)}
<div className="flex items-end gap-2">
<div className="flex-1">
<label htmlFor="add-provider-manual-model" className="mb-1 block text-xs text-muted-foreground">
Add a hidden model manually
</label>
<Input
id="add-provider-manual-model"
value={manualId}
onChange={(e) => setManualId(e.target.value)}
placeholder="upstream model id"
/>
</div>
<Button type="button" variant="outline" onClick={addManualRow} disabled={!manualId.trim()}>
<Plus className="mr-1 size-4" /> Add
</Button>
</div>
<div className="flex justify-between">
<Button type="button" variant="outline" onClick={onBack}>
<ArrowLeft className="mr-1 size-4" /> Back
</Button>
<Button disabled={rows.length === 0} onClick={onCreateModels}>
Create {rows.length} model{rows.length === 1 ? "" : "s"}
</Button>
</div>
</CardContent>
</Card>
);
};
export default ReviewModelsStep;

View file

@ -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<CreationResult["status"], string> = {
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<ProviderStepProps> = ({
providerOptions,
selectedProvider,
onSelectProvider,
credentialName,
onCredentialNameChange,
nameCollision,
onNext,
}) => (
<Card>
<CardContent className="space-y-4">
<Field>
<FieldLabel htmlFor="add-provider-provider">Provider</FieldLabel>
<SearchSelect
inputId="add-provider-provider"
options={providerOptions}
placeholder="Select a provider"
value={selectedProvider ?? ""}
onValueChange={(value) => onSelectProvider(value as Providers)}
/>
</Field>
<Field>
<FieldLabel htmlFor="add-provider-credential-name">Credential name</FieldLabel>
<Input
id="add-provider-credential-name"
value={credentialName}
onChange={(e) => onCredentialNameChange(e.target.value)}
placeholder="e.g. anthropic-prod"
/>
{nameCollision && <p className="text-sm text-destructive">A credential with this name already exists.</p>}
</Field>
<div className="flex justify-end">
<Button disabled={!selectedProvider || !credentialName || nameCollision} onClick={onNext}>
Next <ArrowRight className="ml-1 size-4" />
</Button>
</div>
</CardContent>
</Card>
);
interface JwksStepProps {
jwks: AnthropicJwks | null;
jwksError: string | null;
federationRuleId: string;
onFederationRuleIdChange: (value: string) => void;
onBack: () => void;
onNext: () => void;
}
export const JwksStep: React.FC<JwksStepProps> = ({
jwks,
jwksError,
federationRuleId,
onFederationRuleIdChange,
onBack,
onNext,
}) => (
<Card>
<CardContent className="space-y-4">
<Alert variant="info">
<AlertTitle>Register this JWKS with Anthropic</AlertTitle>
<AlertDescription>
Register this public JWKS as the inline issuer for your federation rule in the Anthropic Console, then
paste the resulting Federation Rule ID below.
</AlertDescription>
</Alert>
{jwksError && (
<Alert variant="destructive">
<AlertTriangle className="size-4" />
<AlertTitle>Could not load JWKS</AlertTitle>
<AlertDescription>{jwksError}</AlertDescription>
</Alert>
)}
{jwks && (
<div className="relative rounded-md border bg-muted p-3">
<CopyButton value={JSON.stringify(jwks, null, 2)} label="Copy JWKS" className="absolute top-2 right-2" />
<pre className="overflow-x-auto text-xs">{JSON.stringify(jwks, null, 2)}</pre>
</div>
)}
<Field>
<FieldLabel htmlFor="add-provider-federation-rule-id">Federation Rule ID</FieldLabel>
<Input
id="add-provider-federation-rule-id"
value={federationRuleId}
onChange={(e) => onFederationRuleIdChange(e.target.value)}
/>
</Field>
<div className="flex justify-between">
<Button type="button" variant="outline" onClick={onBack}>
<ArrowLeft className="mr-1 size-4" /> Back
</Button>
<Button disabled={!federationRuleId} onClick={onNext}>
Next <ArrowRight className="ml-1 size-4" />
</Button>
</div>
</CardContent>
</Card>
);
interface DiscoverStepProps {
isDiscovering: boolean;
discoveryError: string | null;
onBack: () => void;
onRetry: () => void;
}
export const DiscoverStep: React.FC<DiscoverStepProps> = ({ isDiscovering, discoveryError, onBack, onRetry }) => (
<Card>
<CardContent className="space-y-4">
{isDiscovering && (
<p className="flex items-center gap-2 text-sm text-muted-foreground">
<Loader2 className="size-4 animate-spin" /> Discovering models...
</p>
)}
{discoveryError && (
<Alert variant="destructive">
<AlertTriangle className="size-4" />
<AlertTitle>Discovery failed</AlertTitle>
<AlertDescription>{discoveryError}</AlertDescription>
</Alert>
)}
<div className="flex justify-between">
<Button type="button" variant="outline" onClick={onBack}>
<ArrowLeft className="mr-1 size-4" /> Back
</Button>
{discoveryError && (
<Button disabled={isDiscovering} onClick={onRetry}>
Retry
</Button>
)}
</div>
</CardContent>
</Card>
);
interface ResultsStepProps {
isCreating: boolean;
isDone: boolean;
creationResults: CreationResult[];
aliasCollisions: string[];
}
export const ResultsStep: React.FC<ResultsStepProps> = ({ isCreating, isDone, creationResults, aliasCollisions }) => (
<Card>
<CardContent className="space-y-4">
{isCreating && (
<p className="flex items-center gap-2 text-sm text-muted-foreground">
<Loader2 className="size-4 animate-spin" /> Creating models...
</p>
)}
{isDone && (
<>
<ul className="space-y-1 text-sm">
{creationResults.map((result) => (
<li key={result.row.id}>
<span className={CREATION_RESULT_CLASS_NAME[result.status]}>
{result.row.modelName}: {result.status}
{result.detail ? ` (${result.detail})` : ""}
</span>
</li>
))}
</ul>
{aliasCollisions.length > 0 && (
<Alert variant="destructive">
<AlertTriangle className="size-4" />
<AlertTitle>Some alternate names were not saved</AlertTitle>
<AlertDescription>
These alias names already exist and were left unchanged: {aliasCollisions.join(", ")}
</AlertDescription>
</Alert>
)}
</>
)}
</CardContent>
</Card>
);

View file

@ -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" }]);
});
});

View file

@ -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<string, never>;
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<string, ModelGroupAliasValue>;
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 })));

View file

@ -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");
});
});

View file

@ -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<string, ProviderCredentialFieldMetadata>,
values: Record<string, unknown>,
): 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, unknown>): 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,
);
};

View file

@ -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<MountedFormValues>();
return <output data-testid="identity-source">{String(watch("anthropic_identity_source") ?? "")}</output>;
};
it("defaults to the api_key variant and hides WIF fields", async () => {
const queryClient = createQueryClient();
render(
<QueryClientProvider client={queryClient}>
<MountedFormHost>
<ProviderSpecificFields selectedProvider={Providers.Anthropic} />
</MountedFormHost>
</QueryClientProvider>,
);
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(
<QueryClientProvider client={queryClient}>
<MountedFormHost>
<ProviderSpecificFields selectedProvider={Providers.Anthropic} />
</MountedFormHost>
</QueryClientProvider>,
);
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(
<QueryClientProvider client={queryClient}>
<MountedFormHost>
<ProviderSpecificFields selectedProvider={Providers.Anthropic} />
<IdentitySourceProbe />
</MountedFormHost>
</QueryClientProvider>,
);
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(
<QueryClientProvider client={queryClient}>
<MountedFormHost>
<ProviderSpecificFields selectedProvider={Providers.Anthropic} />
</MountedFormHost>
</QueryClientProvider>,
);
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(
<QueryClientProvider client={queryClient}>
<MountedFormHost
defaultValues={{
anthropic_federation_rule_id: "rule-1",
anthropic_organization_id: "org-1",
anthropic_identity_token: "oidc/env/TOKEN",
}}
>
<ProviderSpecificFields selectedProvider={Providers.Anthropic} />
</MountedFormHost>
</QueryClientProvider>,
);
expect(await screen.findByLabelText("Identity Token Reference")).toBeInTheDocument();
expect(screen.queryByLabelText("API Key")).not.toBeInTheDocument();
});
});
});

View file

@ -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<string, ProviderCredentialField[]> = {};
// 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<string, ProviderCredentialVariants | null> = {};
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<MountedFormValues>();
useMountedName(name);
React.useEffect(() => {
form.setValue(name, value);
}, [form, name, value]);
return null;
};
const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({ selectedProvider }) => {
const selectedProviderEnum = Providers[selectedProvider as keyof typeof Providers] as Providers;
const form = useFormContext<MountedFormValues>();
@ -139,45 +163,50 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({ selecte
}
// Compute cache entries keyed by provider display name and identifiers
const entries: Record<string, ProviderCredentialField[]> = {};
const fieldEntries: Record<string, ProviderCredentialField[]> = {};
const variantEntries: Record<string, ProviderCredentialVariants | null> = {};
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<ProviderSpecificFieldsProps> = ({ 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<string | undefined>(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<string | null>(null);
const handleApiBaseChange = React.useCallback(
@ -313,49 +363,83 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({ selecte
);
};
const renderFieldEntry = (field: ProviderCredentialField) => (
<React.Fragment key={field.key}>
<MountedFormField
label={field.tooltip ? labelWithHint(field.label, field.tooltip) : field.label}
name={field.key}
required={field.required}
rules={field.required ? { validate: { required: requiredRule("Required") } } : undefined}
className={field.key === "vertex_credentials" ? "mb-0" : "mb-4"}
>
{(control) => renderFieldControl(field, control)}
</MountedFormField>
{/* Special case for Vertex Credentials help text */}
{field.key === "vertex_credentials" && (
<p className="text-sm mb-3 mt-1">Give a gcp service account(.json file)</p>
)}
{/* Special case for Azure Base Model help text */}
{field.key === "base_model" && (
<div className="grid grid-cols-24">
<p className="col-start-11 col-span-10 text-sm mb-2">
The actual model your azure deployment uses. Used for accurate cost tracking. Select name from{" "}
<a
href="https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json"
target="_blank"
rel="noopener noreferrer"
className="text-primary underline-offset-4 hover:underline"
>
here
</a>
</p>
</div>
)}
</React.Fragment>
);
const activeVariant = variants ? getVariant(variants, activeVariantId) : undefined;
return (
<>
{isLoading && allFields.length === 0 && <p className="text-sm mb-2">Loading provider fields...</p>}
{loadError && allFields.length === 0 && (
{isLoading && currentFields.length === 0 && <p className="text-sm mb-2">Loading provider fields...</p>}
{loadError && currentFields.length === 0 && (
<p className="text-sm mb-2 text-destructive">
{loadError instanceof Error ? loadError.message : "Failed to load provider credential fields"}
</p>
)}
{allFields.map((field) => (
<React.Fragment key={field.key}>
<MountedFormField
label={field.tooltip ? labelWithHint(field.label, field.tooltip) : field.label}
name={field.key}
required={field.required}
rules={field.required ? { validate: { required: requiredRule("Required") } } : undefined}
className={field.key === "vertex_credentials" ? "mb-0" : "mb-4"}
{variants && (
<Field className="mb-4">
<FieldLabel htmlFor="provider-credential-variant">{variants.selector_label}</FieldLabel>
<Select
items={variants.variants.map((variant) => ({ value: variant.id, label: variant.label }))}
value={activeVariantId || null}
onValueChange={(value) => value && setUserChosenVariantId(value)}
>
{(control) => renderFieldControl(field, control)}
</MountedFormField>
{/* Special case for Vertex Credentials help text */}
{field.key === "vertex_credentials" && (
<p className="text-sm mb-3 mt-1">Give a gcp service account(.json file)</p>
)}
{/* Special case for Azure Base Model help text */}
{field.key === "base_model" && (
<div className="grid grid-cols-24">
<p className="col-start-11 col-span-10 text-sm mb-2">
The actual model your azure deployment uses. Used for accurate cost tracking. Select name from{" "}
<a
href="https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json"
target="_blank"
rel="noopener noreferrer"
className="text-primary underline-offset-4 hover:underline"
>
here
</a>
</p>
</div>
)}
</React.Fragment>
))}
<SelectTrigger id="provider-credential-variant" aria-label={variants.selector_label} className="w-full">
<SelectValue />
</SelectTrigger>
<SelectContent>
{variants.variants.map((variant) => (
<SelectItem key={variant.id} value={variant.id}>
{variant.label}
</SelectItem>
))}
</SelectContent>
</Select>
</Field>
)}
{/* 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. */}
<React.Fragment key={variants ? activeVariantId : "static"}>
{currentFields.map(renderFieldEntry)}
{activeVariant &&
Object.entries(activeVariant.fixed_values).map(([key, value]) => (
<FixedValueField key={key} name={key} value={value} />
))}
</React.Fragment>
</>
);
};

View file

@ -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([]);
});
});
});

View file

@ -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();
};

View file

@ -47,12 +47,15 @@ export default function CredentialsPanel() {
const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false);
const [isCredentialDeleting, setIsCredentialDeleting] = useState(false);
const handleUpdateCredential = async (values: Record<string, unknown>) => {
const handleUpdateCredential = async (values: Record<string, unknown>, 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);

View file

@ -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([]);
});
});

View file

@ -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<string, unknown>,
mountedValues: Record<string, unknown>,
): string[] {
return Object.keys(originalValues).filter((key) => !(key in mountedValues));
}

View file

@ -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<string, string>;
}
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<ProviderCreateInfo[]>
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<ProviderModelDiscoveryResponse> => {
/**
* 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<string, string>[];
}
export const getCredentialJwksCall = async (accessToken: string, credentialName: string): Promise<AnthropicJwks> => {
/**
* 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<string, number>;
token_thresholds: Record<string, number>;
@ -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, {