Merge pull request #34072 from BerriAI/litellm_mcp_ema_assertion_store

feat(mcp): store the enterprise IdP identity assertion at SSO login for EMA egress
This commit is contained in:
tin-berri 2026-07-21 15:10:35 -07:00 • committed by GitHub
commit 869ef0cbfb
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 975 additions and 6 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

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

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

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