Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_pricing_auto_update_action

# Conflicts:
#	uv.lock
This commit is contained in:
mateo-berri 2026-07-21 15:16:22 -07:00
commit 35e1b36d11
26 changed files with 1824 additions and 9 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View 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"
)

View file

@ -234,6 +234,7 @@ CONTROL_PLANE_PREFIXES: tuple[str, ...] = (
"/tag",
"/budget",
"/model/",
"/access_group",
"/spend",
"/global",
"/config",

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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 } : {}),
};

View file

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

View file

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

View file

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

View file

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

View 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 };

View file

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

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