mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_pricing_auto_update_action
# Conflicts: # uv.lock
This commit is contained in:
commit
35e1b36d11
26 changed files with 1824 additions and 9 deletions
|
|
@ -0,0 +1,9 @@
|
|||
-- CreateTable
|
||||
CREATE TABLE IF NOT EXISTS "LiteLLM_SSOIdentityAssertion" (
|
||||
"user_id" TEXT NOT NULL,
|
||||
"assertion_b64" TEXT NOT NULL,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
"updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
|
||||
CONSTRAINT "LiteLLM_SSOIdentityAssertion_pkey" PRIMARY KEY ("user_id")
|
||||
);
|
||||
|
|
@ -403,6 +403,15 @@ model LiteLLM_MCPServerOAuthClient {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
}
|
||||
|
||||
// The enterprise IdP identity assertion captured at SSO login, one row per user.
|
||||
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}.
|
||||
model LiteLLM_SSOIdentityAssertion {
|
||||
user_id String @id
|
||||
assertion_b64 String
|
||||
created_at DateTime @default(now()) @map("created_at")
|
||||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
}
|
||||
|
||||
// Generate Tokens for Proxy
|
||||
model LiteLLM_VerificationToken {
|
||||
token String @id
|
||||
|
|
|
|||
|
|
@ -1469,6 +1469,7 @@ _batch_polling_env = os.getenv("PROXY_BATCH_POLLING_ENABLED", "true").lower()
|
|||
PROXY_BATCH_POLLING_ENABLED = _batch_polling_env == "true"
|
||||
PROXY_BUDGET_RESCHEDULER_MAX_TIME = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MAX_TIME", 605))
|
||||
PROXY_BATCH_WRITE_AT = int(os.getenv("PROXY_BATCH_WRITE_AT", 10)) # in seconds, increased from 10
|
||||
PROXY_CONFIG_RELOAD_INTERVAL_SECONDS = get_env_int("PROXY_CONFIG_RELOAD_INTERVAL_SECONDS", 30)
|
||||
|
||||
# APScheduler Configuration - MEMORY LEAK FIX
|
||||
# These settings prevent memory leaks in APScheduler's normalize() and _apply_jitter() functions
|
||||
|
|
|
|||
|
|
@ -0,0 +1,214 @@
|
|||
"""Store for the enterprise IdP identity assertion captured at SSO login (EMA).
|
||||
|
||||
The ``oauth2_id_jag`` egress arm needs the user's IdP ``id_token`` as its RFC 8693
|
||||
``subject_token``. A front-door client holds an identity-only ``llm_session_`` bearer, not an
|
||||
IdP assertion, so the assertion captured at the one SSO login is the only usable subject
|
||||
source for it. This module owns both sides of that state: the SSO callback persists here
|
||||
(write-through to the DB so a login on one pod is visible to every pod) and the resolver
|
||||
seam reads back by ``user_id``. Retention is gated on an ``oauth2_id_jag`` server actually
|
||||
being registered, so a gateway with no EMA upstream never stores bearer material.
|
||||
|
||||
The row is one encrypted payload per user, latest login wins. ``expires_at`` mirrors the
|
||||
id_token ``exp`` claim and is judged by the reader, never enforced by deletion here: an
|
||||
expired assertion with a refresh token is still renewable, and the DB row is the source of
|
||||
truth, the same contract as the per-user OAuth credential store.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import jwt
|
||||
from pydantic import BaseModel, ConfigDict, SecretStr, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
_ASSERTION_DECRYPT_LOG_KEY = "sso_identity_assertion"
|
||||
_STR_ADAPTER: TypeAdapter[str] = TypeAdapter(str)
|
||||
_MAYBE_STR_ADAPTER: TypeAdapter[str | None] = TypeAdapter(str | None)
|
||||
|
||||
|
||||
class SSOIdentityAssertion(BaseModel):
|
||||
"""The IdP material an EMA exchange needs: ``id_token`` is the RFC 8693 subject token,
|
||||
``expires_at`` bounds its usefulness, and the refresh token renews it without re-login."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
id_token: SecretStr
|
||||
refresh_token: SecretStr | None = None
|
||||
issuer: str | None = None
|
||||
expires_at: datetime | None = None
|
||||
|
||||
|
||||
class _IdTokenClaims(BaseModel):
|
||||
exp: float | None = None
|
||||
iss: str | None = None
|
||||
|
||||
|
||||
class _StoredAssertionPayload(BaseModel):
|
||||
id_token: str
|
||||
refresh_token: str | None = None
|
||||
issuer: str | None = None
|
||||
expires_at: datetime | None = None
|
||||
|
||||
|
||||
def assertion_from_sso_login(id_token: object, refresh_token: object) -> SSOIdentityAssertion | None:
|
||||
"""The typed carrier built where the raw token response exists; ``None`` when the provider
|
||||
sent no id_token or sent one that is not a decodable JWT, since neither is exchangeable
|
||||
under EMA. Inputs are ``object`` because they come straight from the provider's untyped
|
||||
token response; this is the one boundary that validates them. The token arrived over TLS
|
||||
from the IdP's own token endpoint, so claims are read without signature verification,
|
||||
matching how the SSO callback already decodes it for identity."""
|
||||
raw_id_token = id_token if isinstance(id_token, str) and id_token else None
|
||||
if raw_id_token is None:
|
||||
return None
|
||||
raw_refresh_token = refresh_token if isinstance(refresh_token, str) and refresh_token else None
|
||||
try:
|
||||
claims = _IdTokenClaims.model_validate(jwt.decode(raw_id_token, options={"verify_signature": False}))
|
||||
expires_at = datetime.fromtimestamp(claims.exp, tz=timezone.utc) if claims.exp is not None else None
|
||||
except Exception: # noqa: BLE001 # decode failure = not retainable; never raise into login
|
||||
verbose_proxy_logger.warning(
|
||||
"SSO id_token could not be decoded or its claims were unusable; not retaining it for EMA egress."
|
||||
)
|
||||
return None
|
||||
return SSOIdentityAssertion(
|
||||
id_token=SecretStr(raw_id_token),
|
||||
refresh_token=SecretStr(raw_refresh_token) if raw_refresh_token else None,
|
||||
issuer=claims.iss,
|
||||
expires_at=expires_at,
|
||||
)
|
||||
|
||||
|
||||
async def ema_assertion_retention_enabled() -> bool:
|
||||
"""Whether any MCP server uses ``oauth2_id_jag``, evaluated per login so the gateway only
|
||||
retains bearer material while an EMA upstream exists to spend it on. Judged against the two
|
||||
configuration authorities: the pod-local config declaration and the shared DB row. The
|
||||
in-memory registry is deliberately not consulted in either direction; it is a per-process
|
||||
snapshot of the DB state that can be stale both ways (a server added on another pod would
|
||||
silently drop the write, one removed on another pod would keep retaining bearer material),
|
||||
and a gate guarding a shared-DB write must judge against that storage's authority."""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids import cycle
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
|
||||
from litellm.types.mcp import MCPAuth # noqa: PLC0415 # runtime global
|
||||
|
||||
config_servers = global_mcp_server_manager.config_mcp_servers.values()
|
||||
if any(server.auth_type == MCPAuth.oauth2_id_jag for server in config_servers):
|
||||
return True
|
||||
if prisma_client is None:
|
||||
return False
|
||||
row = await prisma_client.db.litellm_mcpservertable.find_first(where={"auth_type": MCPAuth.oauth2_id_jag.value})
|
||||
return row is not None
|
||||
|
||||
|
||||
async def persist_sso_identity_assertion(user_id: str, assertion: SSOIdentityAssertion) -> None:
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper # noqa: PLC0415 # runtime global
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
|
||||
|
||||
if prisma_client is None:
|
||||
return
|
||||
payload: dict[str, str] = {
|
||||
"id_token": assertion.id_token.get_secret_value(),
|
||||
**({"refresh_token": assertion.refresh_token.get_secret_value()} if assertion.refresh_token else {}),
|
||||
**({"issuer": assertion.issuer} if assertion.issuer else {}),
|
||||
**({"expires_at": assertion.expires_at.isoformat()} if assertion.expires_at else {}),
|
||||
}
|
||||
encoded = _STR_ADAPTER.validate_python(encrypt_value_helper(json.dumps(payload)))
|
||||
await prisma_client.db.litellm_ssoidentityassertion.upsert(
|
||||
where={"user_id": user_id},
|
||||
data={
|
||||
"create": {"user_id": user_id, "assertion_b64": encoded},
|
||||
"update": {"assertion_b64": encoded},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def fetch_sso_identity_assertion(user_id: str) -> SSOIdentityAssertion | None:
|
||||
"""The stored assertion for ``user_id``, or ``None`` when absent, undecryptable (salt-key
|
||||
rotation), or unparseable. Expiry is not judged here; the reader owns that policy."""
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper # noqa: PLC0415 # runtime global
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
|
||||
|
||||
if prisma_client is None:
|
||||
return None
|
||||
row = await prisma_client.db.litellm_ssoidentityassertion.find_unique(where={"user_id": user_id})
|
||||
if row is None:
|
||||
return None
|
||||
raw = _MAYBE_STR_ADAPTER.validate_python(
|
||||
decrypt_value_helper(row.assertion_b64, _ASSERTION_DECRYPT_LOG_KEY, exception_type="debug")
|
||||
)
|
||||
if raw is None:
|
||||
return None
|
||||
try:
|
||||
payload = _StoredAssertionPayload.model_validate_json(raw)
|
||||
except ValidationError:
|
||||
verbose_proxy_logger.warning(
|
||||
"Stored SSO identity assertion for user_id=%s could not be parsed; treating as absent.", user_id
|
||||
)
|
||||
return None
|
||||
return SSOIdentityAssertion(
|
||||
id_token=SecretStr(payload.id_token),
|
||||
refresh_token=SecretStr(payload.refresh_token) if payload.refresh_token else None,
|
||||
issuer=payload.issuer,
|
||||
expires_at=payload.expires_at,
|
||||
)
|
||||
|
||||
|
||||
async def rotate_sso_identity_assertions_master_key(prisma_client: PrismaClient, new_master_key: str) -> None:
|
||||
"""Re-encrypt every stored assertion under ``new_master_key`` during a salt-key rotation,
|
||||
mirroring the sibling per-user credential tables; an unreadable row is skipped so one
|
||||
corrupt row does not abort the rotation. Rows are decrypted one at a time inside the loop
|
||||
so the whole table's plaintext is never held in memory at once."""
|
||||
from prisma.models import LiteLLM_SSOIdentityAssertion as AssertionRow # noqa: PLC0415 # generated at runtime
|
||||
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: PLC0415 # runtime global
|
||||
decrypt_value_helper,
|
||||
encrypt_value_helper,
|
||||
)
|
||||
|
||||
async def _rotate_row(row: AssertionRow) -> bool:
|
||||
plaintext = _MAYBE_STR_ADAPTER.validate_python(
|
||||
decrypt_value_helper(row.assertion_b64, _ASSERTION_DECRYPT_LOG_KEY, exception_type="debug")
|
||||
)
|
||||
if plaintext is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"rotate_sso_identity_assertions_master_key: could not decrypt assertion for user_id=%s, skipping",
|
||||
row.user_id,
|
||||
)
|
||||
return False
|
||||
re_encrypted = _STR_ADAPTER.validate_python(encrypt_value_helper(plaintext, new_encryption_key=new_master_key))
|
||||
await prisma_client.db.litellm_ssoidentityassertion.update(
|
||||
where={"user_id": row.user_id},
|
||||
data={"assertion_b64": re_encrypted},
|
||||
)
|
||||
return True
|
||||
|
||||
rows = await prisma_client.db.litellm_ssoidentityassertion.find_many()
|
||||
outcomes = [await _rotate_row(row) for row in rows]
|
||||
verbose_proxy_logger.info(
|
||||
"rotate_sso_identity_assertions_master_key: rotated %d row(s), skipped %d",
|
||||
sum(outcomes),
|
||||
len(outcomes) - sum(outcomes),
|
||||
)
|
||||
|
||||
|
||||
async def retain_sso_identity_assertion_for_ema(user_id: str, assertion: SSOIdentityAssertion | None) -> None:
|
||||
"""The SSO-callback hook: a no-op unless there is material AND an EMA server is registered.
|
||||
A store failure is logged and swallowed because the login itself must not fail on an
|
||||
egress-side write; the cost of a miss is a 401 challenge at the EMA upstream, not a lockout."""
|
||||
if assertion is None:
|
||||
return
|
||||
try:
|
||||
if not await ema_assertion_retention_enabled():
|
||||
return
|
||||
await persist_sso_identity_assertion(user_id, assertion)
|
||||
except Exception as exc: # noqa: BLE001 # the login itself must not fail on an egress-side write
|
||||
verbose_proxy_logger.warning(
|
||||
"Failed to persist the SSO identity assertion for EMA egress (user_id=%s): %s", user_id, exc
|
||||
)
|
||||
|
|
@ -2297,6 +2297,11 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
None,
|
||||
description="max response size in MB, if a response is larger than this size it will be rejected",
|
||||
)
|
||||
proxy_config_reload_interval_seconds: int = Field(
|
||||
30,
|
||||
gt=0,
|
||||
description="how often (in seconds) each pod reloads config-in-DB objects (models, credentials, guardrails, etc.) when store_model_in_db is enabled; lower values speed up multi-pod convergence at the cost of more DB load. Applied on proxy startup",
|
||||
)
|
||||
cancel_on_disconnect: Optional[bool] = Field(
|
||||
None,
|
||||
description="cancel the in-flight upstream LLM request (non-streaming) when the client disconnects, freeing backend capacity (e.g. a vLLM GPU slot); the request is logged as a 499 failure",
|
||||
|
|
|
|||
|
|
@ -42,6 +42,9 @@ from litellm.proxy._experimental.mcp_server.db import (
|
|||
rotate_mcp_user_credentials_master_key,
|
||||
rotate_mcp_user_env_vars_master_key,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
|
||||
rotate_sso_identity_assertions_master_key,
|
||||
)
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy._types import LiteLLM_VerificationToken, hash_token
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
|
|
@ -4242,6 +4245,15 @@ async def _rotate_master_key(
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.warning("Failed to rotate MCP user env vars: %s", str(e))
|
||||
|
||||
# 4d. process SSO identity assertion table (EMA subject tokens)
|
||||
try:
|
||||
await rotate_sso_identity_assertions_master_key(
|
||||
prisma_client=prisma_client,
|
||||
new_master_key=new_master_key,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # one store's failure must not abort the master-key rotation
|
||||
verbose_proxy_logger.warning("Failed to rotate SSO identity assertions: %s", str(e))
|
||||
|
||||
# 5. process credentials table
|
||||
try:
|
||||
credentials = await CredentialsRepository(prisma_client).table.find_many()
|
||||
|
|
|
|||
|
|
@ -62,6 +62,11 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
|
||||
SSOIdentityAssertion,
|
||||
assertion_from_sso_login,
|
||||
retain_sso_identity_assertion_for_ema,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
CommonProxyErrors,
|
||||
LiteLLM_UserTable,
|
||||
|
|
@ -1311,12 +1316,15 @@ async def get_generic_sso_response(
|
|||
sso_jwt_handler: Optional[JWTHandler], # sso specific jwt handler - used for restricted sso group access control
|
||||
generic_client_id: str,
|
||||
redirect_url: str,
|
||||
) -> Tuple[Union[OpenID, dict], Optional[dict], Optional[dict]]: # (result, received_response, access_token_payload)
|
||||
) -> tuple[
|
||||
Union[OpenID, dict], dict | None, dict | None, SSOIdentityAssertion | None
|
||||
]: # (result, received_response, access_token_payload, sso_assertion)
|
||||
# make generic sso provider
|
||||
from fastapi_sso.sso.base import DiscoveryDocument
|
||||
from fastapi_sso.sso.generic import create_provider
|
||||
|
||||
received_response: Optional[dict] = None
|
||||
sso_assertion: SSOIdentityAssertion | None = None
|
||||
|
||||
# Setup environment variables
|
||||
(
|
||||
|
|
@ -1450,6 +1458,9 @@ async def get_generic_sso_response(
|
|||
# Assign directly rather than relying on nonlocal mutation so that Pyright
|
||||
# can track that received_response is non-None from this point on.
|
||||
received_response = {k: v for k, v in combined_response.items() if k not in _OAUTH_TOKEN_FIELDS}
|
||||
sso_assertion = assertion_from_sso_login(
|
||||
combined_response.get("id_token"), combined_response.get("refresh_token")
|
||||
)
|
||||
# In the PKCE path verify_and_process is skipped, so generic_sso.access_token
|
||||
# is never set. Read the token directly from the exchange response instead so
|
||||
# process_sso_jwt_access_token can extract JWT-embedded roles/teams.
|
||||
|
|
@ -1461,6 +1472,7 @@ async def get_generic_sso_response(
|
|||
headers=additional_generic_sso_headers_dict,
|
||||
)
|
||||
access_token_str = generic_sso.access_token
|
||||
sso_assertion = assertion_from_sso_login(generic_sso.id_token, generic_sso.refresh_token)
|
||||
|
||||
access_token_payload = process_sso_jwt_access_token(
|
||||
access_token_str, sso_jwt_handler, result, role_mappings=role_mappings
|
||||
|
|
@ -1480,7 +1492,7 @@ async def get_generic_sso_response(
|
|||
additional_generic_sso_headers_dict,
|
||||
)
|
||||
verbose_proxy_logger.debug("generic result: %s", result)
|
||||
return result or {}, received_response, access_token_payload
|
||||
return result or {}, received_response, access_token_payload, sso_assertion
|
||||
|
||||
|
||||
async def create_team_member_add_task(team_id, user_info):
|
||||
|
|
@ -1812,6 +1824,7 @@ async def auth_callback(request: Request, state: Optional[str] = None):
|
|||
generic_client_id = os.getenv("GENERIC_CLIENT_ID", None)
|
||||
received_response: Optional[dict] = None
|
||||
access_token_payload: Optional[dict] = None
|
||||
sso_assertion: SSOIdentityAssertion | None = None
|
||||
# get url from request
|
||||
if master_key is None:
|
||||
raise ProxyException(
|
||||
|
|
@ -1842,6 +1855,7 @@ async def auth_callback(request: Request, state: Optional[str] = None):
|
|||
result,
|
||||
received_response,
|
||||
access_token_payload,
|
||||
sso_assertion,
|
||||
) = await get_generic_sso_response(
|
||||
request=request,
|
||||
jwt_handler=jwt_handler,
|
||||
|
|
@ -1869,6 +1883,7 @@ async def auth_callback(request: Request, state: Optional[str] = None):
|
|||
prefill_user_code=prefill_user_code,
|
||||
result=result,
|
||||
received_response=received_response,
|
||||
sso_assertion=sso_assertion,
|
||||
)
|
||||
|
||||
# Control-plane cross-origin: read return_to from cookie.
|
||||
|
|
@ -1884,6 +1899,7 @@ async def auth_callback(request: Request, state: Optional[str] = None):
|
|||
access_token_payload=access_token_payload,
|
||||
jwt_handler=jwt_handler,
|
||||
return_to=cp_return_to,
|
||||
sso_assertion=sso_assertion,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1943,6 +1959,7 @@ async def _complete_cli_sso_callback_session(
|
|||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
prefill_user_code: str | None = None,
|
||||
sso_assertion: SSOIdentityAssertion | None = None,
|
||||
):
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
||||
|
|
@ -1962,6 +1979,8 @@ async def _complete_cli_sso_callback_session(
|
|||
if not user_info.user_id:
|
||||
raise HTTPException(status_code=500, detail="Failed to retrieve user information from SSO")
|
||||
|
||||
await retain_sso_identity_assertion_for_ema(user_id=user_info.user_id, assertion=sso_assertion)
|
||||
|
||||
teams: List[str] = []
|
||||
if hasattr(user_info, "teams") and user_info.teams:
|
||||
teams = user_info.teams if isinstance(user_info.teams, list) else []
|
||||
|
|
@ -2012,6 +2031,7 @@ async def cli_sso_callback(
|
|||
result: Optional[Union[OpenID, dict]] = None,
|
||||
received_response: Optional[dict] = None,
|
||||
prefill_user_code: str | None = None,
|
||||
sso_assertion: SSOIdentityAssertion | None = None,
|
||||
):
|
||||
"""CLI SSO callback - stores session info for JWT generation on polling"""
|
||||
verbose_proxy_logger.info("CLI SSO callback")
|
||||
|
|
@ -2065,6 +2085,7 @@ async def cli_sso_callback(
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
prefill_user_code=prefill_user_code,
|
||||
sso_assertion=sso_assertion,
|
||||
)
|
||||
except ProxyException:
|
||||
raise
|
||||
|
|
@ -3018,6 +3039,7 @@ class SSOAuthenticationHandler:
|
|||
access_token_payload: Optional[dict] = None,
|
||||
jwt_handler: Optional[JWTHandler] = None,
|
||||
return_to: Optional[str] = None,
|
||||
sso_assertion: SSOIdentityAssertion | None = None,
|
||||
) -> RedirectResponse:
|
||||
import jwt
|
||||
|
||||
|
|
@ -3148,6 +3170,9 @@ class SSOAuthenticationHandler:
|
|||
},
|
||||
)
|
||||
|
||||
if isinstance(user_id, str) and user_id:
|
||||
await retain_sso_identity_assertion_for_ema(user_id=user_id, assertion=sso_assertion)
|
||||
|
||||
disabled_non_admin_personal_key_creation = get_disabled_non_admin_personal_key_creation()
|
||||
litellm_dashboard_ui = get_custom_url(request_base_url=str(request.base_url), route="ui/")
|
||||
|
||||
|
|
@ -4241,6 +4266,7 @@ async def debug_sso_callback(request: Request):
|
|||
result,
|
||||
received_response,
|
||||
access_token_payload,
|
||||
_sso_assertion,
|
||||
) = await get_generic_sso_response(
|
||||
request=request,
|
||||
jwt_handler=jwt_handler,
|
||||
|
|
|
|||
|
|
@ -236,6 +236,7 @@ from litellm.constants import (
|
|||
PROXY_BATCH_WRITE_AT,
|
||||
PROXY_BUDGET_RESCHEDULER_MAX_TIME,
|
||||
PROXY_BUDGET_RESCHEDULER_MIN_TIME,
|
||||
PROXY_CONFIG_RELOAD_INTERVAL_SECONDS,
|
||||
)
|
||||
from litellm.exceptions import RejectedRequestError
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
|
|
@ -2001,6 +2002,7 @@ proxy_budget_rescheduler_min_time = PROXY_BUDGET_RESCHEDULER_MIN_TIME
|
|||
proxy_budget_rescheduler_max_time = PROXY_BUDGET_RESCHEDULER_MAX_TIME
|
||||
proxy_batch_polling_interval = PROXY_BATCH_POLLING_INTERVAL
|
||||
proxy_batch_write_at = PROXY_BATCH_WRITE_AT
|
||||
proxy_config_reload_interval_seconds = PROXY_CONFIG_RELOAD_INTERVAL_SECONDS
|
||||
litellm_master_key_hash = None
|
||||
disable_spend_logs = False
|
||||
jwt_handler = JWTHandler()
|
||||
|
|
@ -4337,6 +4339,7 @@ class ProxyConfig:
|
|||
open_telemetry_logger, \
|
||||
health_check_details, \
|
||||
proxy_batch_polling_interval, \
|
||||
proxy_config_reload_interval_seconds, \
|
||||
config_passthrough_endpoints
|
||||
|
||||
config: dict = await self.get_config(config_file_path=config_file_path)
|
||||
|
|
@ -4826,6 +4829,10 @@ class ProxyConfig:
|
|||
)
|
||||
## BATCH WRITER ##
|
||||
proxy_batch_write_at = general_settings.get("proxy_batch_write_at", proxy_batch_write_at)
|
||||
## DB CONFIG RELOAD INTERVAL ##
|
||||
proxy_config_reload_interval_seconds = general_settings.get(
|
||||
"proxy_config_reload_interval_seconds", proxy_config_reload_interval_seconds
|
||||
)
|
||||
## DISABLE SPEND LOGS ## - gives a perf improvement
|
||||
disable_spend_logs = general_settings.get("disable_spend_logs", disable_spend_logs)
|
||||
### BACKGROUND HEALTH CHECKS ###
|
||||
|
|
@ -7998,12 +8005,20 @@ class ProxyStartupEvent:
|
|||
verbose_proxy_logger.debug("Failed to check DB for store_model_in_db: %s", str(e))
|
||||
|
||||
if store_model_in_db is True:
|
||||
config_reload_interval_seconds = proxy_config_reload_interval_seconds
|
||||
if not isinstance(config_reload_interval_seconds, int) or config_reload_interval_seconds <= 0:
|
||||
verbose_proxy_logger.warning(
|
||||
"proxy_config_reload_interval_seconds=%s must be a positive integer; falling back to 30s",
|
||||
config_reload_interval_seconds,
|
||||
)
|
||||
config_reload_interval_seconds = 30
|
||||
|
||||
# MEMORY LEAK FIX: Increase interval from 10s to 30s minimum
|
||||
# Frequent polling was causing excessive memory allocations
|
||||
scheduler.add_job(
|
||||
proxy_config.add_deployment,
|
||||
"interval",
|
||||
seconds=30, # increased from 10s to reduce memory pressure
|
||||
seconds=config_reload_interval_seconds,
|
||||
# REMOVED jitter parameter - major cause of memory leak
|
||||
args=[prisma_client, proxy_logging_obj],
|
||||
id="add_deployment_job",
|
||||
|
|
@ -8018,7 +8033,7 @@ class ProxyStartupEvent:
|
|||
scheduler.add_job(
|
||||
proxy_config.get_credentials,
|
||||
"interval",
|
||||
seconds=30, # increased from 10s to reduce memory pressure
|
||||
seconds=config_reload_interval_seconds,
|
||||
# REMOVED jitter parameter - major cause of memory leak
|
||||
args=[prisma_client],
|
||||
id="get_credentials_job",
|
||||
|
|
@ -15052,6 +15067,7 @@ async def get_config_list(
|
|||
"global_max_parallel_requests": {"type": "Integer"},
|
||||
"max_request_size_mb": {"type": "Integer"},
|
||||
"max_response_size_mb": {"type": "Integer"},
|
||||
"proxy_config_reload_interval_seconds": {"type": "Integer"},
|
||||
"pass_through_endpoints": {"type": "PydanticModel"},
|
||||
"store_model_in_db": {"type": "Boolean"},
|
||||
"store_prompts_in_spend_logs": {"type": "Boolean"},
|
||||
|
|
|
|||
|
|
@ -403,6 +403,15 @@ model LiteLLM_MCPServerOAuthClient {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
}
|
||||
|
||||
// The enterprise IdP identity assertion captured at SSO login, one row per user.
|
||||
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}.
|
||||
model LiteLLM_SSOIdentityAssertion {
|
||||
user_id String @id
|
||||
assertion_b64 String
|
||||
created_at DateTime @default(now()) @map("created_at")
|
||||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
}
|
||||
|
||||
// Generate Tokens for Proxy
|
||||
model LiteLLM_VerificationToken {
|
||||
token String @id
|
||||
|
|
|
|||
|
|
@ -403,6 +403,15 @@ model LiteLLM_MCPServerOAuthClient {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
}
|
||||
|
||||
// The enterprise IdP identity assertion captured at SSO login, one row per user.
|
||||
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}.
|
||||
model LiteLLM_SSOIdentityAssertion {
|
||||
user_id String @id
|
||||
assertion_b64 String
|
||||
created_at DateTime @default(now()) @map("created_at")
|
||||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
}
|
||||
|
||||
// Generate Tokens for Proxy
|
||||
model LiteLLM_VerificationToken {
|
||||
token String @id
|
||||
|
|
|
|||
385
tests/e2e/management/test_model_tag_accessgroup_e2e.py
Normal file
385
tests/e2e/management/test_model_tag_accessgroup_e2e.py
Normal file
|
|
@ -0,0 +1,385 @@
|
|||
"""Live e2e: the model, tag, and model-access-group management routes.
|
||||
|
||||
Each test creates its resources under unique names (deleted on teardown) and
|
||||
asserts the route's contract against a live proxy: the admin-only guard on
|
||||
adding a global model, the tag inventory round-trip through /tag/list and
|
||||
/tag/delete, and creating a model access group then reading it back through
|
||||
/access_group/{name}/info. Reads that lag a write poll to a deadline instead of
|
||||
asserting once.
|
||||
|
||||
Request bodies for /model/new are the shared pydantic models; every response
|
||||
this suite reads is modelled locally so the file is self-contained and no
|
||||
untyped dict crosses the boundary.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel, ConfigDict, RootModel
|
||||
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import NoBody, unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from management_client import ManagementClient
|
||||
from models import KeyGenerateBody, LiteLLMParamsBody, ModelInfoBody, ModelNewBody
|
||||
from proxy_client import ProxyClient
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
_MODEL_PERMISSION_DENIED_MARKER = "does not have permission to make this model call"
|
||||
_DUMMY_MODEL = "openai/gpt-5.5"
|
||||
_DUMMY_API_KEY = "e2e-dummy-key"
|
||||
|
||||
|
||||
def _poll[T](proxy: ProxyClient, attempt: Callable[[], T | None], failure: str) -> T:
|
||||
deadline = time.monotonic() + proxy.poll_timeout
|
||||
while time.monotonic() < deadline:
|
||||
found = attempt()
|
||||
if found is not None:
|
||||
return found
|
||||
time.sleep(proxy.poll_interval)
|
||||
pytest.fail(failure)
|
||||
|
||||
|
||||
# ---------- tag route models / helpers ----------
|
||||
|
||||
|
||||
class TagCreateBody(BaseModel):
|
||||
name: str
|
||||
description: str | None = None
|
||||
|
||||
|
||||
class TagDeleteBody(BaseModel):
|
||||
name: str
|
||||
|
||||
|
||||
class TagEntry(BaseModel):
|
||||
name: str
|
||||
description: str | None = None
|
||||
|
||||
|
||||
class TagCatalog(RootModel[list[TagEntry]]):
|
||||
"""GET /tag/list answers with a bare array of tag configs, not an object
|
||||
wrapping them; read the rows off .root."""
|
||||
|
||||
|
||||
def _tag_list(client: ManagementClient) -> tuple[TagEntry, ...]:
|
||||
return tuple(
|
||||
unwrap(
|
||||
client.proxy.transport.get(
|
||||
"/tag/list",
|
||||
headers=client.proxy.transport.master,
|
||||
params=NoBody(),
|
||||
response_type=TagCatalog,
|
||||
)
|
||||
).root
|
||||
)
|
||||
|
||||
|
||||
def _create_tag(client: ManagementClient, body: TagCreateBody) -> None:
|
||||
_ = unwrap(
|
||||
client.proxy.transport.post(
|
||||
"/tag/new",
|
||||
headers=client.proxy.transport.master,
|
||||
json=body,
|
||||
response_type=NoBody,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _delete_tag(client: ManagementClient, name: str) -> None:
|
||||
"""Best-effort delete for teardown: a repeat /tag/delete on an already-deleted
|
||||
tag is a no-op the warn-only teardown absorbs."""
|
||||
_ = client.proxy.transport.post(
|
||||
"/tag/delete",
|
||||
headers=client.proxy.transport.master,
|
||||
json=TagDeleteBody(name=name),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
||||
|
||||
def _delete_tag_strict(client: ManagementClient, name: str) -> None:
|
||||
"""Strict delete for the act phase: a failed /tag/delete is a hard failure."""
|
||||
_ = unwrap(
|
||||
client.proxy.transport.post(
|
||||
"/tag/delete",
|
||||
headers=client.proxy.transport.master,
|
||||
json=TagDeleteBody(name=name),
|
||||
response_type=NoBody,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
# ---------- access group route models / helpers ----------
|
||||
|
||||
|
||||
class AccessGroupNewBody(BaseModel):
|
||||
access_group: str
|
||||
model_names: list[str]
|
||||
|
||||
|
||||
class AccessGroupNewResponse(BaseModel):
|
||||
access_group: str
|
||||
models_updated: int
|
||||
|
||||
|
||||
class AccessGroupInfoResponse(BaseModel):
|
||||
access_group: str
|
||||
model_names: list[str]
|
||||
deployment_count: int
|
||||
|
||||
|
||||
def _create_access_group(client: ManagementClient, body: AccessGroupNewBody) -> AccessGroupNewResponse:
|
||||
return unwrap(
|
||||
client.proxy.transport.post(
|
||||
"/access_group/new",
|
||||
headers=client.proxy.transport.master,
|
||||
json=body,
|
||||
response_type=AccessGroupNewResponse,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _access_group_info(client: ManagementClient, access_group: str) -> AccessGroupInfoResponse | None:
|
||||
result = client.proxy.transport.get(
|
||||
f"/access_group/{access_group}/info",
|
||||
headers=client.proxy.transport.master,
|
||||
params=NoBody(),
|
||||
response_type=AccessGroupInfoResponse,
|
||||
)
|
||||
return unwrap(result) if result.kind == "success" else None
|
||||
|
||||
|
||||
def _delete_access_group(client: ManagementClient, access_group: str) -> None:
|
||||
"""Best-effort delete for teardown; deleting the model behind it removes the
|
||||
access group too, so a repeat delete is a no-op the teardown absorbs."""
|
||||
_ = client.proxy.transport.delete(
|
||||
f"/access_group/{access_group}/delete",
|
||||
headers=client.proxy.transport.master,
|
||||
json=NoBody(),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
||||
|
||||
def _create_db_model(client: ManagementClient, resources: ResourceManager, model_name: str) -> str:
|
||||
model_id = client.proxy.create_model(
|
||||
model_name, LiteLLMParamsBody(model=_DUMMY_MODEL, api_key=_DUMMY_API_KEY)
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_model(model_id))
|
||||
return model_id
|
||||
|
||||
|
||||
# ---------- model block route models / helpers ----------
|
||||
|
||||
|
||||
class ModelBlockBody(BaseModel):
|
||||
model_config = ConfigDict(protected_namespaces=())
|
||||
model_id: str
|
||||
|
||||
|
||||
class ModelInfoBlockDetail(BaseModel):
|
||||
id: str | None = None
|
||||
blocked: bool | None = None
|
||||
|
||||
|
||||
class ModelInfoBlockEntry(BaseModel):
|
||||
model_config = ConfigDict(protected_namespaces=())
|
||||
model_name: str
|
||||
model_info: ModelInfoBlockDetail = ModelInfoBlockDetail()
|
||||
|
||||
|
||||
class ModelInfoCatalog(BaseModel):
|
||||
data: list[ModelInfoBlockEntry] = []
|
||||
|
||||
|
||||
def _model_blocked_flag(client: ManagementClient, model_id: str) -> bool | None:
|
||||
catalog = unwrap(
|
||||
client.proxy.transport.get(
|
||||
"/model/info",
|
||||
headers=client.proxy.transport.master,
|
||||
params=NoBody(),
|
||||
response_type=ModelInfoCatalog,
|
||||
)
|
||||
)
|
||||
entry = next((row for row in catalog.data if row.model_info.id == model_id), None)
|
||||
return entry.model_info.blocked if entry is not None else None
|
||||
|
||||
|
||||
class TestModelRoutes:
|
||||
@pytest.mark.covers("mgmt.model.add.admin_only")
|
||||
def test_non_admin_key_cannot_add_global_model(
|
||||
self, client: ManagementClient, resources: ResourceManager
|
||||
) -> None:
|
||||
key = client.proxy.generate_key(KeyGenerateBody(models=[]))
|
||||
resources.defer(lambda: client.proxy.delete_key(key))
|
||||
|
||||
model_name = f"e2e-mgmt-model-forbidden-{unique_marker()}"
|
||||
outcome = client.proxy.transport.send(
|
||||
"/model/new",
|
||||
headers=client.proxy.transport.bearer(key),
|
||||
json=ModelNewBody(
|
||||
model_name=model_name,
|
||||
litellm_params=LiteLLMParamsBody(model=_DUMMY_MODEL, api_key=_DUMMY_API_KEY),
|
||||
model_info=ModelInfoBody(),
|
||||
),
|
||||
)
|
||||
|
||||
assert outcome.status_code == 403, (
|
||||
f"non-admin key adding a global model (no team_id) must be denied 403, got "
|
||||
f"{outcome.status_code}: {outcome.body[:300]}"
|
||||
)
|
||||
assert _MODEL_PERMISSION_DENIED_MARKER in outcome.body, (
|
||||
f"403 body must be the model-permission denial, got: {outcome.body[:300]}"
|
||||
)
|
||||
|
||||
cataloged = [entry.model_name for entry in client.proxy.model_info()]
|
||||
assert model_name not in cataloged, (
|
||||
f"{model_name!r} was registered in /model/info despite the 403; the admin-only "
|
||||
f"guard did not block the write"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("mgmt.model.block.persists")
|
||||
def test_block_then_unblock_persists_to_model_info(
|
||||
self, client: ManagementClient, resources: ResourceManager
|
||||
) -> None:
|
||||
"""The blocked flag's persistence is read back from /model/info, not from the
|
||||
/model/block response: that route currently returns a non-2xx serialization
|
||||
envelope even though the DB write lands, so the /model/info read-back is the
|
||||
authoritative persistence contract and keeps this test valid once the
|
||||
response shape is fixed."""
|
||||
model_name = f"e2e-mgmt-model-block-{unique_marker()}"
|
||||
model_id = _create_db_model(client, resources, model_name)
|
||||
|
||||
assert _model_blocked_flag(client, model_id) is not True, (
|
||||
f"{model_name!r} already reports blocked in /model/info before /model/block ran"
|
||||
)
|
||||
|
||||
_ = client.proxy.transport.send(
|
||||
"/model/block",
|
||||
headers=client.proxy.transport.master,
|
||||
json=ModelBlockBody(model_id=model_id),
|
||||
)
|
||||
_ = _poll(
|
||||
client.proxy,
|
||||
lambda: True if _model_blocked_flag(client, model_id) is True else None,
|
||||
f"/model/info never reported {model_name!r} blocked after /model/block",
|
||||
)
|
||||
|
||||
_ = client.proxy.transport.send(
|
||||
"/model/unblock",
|
||||
headers=client.proxy.transport.master,
|
||||
json=ModelBlockBody(model_id=model_id),
|
||||
)
|
||||
_ = _poll(
|
||||
client.proxy,
|
||||
lambda: True if _model_blocked_flag(client, model_id) is not True else None,
|
||||
f"/model/info never cleared blocked for {model_name!r} after /model/unblock",
|
||||
)
|
||||
|
||||
|
||||
class TestTagRoutes:
|
||||
@pytest.mark.covers("mgmt.tag.list.happy_path")
|
||||
def test_tag_list_reports_created_tag(self, client: ManagementClient, resources: ResourceManager) -> None:
|
||||
name = f"e2e-mgmt-tag-{unique_marker()}"
|
||||
description = "coverage: tag inventory"
|
||||
assert all(entry.name != name for entry in _tag_list(client)), (
|
||||
f"tag {name!r} was already listed by /tag/list before /tag/new created it"
|
||||
)
|
||||
|
||||
_create_tag(client, TagCreateBody(name=name, description=description))
|
||||
resources.defer(lambda: _delete_tag(client, name))
|
||||
|
||||
entry = _poll(
|
||||
client.proxy,
|
||||
lambda: next((entry for entry in _tag_list(client) if entry.name == name), None),
|
||||
f"/tag/list never listed {name!r} after /tag/new",
|
||||
)
|
||||
assert entry.description == description, (
|
||||
f"/tag/list reports description {entry.description!r} for {name!r}, configured {description!r}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("mgmt.tag.delete.persists")
|
||||
def test_tag_delete_removes_from_list(self, client: ManagementClient, resources: ResourceManager) -> None:
|
||||
"""The teardown's deferred delete fires again on the already-deleted tag by
|
||||
design: it is the safety net if this test fails before the in-body delete,
|
||||
and a repeat /tag/delete is a warn-only no-op the teardown absorbs."""
|
||||
name = f"e2e-mgmt-tag-{unique_marker()}"
|
||||
_create_tag(client, TagCreateBody(name=name))
|
||||
resources.defer(lambda: _delete_tag(client, name))
|
||||
|
||||
_ = _poll(
|
||||
client.proxy,
|
||||
lambda: True if any(entry.name == name for entry in _tag_list(client)) else None,
|
||||
f"/tag/list never listed {name!r} after /tag/new; cannot prove deletion removes it",
|
||||
)
|
||||
|
||||
_delete_tag_strict(client, name)
|
||||
|
||||
_ = _poll(
|
||||
client.proxy,
|
||||
lambda: True if all(entry.name != name for entry in _tag_list(client)) else None,
|
||||
f"{name!r} still present in /tag/list after /tag/delete at the deadline",
|
||||
)
|
||||
|
||||
|
||||
class TestModelAccessGroupRoutes:
|
||||
@pytest.mark.covers("mgmt.access_group.new.happy_path")
|
||||
def test_new_access_group_tags_the_deployment(
|
||||
self, client: ManagementClient, resources: ResourceManager
|
||||
) -> None:
|
||||
model_name = f"e2e-mgmt-agmodel-{unique_marker()}"
|
||||
_ = _create_db_model(client, resources, model_name)
|
||||
|
||||
access_group = f"e2e-mgmt-ag-{unique_marker()}"
|
||||
created = _create_access_group(
|
||||
client, AccessGroupNewBody(access_group=access_group, model_names=[model_name])
|
||||
)
|
||||
resources.defer(lambda: _delete_access_group(client, access_group))
|
||||
|
||||
assert created.access_group == access_group, (
|
||||
f"/access_group/new echoed access_group {created.access_group!r}, requested {access_group!r}"
|
||||
)
|
||||
assert created.models_updated >= 1, (
|
||||
f"/access_group/new tagged {created.models_updated} deployments for {model_name!r}, expected >= 1"
|
||||
)
|
||||
|
||||
info = _poll(
|
||||
client.proxy,
|
||||
lambda: _access_group_info(client, access_group),
|
||||
f"/access_group/{access_group}/info never resolved the group created by /access_group/new",
|
||||
)
|
||||
assert model_name in info.model_names, (
|
||||
f"the group created by /access_group/new does not list {model_name!r} on read-back; "
|
||||
f"/access_group/info reports members {info.model_names}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("mgmt.access_group.info.happy_path")
|
||||
def test_access_group_info_reports_membership(
|
||||
self, client: ManagementClient, resources: ResourceManager
|
||||
) -> None:
|
||||
model_name = f"e2e-mgmt-agmodel-{unique_marker()}"
|
||||
_ = _create_db_model(client, resources, model_name)
|
||||
|
||||
access_group = f"e2e-mgmt-ag-{unique_marker()}"
|
||||
_ = _create_access_group(
|
||||
client, AccessGroupNewBody(access_group=access_group, model_names=[model_name])
|
||||
)
|
||||
resources.defer(lambda: _delete_access_group(client, access_group))
|
||||
|
||||
info = _poll(
|
||||
client.proxy,
|
||||
lambda: _access_group_info(client, access_group),
|
||||
f"/access_group/{access_group}/info never resolved the created access group",
|
||||
)
|
||||
assert info.access_group == access_group, (
|
||||
f"/access_group/info reports access_group {info.access_group!r}, created {access_group!r}"
|
||||
)
|
||||
assert model_name in info.model_names, (
|
||||
f"/access_group/info reports members {info.model_names}, expected to include {model_name!r}"
|
||||
)
|
||||
assert info.deployment_count >= 1, (
|
||||
f"/access_group/info reports deployment_count {info.deployment_count}, expected >= 1"
|
||||
)
|
||||
|
|
@ -234,6 +234,7 @@ CONTROL_PLANE_PREFIXES: tuple[str, ...] = (
|
|||
"/tag",
|
||||
"/budget",
|
||||
"/model/",
|
||||
"/access_group",
|
||||
"/spend",
|
||||
"/global",
|
||||
"/config",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,343 @@
|
|||
"""Tests for the SSO identity assertion store (EMA subject-token capture).
|
||||
|
||||
Pins the contract of the store that PR 2's ``_id_jag`` subject-sourcing seam will read:
|
||||
the carrier validates untyped IdP token-response values at the boundary, retention is
|
||||
gated on an ``oauth2_id_jag`` server being registered, the row is encrypted at rest and
|
||||
round-trips exactly, a store failure never escapes into the login path, and a salt-key
|
||||
rotation re-encrypts stored rows like the sibling per-user credential tables.
|
||||
"""
|
||||
|
||||
import json
|
||||
import time
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import jwt as pyjwt
|
||||
import pytest
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
|
||||
assertion_from_sso_login,
|
||||
ema_assertion_retention_enabled,
|
||||
fetch_sso_identity_assertion,
|
||||
persist_sso_identity_assertion,
|
||||
retain_sso_identity_assertion_for_ema,
|
||||
rotate_sso_identity_assertions_master_key,
|
||||
)
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
SALT_KEY = "test-salt-key-for-sso-assertion-tests-1234"
|
||||
SIGNING_KEY = "test-idp-signing-key-32-bytes-long-xxxx"
|
||||
ISSUER = "https://idp.example.com"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _set_salt_key(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", SALT_KEY)
|
||||
|
||||
|
||||
def _make_id_token(exp_offset: int = 3600, iss: str = ISSUER) -> str:
|
||||
return pyjwt.encode(
|
||||
{"iss": iss, "sub": "u1", "exp": int(time.time()) + exp_offset},
|
||||
SIGNING_KEY,
|
||||
algorithm="HS256",
|
||||
)
|
||||
|
||||
|
||||
def _make_prisma(stored: dict, db_has_id_jag_server: bool = False):
|
||||
"""A fake prisma client whose sso-assertion table reads and writes ``stored``
|
||||
(user_id -> assertion_b64), covering upsert, find_unique, find_many, and update.
|
||||
``db_has_id_jag_server`` drives the retention gate's authoritative DB fallback;
|
||||
it is wired explicitly so the gate never reads a truthy bare MagicMock."""
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpservertable.find_first = AsyncMock(
|
||||
return_value=MagicMock() if db_has_id_jag_server else None
|
||||
)
|
||||
|
||||
async def _upsert(where, data):
|
||||
stored[where["user_id"]] = data["update"]["assertion_b64"]
|
||||
|
||||
async def _find_unique(where):
|
||||
blob = stored.get(where["user_id"])
|
||||
if blob is None:
|
||||
return None
|
||||
row = MagicMock()
|
||||
row.user_id = where["user_id"]
|
||||
row.assertion_b64 = blob
|
||||
return row
|
||||
|
||||
async def _find_many():
|
||||
rows = []
|
||||
for user_id, blob in stored.items():
|
||||
row = MagicMock()
|
||||
row.user_id = user_id
|
||||
row.assertion_b64 = blob
|
||||
rows.append(row)
|
||||
return rows
|
||||
|
||||
async def _update(where, data):
|
||||
stored[where["user_id"]] = data["assertion_b64"]
|
||||
|
||||
prisma.db.litellm_ssoidentityassertion.upsert = AsyncMock(side_effect=_upsert)
|
||||
prisma.db.litellm_ssoidentityassertion.find_unique = AsyncMock(side_effect=_find_unique)
|
||||
prisma.db.litellm_ssoidentityassertion.find_many = AsyncMock(side_effect=_find_many)
|
||||
prisma.db.litellm_ssoidentityassertion.update = AsyncMock(side_effect=_update)
|
||||
return prisma
|
||||
|
||||
|
||||
def _server_with_auth(auth_type):
|
||||
server = MagicMock()
|
||||
server.auth_type = auth_type
|
||||
return server
|
||||
|
||||
|
||||
def test_assertion_from_sso_login_happy_path():
|
||||
token = _make_id_token()
|
||||
assertion = assertion_from_sso_login(token, "rt_1")
|
||||
assert assertion is not None
|
||||
assert assertion.id_token.get_secret_value() == token
|
||||
assert assertion.refresh_token is not None
|
||||
assert assertion.refresh_token.get_secret_value() == "rt_1"
|
||||
assert assertion.issuer == ISSUER
|
||||
assert assertion.expires_at is not None
|
||||
assert assertion.expires_at.timestamp() == pytest.approx(time.time() + 3600, abs=5)
|
||||
|
||||
|
||||
def test_assertion_repr_never_leaks_token_material():
|
||||
token = _make_id_token()
|
||||
assertion = assertion_from_sso_login(token, "rt_secret_value")
|
||||
rendered = repr(assertion) + str(assertion)
|
||||
assert token not in rendered
|
||||
assert "rt_secret_value" not in rendered
|
||||
|
||||
|
||||
@pytest.mark.parametrize("id_token", [None, "", "not-a-jwt", 12345, ["x"], {"a": 1}])
|
||||
def test_assertion_from_sso_login_rejects_unusable_id_token(id_token):
|
||||
assert assertion_from_sso_login(id_token, "rt") is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("refresh_token", [None, "", 123, ["rt"], {"rt": 1}])
|
||||
def test_assertion_from_sso_login_drops_malformed_refresh_token(refresh_token):
|
||||
assertion = assertion_from_sso_login(_make_id_token(), refresh_token)
|
||||
assert assertion is not None
|
||||
assert assertion.refresh_token is None
|
||||
|
||||
|
||||
def test_assertion_without_exp_or_iss_still_retained():
|
||||
token = pyjwt.encode({"sub": "u1"}, SIGNING_KEY, algorithm="HS256")
|
||||
assertion = assertion_from_sso_login(token, None)
|
||||
assert assertion is not None
|
||||
assert assertion.expires_at is None
|
||||
assert assertion.issuer is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retention_gate_requires_an_id_jag_server():
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
|
||||
patch("litellm.proxy.proxy_server.prisma_client", _make_prisma({}, db_has_id_jag_server=False)),
|
||||
):
|
||||
manager.config_mcp_servers = {
|
||||
"s1": _server_with_auth(MCPAuth.oauth2),
|
||||
"s2": _server_with_auth(None),
|
||||
}
|
||||
assert await ema_assertion_retention_enabled() is False
|
||||
manager.config_mcp_servers = {
|
||||
"s1": _server_with_auth(MCPAuth.oauth2),
|
||||
"s2": _server_with_auth(MCPAuth.oauth2_id_jag),
|
||||
}
|
||||
assert await ema_assertion_retention_enabled() is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retention_gate_reads_the_db_when_config_declares_no_id_jag_server():
|
||||
"""A DB-backed server added on another pod (or before this pod's DB load) must still enable
|
||||
retention off the authoritative DB row; False only when neither authority knows one."""
|
||||
with patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager:
|
||||
manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2)}
|
||||
db_backed = _make_prisma({}, db_has_id_jag_server=True)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", db_backed):
|
||||
assert await ema_assertion_retention_enabled() is True
|
||||
db_backed.db.litellm_mcpservertable.find_first.assert_awaited_once_with(
|
||||
where={"auth_type": MCPAuth.oauth2_id_jag.value}
|
||||
)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", None):
|
||||
assert await ema_assertion_retention_enabled() is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retention_gate_never_consults_the_registry_snapshot():
|
||||
"""The registry is a per-process snapshot of DB state, stale in either direction: trusting
|
||||
it positively would keep retaining bearer material after the last EMA server was removed on
|
||||
another pod, trusting it negatively would drop writes for one added elsewhere. The gate must
|
||||
judge only the config declaration and the DB row, so a stale snapshot listing an id_jag
|
||||
server changes nothing."""
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
|
||||
patch("litellm.proxy.proxy_server.prisma_client", _make_prisma({}, db_has_id_jag_server=False)),
|
||||
):
|
||||
manager.config_mcp_servers = {}
|
||||
manager.get_registry.return_value = {"stale": _server_with_auth(MCPAuth.oauth2_id_jag)}
|
||||
assert await ema_assertion_retention_enabled() is False
|
||||
manager.get_registry.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retain_persists_when_only_the_db_knows_the_id_jag_server():
|
||||
stored = {}
|
||||
prisma = _make_prisma(stored, db_has_id_jag_server=True)
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma),
|
||||
):
|
||||
manager.config_mcp_servers = {}
|
||||
await retain_sso_identity_assertion_for_ema(
|
||||
user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None)
|
||||
)
|
||||
assert "user-a" in stored
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_persist_and_fetch_round_trip_encrypted_at_rest():
|
||||
stored = {}
|
||||
prisma = _make_prisma(stored)
|
||||
token = _make_id_token()
|
||||
assertion = assertion_from_sso_login(token, "rt_1")
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
await persist_sso_identity_assertion("user-a", assertion)
|
||||
fetched = await fetch_sso_identity_assertion("user-a")
|
||||
assert fetched is not None
|
||||
assert fetched.id_token.get_secret_value() == token
|
||||
assert fetched.refresh_token is not None
|
||||
assert fetched.refresh_token.get_secret_value() == "rt_1"
|
||||
assert fetched.issuer == assertion.issuer
|
||||
assert fetched.expires_at == assertion.expires_at
|
||||
assert token not in stored["user-a"]
|
||||
assert "rt_1" not in stored["user-a"]
|
||||
decrypted = decrypt_value_helper(stored["user-a"], "test", exception_type="debug")
|
||||
assert json.loads(decrypted)["id_token"] == token
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_persist_overwrites_previous_login():
|
||||
stored = {}
|
||||
prisma = _make_prisma(stored)
|
||||
first = _make_id_token(exp_offset=100)
|
||||
second = _make_id_token(exp_offset=7200)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
await persist_sso_identity_assertion("user-a", assertion_from_sso_login(first, None))
|
||||
await persist_sso_identity_assertion("user-a", assertion_from_sso_login(second, "rt_new"))
|
||||
fetched = await fetch_sso_identity_assertion("user-a")
|
||||
assert fetched is not None
|
||||
assert fetched.id_token.get_secret_value() == second
|
||||
assert fetched.refresh_token is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_missing_row_returns_none():
|
||||
prisma = _make_prisma({})
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
assert await fetch_sso_identity_assertion("nobody") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_undecryptable_row_returns_none():
|
||||
prisma = _make_prisma({"user-a": "not-an-encrypted-blob"})
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
assert await fetch_sso_identity_assertion("user-a") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_unparseable_payload_returns_none():
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
|
||||
|
||||
prisma = _make_prisma({"user-a": encrypt_value_helper("]]not json")})
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
assert await fetch_sso_identity_assertion("user-a") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retain_noop_when_no_id_jag_server():
|
||||
stored = {}
|
||||
prisma = _make_prisma(stored)
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma),
|
||||
):
|
||||
manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2)}
|
||||
await retain_sso_identity_assertion_for_ema(
|
||||
user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None)
|
||||
)
|
||||
prisma.db.litellm_ssoidentityassertion.upsert.assert_not_called()
|
||||
assert stored == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retain_persists_when_id_jag_server_registered():
|
||||
stored = {}
|
||||
prisma = _make_prisma(stored)
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma),
|
||||
):
|
||||
manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2_id_jag)}
|
||||
await retain_sso_identity_assertion_for_ema(
|
||||
user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None)
|
||||
)
|
||||
assert "user-a" in stored
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retain_none_assertion_never_consults_gate_or_store():
|
||||
gate = MagicMock()
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store.ema_assertion_retention_enabled",
|
||||
gate,
|
||||
):
|
||||
await retain_sso_identity_assertion_for_ema(user_id="user-a", assertion=None)
|
||||
gate.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retain_swallows_store_failure():
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_ssoidentityassertion.upsert = AsyncMock(side_effect=RuntimeError("db down"))
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma),
|
||||
):
|
||||
manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2_id_jag)}
|
||||
await retain_sso_identity_assertion_for_ema(
|
||||
user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rotation_reencrypts_under_new_key(monkeypatch):
|
||||
stored = {}
|
||||
prisma = _make_prisma(stored)
|
||||
token = _make_id_token()
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
await persist_sso_identity_assertion("user-a", assertion_from_sso_login(token, None))
|
||||
original_blob = stored["user-a"]
|
||||
|
||||
new_key = "rotated-sso-assertion-salt-key-5678"
|
||||
await rotate_sso_identity_assertions_master_key(prisma_client=prisma, new_master_key=new_key)
|
||||
assert stored["user-a"] != original_blob
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", new_key)
|
||||
decrypted = decrypt_value_helper(stored["user-a"], "test", exception_type="debug")
|
||||
assert decrypted is not None
|
||||
assert json.loads(decrypted)["id_token"] == token
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rotation_skips_unreadable_rows_but_rotates_readable_ones():
|
||||
stored = {"good": None, "bad": "garbage-blob"}
|
||||
prisma = _make_prisma(stored)
|
||||
token = _make_id_token()
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
await persist_sso_identity_assertion("good", assertion_from_sso_login(token, None))
|
||||
good_blob_before = stored["good"]
|
||||
await rotate_sso_identity_assertions_master_key(prisma_client=prisma, new_master_key="another-new-salt-key-0000")
|
||||
assert stored["bad"] == "garbage-blob"
|
||||
assert stored["good"] != good_blob_before
|
||||
|
|
@ -14994,3 +14994,76 @@ async def test_list_keys_without_expires_param_forwards_none():
|
|||
|
||||
mock_helper.assert_called_once()
|
||||
assert mock_helper.call_args.kwargs["expires_filter"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.rotate_sso_identity_assertions_master_key"
|
||||
)
|
||||
@patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_user_env_vars_master_key"
|
||||
)
|
||||
@patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_user_credentials_master_key"
|
||||
)
|
||||
@patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_server_credentials_master_key"
|
||||
)
|
||||
async def test_rotate_master_key_rotates_sso_identity_assertions(
|
||||
mock_rotate_mcp_server,
|
||||
mock_rotate_mcp_user,
|
||||
mock_rotate_env_vars,
|
||||
mock_rotate_sso,
|
||||
):
|
||||
"""Master-key rotation must re-encrypt the SSO identity assertion store alongside
|
||||
the sibling per-user encrypted tables, or a salt rotation orphans every stored
|
||||
assertion (step 4d)."""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
_rotate_master_key,
|
||||
)
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db = MagicMock()
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
|
||||
mock_tx = AsyncMock()
|
||||
mock_tx.litellm_proxymodeltable = MagicMock()
|
||||
mock_tx.litellm_proxymodeltable.delete_many = AsyncMock()
|
||||
mock_tx.litellm_proxymodeltable.create_many = AsyncMock()
|
||||
mock_prisma_client.db.tx = MagicMock(
|
||||
return_value=AsyncMock(
|
||||
__aenter__=AsyncMock(return_value=mock_tx),
|
||||
__aexit__=AsyncMock(return_value=False),
|
||||
)
|
||||
)
|
||||
mock_prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(
|
||||
return_value=[]
|
||||
)
|
||||
|
||||
mock_proxy_config = MagicMock()
|
||||
mock_proxy_config.decrypt_model_list_from_db.return_value = []
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="test-user",
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.proxy_config",
|
||||
mock_proxy_config,
|
||||
):
|
||||
await _rotate_master_key(
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
current_master_key="sk-old-master-key",
|
||||
new_master_key="sk-new-master-key",
|
||||
)
|
||||
|
||||
mock_rotate_sso.assert_awaited_once_with(
|
||||
prisma_client=mock_prisma_client,
|
||||
new_master_key="sk-new-master-key",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1458,7 +1458,7 @@ async def test_get_generic_sso_response_with_additional_headers():
|
|||
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
|
||||
):
|
||||
# Act
|
||||
result, received_response, _ = await get_generic_sso_response(
|
||||
result, received_response, _, _ = await get_generic_sso_response(
|
||||
request=mock_request,
|
||||
jwt_handler=mock_jwt_handler,
|
||||
generic_client_id=generic_client_id,
|
||||
|
|
@ -1522,7 +1522,7 @@ async def test_get_generic_sso_response_with_empty_headers():
|
|||
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
|
||||
):
|
||||
# Act
|
||||
result, received_response, _ = await get_generic_sso_response(
|
||||
result, received_response, _, _ = await get_generic_sso_response(
|
||||
request=mock_request,
|
||||
jwt_handler=mock_jwt_handler,
|
||||
generic_client_id=generic_client_id,
|
||||
|
|
@ -2893,6 +2893,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
prefill_user_code=None,
|
||||
result=mock_result,
|
||||
received_response=None,
|
||||
sso_assertion=None,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -2933,6 +2934,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
prefill_user_code="WXYZ-2345",
|
||||
result=mock_result,
|
||||
received_response=None,
|
||||
sso_assertion=None,
|
||||
)
|
||||
|
||||
def test_get_redirect_url_does_not_include_existing_key_in_url(self):
|
||||
|
|
@ -7019,7 +7021,7 @@ class TestPKCEStateCookieBinding:
|
|||
):
|
||||
jwt_handler = MagicMock(spec=JWTHandler)
|
||||
jwt_handler.get_team_ids_from_jwt.return_value = []
|
||||
result, _, _ = await get_generic_sso_response(
|
||||
result, _, _, _ = await get_generic_sso_response(
|
||||
request=mock_request,
|
||||
jwt_handler=jwt_handler,
|
||||
generic_client_id="cid",
|
||||
|
|
@ -7078,7 +7080,7 @@ async def test_debug_sso_callback_renders_full_jwt_claims():
|
|||
}
|
||||
|
||||
async def fake_get_generic_sso_response(**kwargs):
|
||||
return parsed_openid, raw_userinfo_with_leaked_token, access_token_payload
|
||||
return parsed_openid, raw_userinfo_with_leaked_token, access_token_payload, None
|
||||
|
||||
with (
|
||||
patch.dict(
|
||||
|
|
@ -7374,3 +7376,266 @@ async def test_auth_callback_without_oauth_error_proceeds_to_normal_flow():
|
|||
|
||||
assert exc_info.value.status_code == 500
|
||||
assert "DB not connected" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
# ── SSO identity assertion capture + persist wiring (EMA) ─────────────────────
|
||||
|
||||
|
||||
def _ema_id_token(sub: str = "u1") -> str:
|
||||
import time as _time
|
||||
|
||||
import jwt as _pyjwt
|
||||
|
||||
return _pyjwt.encode(
|
||||
{"iss": "https://idp.example.com", "sub": sub, "exp": int(_time.time()) + 3600},
|
||||
"test-idp-signing-key-32-bytes-long-xxxx",
|
||||
algorithm="HS256",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_arm_captures_sso_assertion():
|
||||
"""The PKCE token exchange strips bearer fields from received_response for safety;
|
||||
the typed assertion carrier must still capture id_token + refresh_token."""
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
SSOAuthenticationHandler,
|
||||
get_generic_sso_response,
|
||||
)
|
||||
|
||||
id_token = _ema_id_token()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.query_params = {"state": "matched-state", "code": "auth-code"}
|
||||
mock_request.cookies = {"litellm_oauth_state": "matched-state"}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
SSOAuthenticationHandler,
|
||||
"prepare_token_exchange_parameters",
|
||||
AsyncMock(
|
||||
return_value={
|
||||
"code_verifier": "verifier",
|
||||
"_pkce_cache_key": "pkce_verifier:matched-state",
|
||||
}
|
||||
),
|
||||
),
|
||||
patch.object(
|
||||
SSOAuthenticationHandler,
|
||||
"_pkce_token_exchange",
|
||||
AsyncMock(
|
||||
return_value={
|
||||
"access_token": "tok",
|
||||
"id_token": id_token,
|
||||
"refresh_token": "rt_from_idp",
|
||||
"sub": "user@example.com",
|
||||
"email": "user@example.com",
|
||||
}
|
||||
),
|
||||
),
|
||||
patch.object(SSOAuthenticationHandler, "_delete_pkce_verifier", AsyncMock()),
|
||||
patch("fastapi_sso.sso.base.DiscoveryDocument"),
|
||||
patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()),
|
||||
patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"GENERIC_CLIENT_SECRET": "x",
|
||||
"GENERIC_AUTHORIZATION_ENDPOINT": "https://idp.example.com/auth",
|
||||
"GENERIC_TOKEN_ENDPOINT": "https://idp.example.com/token",
|
||||
"GENERIC_USERINFO_ENDPOINT": "https://idp.example.com/userinfo",
|
||||
"GENERIC_CLIENT_USE_PKCE": "true",
|
||||
},
|
||||
),
|
||||
):
|
||||
jwt_handler = MagicMock(spec=JWTHandler)
|
||||
jwt_handler.get_team_ids_from_jwt.return_value = []
|
||||
result, received_response, _, sso_assertion = await get_generic_sso_response(
|
||||
request=mock_request,
|
||||
jwt_handler=jwt_handler,
|
||||
generic_client_id="cid",
|
||||
redirect_url="https://proxy.example.com/sso/callback",
|
||||
sso_jwt_handler=None,
|
||||
)
|
||||
|
||||
assert sso_assertion is not None
|
||||
assert sso_assertion.id_token.get_secret_value() == id_token
|
||||
assert sso_assertion.refresh_token is not None
|
||||
assert sso_assertion.refresh_token.get_secret_value() == "rt_from_idp"
|
||||
# The sanitized received_response must still not carry bearer material.
|
||||
assert "id_token" not in (received_response or {})
|
||||
assert "refresh_token" not in (received_response or {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_verify_and_process_arm_captures_sso_assertion():
|
||||
"""The non-PKCE generic arm reads the raw bearer fields off the fastapi-sso client."""
|
||||
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
|
||||
|
||||
id_token = _ema_id_token()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_jwt_handler = MagicMock(spec=JWTHandler)
|
||||
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
|
||||
|
||||
mock_sso_instance = MagicMock()
|
||||
mock_sso_instance.verify_and_process = AsyncMock(
|
||||
return_value={"sub": "u1", "email": "u@example.com"}
|
||||
)
|
||||
mock_sso_instance.access_token = None
|
||||
mock_sso_instance.id_token = id_token
|
||||
mock_sso_instance.refresh_token = "rt_from_idp"
|
||||
mock_sso_class = MagicMock(return_value=mock_sso_instance)
|
||||
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"GENERIC_CLIENT_SECRET": "test_secret",
|
||||
"GENERIC_AUTHORIZATION_ENDPOINT": "https://auth.example.com/auth",
|
||||
"GENERIC_TOKEN_ENDPOINT": "https://auth.example.com/token",
|
||||
"GENERIC_USERINFO_ENDPOINT": "https://auth.example.com/userinfo",
|
||||
},
|
||||
):
|
||||
with patch("fastapi_sso.sso.base.DiscoveryDocument"):
|
||||
with patch(
|
||||
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
|
||||
):
|
||||
_, _, _, sso_assertion = await get_generic_sso_response(
|
||||
request=mock_request,
|
||||
jwt_handler=mock_jwt_handler,
|
||||
generic_client_id="test_client_id",
|
||||
redirect_url="http://test.com/callback",
|
||||
sso_jwt_handler=None,
|
||||
)
|
||||
|
||||
assert sso_assertion is not None
|
||||
assert sso_assertion.id_token.get_secret_value() == id_token
|
||||
assert sso_assertion.refresh_token is not None
|
||||
assert sso_assertion.refresh_token.get_secret_value() == "rt_from_idp"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redirect_from_openid_persists_assertion_under_canonical_user_id():
|
||||
"""The browser funnel persists the captured assertion AFTER canonical user
|
||||
resolution, keyed by the user_id admission will later resolve (the key-generation
|
||||
response user_id), not the raw IdP subject."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
|
||||
assertion_from_sso_login,
|
||||
)
|
||||
|
||||
assertion = assertion_from_sso_login(_ema_id_token(), "rt_1")
|
||||
assert assertion is not None
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "http://localhost:4000/"
|
||||
mock_request.cookies = {}
|
||||
|
||||
retain_mock = AsyncMock()
|
||||
with (
|
||||
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
|
||||
patch("litellm.proxy.proxy_server.general_settings", {}),
|
||||
patch("litellm.proxy.proxy_server.premium_user", False),
|
||||
patch("litellm.proxy.proxy_server.user_custom_sso", None),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.generate_key_helper_fn",
|
||||
AsyncMock(
|
||||
return_value={"token": "sk-ui-key", "user_id": "canonical-user-id"}
|
||||
),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
|
||||
AsyncMock(return_value=None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.ui_sso.check_and_update_if_proxy_admin_id",
|
||||
AsyncMock(return_value="internal_user"),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
|
||||
retain_mock,
|
||||
),
|
||||
):
|
||||
response = await SSOAuthenticationHandler.get_redirect_response_from_openid(
|
||||
result=CustomOpenID(
|
||||
id="raw-idp-subject",
|
||||
email="u@example.com",
|
||||
first_name="U",
|
||||
last_name="Ser",
|
||||
display_name="U Ser",
|
||||
provider="generic",
|
||||
team_ids=[],
|
||||
user_role=None,
|
||||
),
|
||||
request=mock_request,
|
||||
received_response=None,
|
||||
generic_client_id="cid",
|
||||
ui_access_mode=None,
|
||||
access_token_payload=None,
|
||||
jwt_handler=None,
|
||||
sso_assertion=assertion,
|
||||
)
|
||||
|
||||
retain_mock.assert_awaited_once_with(
|
||||
user_id="canonical-user-id", assertion=assertion
|
||||
)
|
||||
assert response is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cli_completion_persists_assertion_under_db_user_id():
|
||||
"""The CLI funnel persists the captured assertion under the DB-resolved user_id."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
|
||||
assertion_from_sso_login,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
_complete_cli_sso_callback_session,
|
||||
)
|
||||
|
||||
assertion = assertion_from_sso_login(_ema_id_token(), None)
|
||||
assert assertion is not None
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "http://localhost:4000/"
|
||||
|
||||
user_info = MagicMock()
|
||||
user_info.user_id = "cli-user-id"
|
||||
user_info.user_role = "internal_user"
|
||||
user_info.models = []
|
||||
user_info.teams = []
|
||||
|
||||
retain_mock = AsyncMock()
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
|
||||
AsyncMock(return_value=user_info),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.ui_sso._fetch_cli_sso_team_details",
|
||||
AsyncMock(return_value=[]),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.ui_sso.build_cli_sso_attribution_metadata",
|
||||
return_value={},
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
|
||||
retain_mock,
|
||||
),
|
||||
):
|
||||
response = await _complete_cli_sso_callback_session(
|
||||
request=mock_request,
|
||||
key="cli-login-id",
|
||||
flow={},
|
||||
result={"sub": "raw-idp-subject"},
|
||||
parsed_openid_result={
|
||||
"user_id": "raw-idp-subject",
|
||||
"user_email": "u@example.com",
|
||||
"user_role": None,
|
||||
},
|
||||
user_defined_values=None,
|
||||
prisma_client=MagicMock(),
|
||||
user_api_key_cache=MagicMock(),
|
||||
proxy_logging_obj=MagicMock(),
|
||||
sso_assertion=assertion,
|
||||
)
|
||||
|
||||
retain_mock.assert_awaited_once_with(user_id="cli-user-id", assertion=assertion)
|
||||
assert response.status_code == 200
|
||||
|
|
|
|||
|
|
@ -1069,6 +1069,32 @@ async def test_ProxyConfig_load_config_wires_general_settings_url_validation(tmp
|
|||
litellm.provider_url_destination_allowed_hosts = original_provider_hosts
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig_load_config_wires_config_reload_interval(tmp_path, monkeypatch):
|
||||
"""general_settings.proxy_config_reload_interval_seconds must reach the proxy_server
|
||||
module global that schedules the DB config-reload jobs, so operators can tune multi-pod
|
||||
convergence from config.yaml."""
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
f = tmp_path / "c.yaml"
|
||||
f.write_text(
|
||||
"model_list: []\n"
|
||||
"general_settings:\n"
|
||||
" proxy_config_reload_interval_seconds: 47\n"
|
||||
"litellm_settings: {}\n"
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
|
||||
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
|
||||
|
||||
original = proxy_server.proxy_config_reload_interval_seconds
|
||||
try:
|
||||
await ProxyConfig().load_config(router=None, config_file_path=str(f))
|
||||
assert proxy_server.proxy_config_reload_interval_seconds == 47
|
||||
finally:
|
||||
proxy_server.proxy_config_reload_interval_seconds = original
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig_load_config_missing_file_raises(monkeypatch):
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ Routes covered:
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from .conftest import VOLATILE_KEYS, normalize
|
||||
|
|
@ -473,6 +474,83 @@ def test_config_list_happy_admin(client, auth_as, mock_prisma, monkeypatch):
|
|||
}
|
||||
|
||||
|
||||
def test_config_list_exposes_config_reload_interval(client, auth_as, mock_prisma, monkeypatch):
|
||||
"""proxy_config_reload_interval_seconds must surface in the admin UI general-settings
|
||||
list as an Integer field defaulting to 30, so operators can tune multi-pod convergence
|
||||
from the dashboard."""
|
||||
from litellm.proxy import proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
table = _install_litellm_config(mock_prisma)
|
||||
row = MagicMock()
|
||||
row.param_value = {}
|
||||
table.find_first = AsyncMock(return_value=row)
|
||||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||||
|
||||
with auth_as(LitellmUserRoles.PROXY_ADMIN):
|
||||
response = client.get("/config/list", params={"config_type": "general_settings"})
|
||||
assert response.status_code == 200
|
||||
by_name = {entry["field_name"]: entry for entry in response.json()}
|
||||
assert "proxy_config_reload_interval_seconds" in by_name
|
||||
entry = by_name["proxy_config_reload_interval_seconds"]
|
||||
assert entry["field_type"] == "Integer"
|
||||
assert entry["field_default_value"] == 30
|
||||
|
||||
|
||||
def test_config_field_update_accepts_config_reload_interval(client, auth_as, mock_prisma, monkeypatch):
|
||||
"""POST /config/field/update accepts proxy_config_reload_interval_seconds and persists
|
||||
it to the DB general_settings row for all pods to pick up."""
|
||||
from litellm.proxy import proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
table = _install_litellm_config(mock_prisma)
|
||||
table.find_first = AsyncMock(return_value=None)
|
||||
upsert_row = {
|
||||
"param_name": "general_settings",
|
||||
"param_value": {"proxy_config_reload_interval_seconds": 45},
|
||||
"id": "row-1",
|
||||
}
|
||||
table.upsert = AsyncMock(return_value=upsert_row)
|
||||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||||
|
||||
with auth_as(LitellmUserRoles.PROXY_ADMIN):
|
||||
response = client.post(
|
||||
"/config/field/update",
|
||||
json={
|
||||
"field_name": "proxy_config_reload_interval_seconds",
|
||||
"field_value": 45,
|
||||
"config_type": "general_settings",
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
upserted = table.upsert.call_args.kwargs["data"]["create"]["param_value"]
|
||||
assert json.loads(upserted)["proxy_config_reload_interval_seconds"] == 45
|
||||
|
||||
|
||||
def test_config_field_update_rejects_non_positive_config_reload_interval(client, auth_as, mock_prisma, monkeypatch):
|
||||
"""A non-positive proxy_config_reload_interval_seconds from the UI is rejected with a 400
|
||||
and never persisted, since APScheduler requires a positive interval."""
|
||||
from litellm.proxy import proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
table = _install_litellm_config(mock_prisma)
|
||||
table.find_first = AsyncMock(return_value=None)
|
||||
table.upsert = AsyncMock()
|
||||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||||
|
||||
with auth_as(LitellmUserRoles.PROXY_ADMIN):
|
||||
response = client.post(
|
||||
"/config/field/update",
|
||||
json={
|
||||
"field_name": "proxy_config_reload_interval_seconds",
|
||||
"field_value": 0,
|
||||
"config_type": "general_settings",
|
||||
},
|
||||
)
|
||||
assert response.status_code == 400
|
||||
table.upsert.assert_not_called()
|
||||
|
||||
|
||||
def test_config_list_non_admin_rejected(client, auth_as, mock_prisma, monkeypatch):
|
||||
"""Non-admin gets a 400 with the role embedded in the error message."""
|
||||
from litellm.proxy import proxy_server as ps
|
||||
|
|
|
|||
|
|
@ -751,6 +751,97 @@ async def test_initialize_scheduled_jobs_credentials(monkeypatch):
|
|||
assert len(mock_scheduler_calls) > 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_scheduled_jobs_uses_configured_config_reload_interval(monkeypatch):
|
||||
"""
|
||||
The DB config-reload jobs (add_deployment, get_credentials) that keep multi-pod
|
||||
deployments in sync must be scheduled at the configured
|
||||
proxy_config_reload_interval_seconds, not a hardcoded value.
|
||||
"""
|
||||
monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False)
|
||||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||||
mock_proxy_config = AsyncMock()
|
||||
mock_scheduler = MagicMock()
|
||||
|
||||
configured_interval = 47
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||||
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True),
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.proxy_config_reload_interval_seconds",
|
||||
configured_interval,
|
||||
),
|
||||
patch("litellm.proxy.proxy_server.AsyncIOScheduler", return_value=mock_scheduler),
|
||||
):
|
||||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||||
general_settings={},
|
||||
prisma_client=mock_prisma_client,
|
||||
proxy_budget_rescheduler_min_time=1,
|
||||
proxy_budget_rescheduler_max_time=2,
|
||||
proxy_batch_write_at=5,
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
)
|
||||
|
||||
scheduled_seconds = {
|
||||
job_call.kwargs["id"]: job_call.kwargs.get("seconds")
|
||||
for job_call in mock_scheduler.add_job.call_args_list
|
||||
if "id" in job_call.kwargs
|
||||
}
|
||||
assert scheduled_seconds["add_deployment_job"] == configured_interval
|
||||
assert scheduled_seconds["get_credentials_job"] == configured_interval
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_scheduled_jobs_rejects_non_positive_config_reload_interval(monkeypatch):
|
||||
"""
|
||||
A non-positive proxy_config_reload_interval_seconds (misconfig via env/config/DB) would
|
||||
make APScheduler reject the job and crash startup, so the scheduler must fall back to the
|
||||
30s default instead of forwarding the bad value.
|
||||
"""
|
||||
monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False)
|
||||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||||
mock_proxy_config = AsyncMock()
|
||||
mock_scheduler = MagicMock()
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||||
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True),
|
||||
patch("litellm.proxy.proxy_server.proxy_config_reload_interval_seconds", 0),
|
||||
patch("litellm.proxy.proxy_server.AsyncIOScheduler", return_value=mock_scheduler),
|
||||
):
|
||||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||||
general_settings={},
|
||||
prisma_client=mock_prisma_client,
|
||||
proxy_budget_rescheduler_min_time=1,
|
||||
proxy_budget_rescheduler_max_time=2,
|
||||
proxy_batch_write_at=5,
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
)
|
||||
|
||||
scheduled_seconds = {
|
||||
job_call.kwargs["id"]: job_call.kwargs.get("seconds")
|
||||
for job_call in mock_scheduler.add_job.call_args_list
|
||||
if "id" in job_call.kwargs
|
||||
}
|
||||
assert scheduled_seconds["add_deployment_job"] == 30
|
||||
assert scheduled_seconds["get_credentials_job"] == 30
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_scheduled_jobs_hydrates_mcp_when_store_model_in_db_false(monkeypatch):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ import {
|
|||
type OnChangeFn,
|
||||
type Row,
|
||||
type RowData,
|
||||
type RowSelectionState,
|
||||
type Table,
|
||||
type TableOptions,
|
||||
useReactTable,
|
||||
|
|
@ -70,6 +71,8 @@ export function validateDataTableConfig<TData extends RowData, TValue>(
|
|||
const bothSortingSources = props.defaultSorting !== undefined && props.sorting !== undefined;
|
||||
const bothFilterSources = props.defaultColumnFilters !== undefined && props.columnFilters !== undefined;
|
||||
|
||||
const controlledSelectionIncomplete = props.rowSelection !== undefined && props.onRowSelectionChange === undefined;
|
||||
|
||||
return [
|
||||
serverSortingIncomplete ? "sortingMode='server' requires both `sorting` and `onSortingChange`." : null,
|
||||
serverPaginationIncomplete
|
||||
|
|
@ -80,6 +83,9 @@ export function validateDataTableConfig<TData extends RowData, TValue>(
|
|||
bothFilterSources
|
||||
? "Provide either `defaultColumnFilters` (uncontrolled) or `columnFilters` (controlled), not both."
|
||||
: null,
|
||||
controlledSelectionIncomplete
|
||||
? "Controlled `rowSelection` requires `onRowSelectionChange`; without it selection changes are dropped."
|
||||
: null,
|
||||
].filter((message): message is string => message !== null);
|
||||
}
|
||||
|
||||
|
|
@ -448,6 +454,9 @@ function useDataTableInstance<TData extends RowData, TValue>(props: DataTablePro
|
|||
renderSubComponent,
|
||||
expanded,
|
||||
onExpandedChange,
|
||||
enableRowSelection,
|
||||
rowSelection,
|
||||
onRowSelectionChange,
|
||||
} = props;
|
||||
|
||||
const sortingState = useControllable(sorting, onSortingChange, defaultSorting ?? []);
|
||||
|
|
@ -462,6 +471,7 @@ function useDataTableInstance<TData extends RowData, TValue>(props: DataTablePro
|
|||
);
|
||||
const globalFilterState = useControllable<string>(globalFilter, onGlobalFilterChange, "");
|
||||
const expandedState = useControllable<ExpandedState>(expanded, onExpandedChange, {});
|
||||
const rowSelectionState = useControllable<RowSelectionState>(rowSelection, onRowSelectionChange, {});
|
||||
const [columnVisibility, setColumnVisibility] = useState<VisibilityState>(defaultColumnVisibility ?? {});
|
||||
const [columnSizing, setColumnSizing] = useState<ColumnSizingState>({});
|
||||
const columnPinning = React.useMemo(() => derivePinning(columns), [columns]);
|
||||
|
|
@ -476,6 +486,7 @@ function useDataTableInstance<TData extends RowData, TValue>(props: DataTablePro
|
|||
columnFilters: filterState.value,
|
||||
globalFilter: globalFilterState.value,
|
||||
expanded: expandedState.value,
|
||||
rowSelection: rowSelectionState.value,
|
||||
columnVisibility,
|
||||
columnSizing,
|
||||
},
|
||||
|
|
@ -491,11 +502,13 @@ function useDataTableInstance<TData extends RowData, TValue>(props: DataTablePro
|
|||
onColumnFiltersChange: filterState.onChange,
|
||||
onGlobalFilterChange: globalFilterState.onChange,
|
||||
onExpandedChange: expandedState.onChange,
|
||||
onRowSelectionChange: rowSelectionState.onChange,
|
||||
onColumnVisibilityChange: setColumnVisibility,
|
||||
onColumnSizingChange: setColumnSizing,
|
||||
getCoreRowModel: getCoreRowModel(),
|
||||
...buildRowModels(sortingMode, paginationMode, filterMode, expansionGuard),
|
||||
...(getRowId !== undefined ? { getRowId } : {}),
|
||||
...(enableRowSelection !== undefined ? { enableRowSelection } : {}),
|
||||
...(paginationMode === "server" && rowCount !== undefined ? { rowCount } : {}),
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,137 @@
|
|||
import type { ColumnDef, RowSelectionState } from "@tanstack/react-table";
|
||||
import { render, screen } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { useState } from "react";
|
||||
import { describe, expect, it } from "vitest";
|
||||
|
||||
import { createSelectionColumn, DataTable, validateDataTableConfig } from "./index";
|
||||
|
||||
interface Model {
|
||||
id: string;
|
||||
name: string;
|
||||
}
|
||||
|
||||
const data: Model[] = [
|
||||
{ id: "m1", name: "Alpha" },
|
||||
{ id: "m2", name: "Beta" },
|
||||
{ id: "m3", name: "Gamma" },
|
||||
];
|
||||
|
||||
const columns: ColumnDef<Model, unknown>[] = [
|
||||
createSelectionColumn<Model>({ rowAriaLabel: (row) => `Select ${row.original.name}` }),
|
||||
{ id: "name", accessorKey: "name", header: "Name", enableSorting: false },
|
||||
];
|
||||
|
||||
const selectAll = () => screen.getByTestId("datatable-select-all");
|
||||
const rowBox = (id: string) => screen.getByTestId(`datatable-select-row-${id}`);
|
||||
const selectedCount = () => screen.getByTestId("count");
|
||||
|
||||
function ControlledHarness() {
|
||||
const [rowSelection, setRowSelection] = useState<RowSelectionState>({});
|
||||
|
||||
return (
|
||||
<>
|
||||
<span data-testid="keys">
|
||||
{Object.keys(rowSelection)
|
||||
.filter((key) => rowSelection[key])
|
||||
.sort()
|
||||
.join(",")}
|
||||
</span>
|
||||
<button type="button" data-testid="clear" onClick={() => setRowSelection({})}>
|
||||
clear
|
||||
</button>
|
||||
<DataTable
|
||||
data={data}
|
||||
columns={columns}
|
||||
getRowId={(row) => row.id}
|
||||
rowSelection={rowSelection}
|
||||
onRowSelectionChange={setRowSelection}
|
||||
/>
|
||||
</>
|
||||
);
|
||||
}
|
||||
|
||||
describe("DataTable row selection", () => {
|
||||
it("supports uncontrolled per-row toggle, select-all, and indeterminate", async () => {
|
||||
const user = userEvent.setup();
|
||||
|
||||
render(
|
||||
<DataTable
|
||||
data={data}
|
||||
columns={columns}
|
||||
getRowId={(row) => row.id}
|
||||
toolbar={(table) => <span data-testid="count">{table.getSelectedRowModel().rows.length}</span>}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(selectedCount()).toHaveTextContent("0");
|
||||
|
||||
await user.click(rowBox("m1"));
|
||||
expect(selectedCount()).toHaveTextContent("1");
|
||||
expect(selectAll()).toHaveAttribute("aria-checked", "mixed");
|
||||
|
||||
await user.click(selectAll());
|
||||
expect(selectedCount()).toHaveTextContent("3");
|
||||
expect(selectAll()).toHaveAttribute("aria-checked", "true");
|
||||
|
||||
await user.click(selectAll());
|
||||
expect(selectedCount()).toHaveTextContent("0");
|
||||
});
|
||||
|
||||
it("keys controlled selection by getRowId so the parent can map back to entities", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<ControlledHarness />);
|
||||
|
||||
await user.click(rowBox("m2"));
|
||||
expect(screen.getByTestId("keys")).toHaveTextContent("m2");
|
||||
|
||||
await user.click(rowBox("m3"));
|
||||
expect(screen.getByTestId("keys")).toHaveTextContent("m2,m3");
|
||||
});
|
||||
|
||||
it("lets the parent clear the selection, the pattern an external pager needs", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<ControlledHarness />);
|
||||
|
||||
await user.click(selectAll());
|
||||
expect(screen.getByTestId("keys")).toHaveTextContent("m1,m2,m3");
|
||||
|
||||
await user.click(screen.getByTestId("clear"));
|
||||
expect(screen.getByTestId("keys")).toBeEmptyDOMElement();
|
||||
expect(rowBox("m1")).toHaveAttribute("aria-checked", "false");
|
||||
});
|
||||
|
||||
it("respects an enableRowSelection predicate", async () => {
|
||||
const user = userEvent.setup();
|
||||
|
||||
render(
|
||||
<DataTable
|
||||
data={data}
|
||||
columns={columns}
|
||||
getRowId={(row) => row.id}
|
||||
enableRowSelection={(row) => row.original.id !== "m2"}
|
||||
toolbar={(table) => <span data-testid="count">{table.getSelectedRowModel().rows.length}</span>}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(rowBox("m2")).toHaveAttribute("aria-disabled", "true");
|
||||
|
||||
await user.click(rowBox("m2"));
|
||||
expect(selectedCount()).toHaveTextContent("0");
|
||||
|
||||
await user.click(rowBox("m1"));
|
||||
expect(selectedCount()).toHaveTextContent("1");
|
||||
});
|
||||
|
||||
it("rejects controlled rowSelection without onRowSelectionChange", () => {
|
||||
const errors = validateDataTableConfig<Model, unknown>({ data, columns, rowSelection: { m1: true } });
|
||||
|
||||
expect(errors).toContain(
|
||||
"Controlled `rowSelection` requires `onRowSelectionChange`; without it selection changes are dropped.",
|
||||
);
|
||||
});
|
||||
|
||||
it("does not complain when selection is left uncontrolled", () => {
|
||||
expect(validateDataTableConfig<Model, unknown>({ data, columns })).toHaveLength(0);
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,53 @@
|
|||
"use client";
|
||||
|
||||
import type { ColumnDef, Row, RowData, Table } from "@tanstack/react-table";
|
||||
|
||||
import { Checkbox } from "@/components/ui/checkbox";
|
||||
|
||||
interface SelectionColumnOptions<TData> {
|
||||
rowAriaLabel?: (row: Row<TData>) => string;
|
||||
}
|
||||
|
||||
function SelectAllCheckbox<TData>({ table }: { table: Table<TData> }) {
|
||||
const allSelected = table.getIsAllPageRowsSelected();
|
||||
const someSelected = table.getIsSomePageRowsSelected();
|
||||
|
||||
return (
|
||||
<Checkbox
|
||||
aria-label="Select all rows"
|
||||
data-testid="datatable-select-all"
|
||||
checked={allSelected}
|
||||
indeterminate={someSelected && !allSelected}
|
||||
onCheckedChange={(checked) => table.toggleAllPageRowsSelected(Boolean(checked))}
|
||||
/>
|
||||
);
|
||||
}
|
||||
|
||||
function SelectRowCheckbox<TData>({ row, label }: { row: Row<TData>; label: string }) {
|
||||
return (
|
||||
<Checkbox
|
||||
aria-label={label}
|
||||
data-testid={`datatable-select-row-${row.id}`}
|
||||
checked={row.getIsSelected()}
|
||||
disabled={!row.getCanSelect()}
|
||||
onCheckedChange={(checked) => row.toggleSelected(Boolean(checked))}
|
||||
/>
|
||||
);
|
||||
}
|
||||
|
||||
export function createSelectionColumn<TData extends RowData>(
|
||||
options: SelectionColumnOptions<TData> = {},
|
||||
): ColumnDef<TData, unknown> {
|
||||
const { rowAriaLabel } = options;
|
||||
|
||||
return {
|
||||
id: "select",
|
||||
size: 44,
|
||||
enableSorting: false,
|
||||
enableHiding: false,
|
||||
enableResizing: false,
|
||||
meta: { title: "Select", className: "w-11", headerClassName: "w-11" },
|
||||
header: ({ table }) => <SelectAllCheckbox table={table} />,
|
||||
cell: ({ row }) => <SelectRowCheckbox row={row} label={rowAriaLabel?.(row) ?? "Select row"} />,
|
||||
};
|
||||
}
|
||||
|
|
@ -3,6 +3,7 @@ import "./columnMeta";
|
|||
export { DataTable, DataTableConfigError, validateDataTableConfig } from "./DataTable";
|
||||
export { DataTableFilterDrawer, DataTableFilterField, type FilterDraft } from "./DataTableFilterDrawer";
|
||||
export { DataTablePagination, DEFAULT_PAGE_SIZE_OPTIONS } from "./DataTablePagination";
|
||||
export { createSelectionColumn } from "./DataTableSelectionColumn";
|
||||
export { DataTableToolbar } from "./DataTableToolbar";
|
||||
export { DataTableViewOptions } from "./DataTableViewOptions";
|
||||
export {
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import type {
|
|||
PaginationState,
|
||||
Row,
|
||||
RowData,
|
||||
RowSelectionState,
|
||||
SortingState,
|
||||
Table,
|
||||
VisibilityState,
|
||||
|
|
@ -59,6 +60,10 @@ export interface DataTableProps<TData extends RowData, TValue> {
|
|||
expanded?: ExpandedState;
|
||||
onExpandedChange?: OnChangeFn<ExpandedState>;
|
||||
|
||||
enableRowSelection?: boolean | ((row: Row<TData>) => boolean);
|
||||
rowSelection?: RowSelectionState;
|
||||
onRowSelectionChange?: OnChangeFn<RowSelectionState>;
|
||||
|
||||
onRowClick?: (row: TData) => void;
|
||||
|
||||
rowClassName?: (row: Row<TData>) => string;
|
||||
|
|
|
|||
28
ui/litellm-dashboard/src/components/ui/checkbox.tsx
Normal file
28
ui/litellm-dashboard/src/components/ui/checkbox.tsx
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
"use client";
|
||||
|
||||
import { Checkbox as CheckboxPrimitive } from "@base-ui/react/checkbox";
|
||||
|
||||
import { cn } from "@/lib/cva.config";
|
||||
import { CheckIcon } from "lucide-react";
|
||||
|
||||
function Checkbox({ className, ...props }: CheckboxPrimitive.Root.Props) {
|
||||
return (
|
||||
<CheckboxPrimitive.Root
|
||||
data-slot="checkbox"
|
||||
className={cn(
|
||||
"peer relative flex size-4 shrink-0 items-center justify-center rounded-[4px] border border-input shadow-xs transition-shadow outline-none group-has-disabled/field:opacity-50 after:absolute after:-inset-x-3 after:-inset-y-2 focus-visible:border-ring focus-visible:ring-3 focus-visible:ring-ring/50 disabled:cursor-not-allowed disabled:opacity-50 aria-invalid:border-destructive aria-invalid:ring-3 aria-invalid:ring-destructive/20 aria-invalid:aria-checked:border-primary dark:bg-input/30 dark:aria-invalid:border-destructive/50 dark:aria-invalid:ring-destructive/40 data-checked:border-primary data-checked:bg-primary data-checked:text-primary-foreground dark:data-checked:bg-primary",
|
||||
className,
|
||||
)}
|
||||
{...props}
|
||||
>
|
||||
<CheckboxPrimitive.Indicator
|
||||
data-slot="checkbox-indicator"
|
||||
className="grid place-content-center text-current transition-none [&>svg]:size-3.5"
|
||||
>
|
||||
<CheckIcon />
|
||||
</CheckboxPrimitive.Indicator>
|
||||
</CheckboxPrimitive.Root>
|
||||
);
|
||||
}
|
||||
|
||||
export { Checkbox };
|
||||
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -22629,6 +22629,12 @@ export interface components {
|
|||
* @description Allowlist of hosts a request may redirect a provider call's destination URL to.
|
||||
*/
|
||||
provider_url_destination_allowed_hosts?: string[] | null;
|
||||
/**
|
||||
* Proxy Config Reload Interval Seconds
|
||||
* @description how often (in seconds) each pod reloads config-in-DB objects (models, credentials, guardrails, etc.) when store_model_in_db is enabled; lower values speed up multi-pod convergence at the cost of more DB load. Applied on proxy startup
|
||||
* @default 30
|
||||
*/
|
||||
proxy_config_reload_interval_seconds: number;
|
||||
/**
|
||||
* Reject Clientside Metadata Tags
|
||||
* @description When set to True, rejects requests that contain client-side 'metadata.tags' to prevent users from influencing budgets by sending different tags. Tags can only be inherited from the API key metadata.
|
||||
|
|
|
|||
2
uv.lock
generated
2
uv.lock
generated
|
|
@ -10,7 +10,7 @@ resolution-markers = [
|
|||
]
|
||||
|
||||
[options]
|
||||
exclude-newer = "2026-07-18T21:51:03.658458698Z"
|
||||
exclude-newer = "2026-07-18T21:57:43.13625Z"
|
||||
exclude-newer-span = "P3D"
|
||||
|
||||
[manifest]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue