diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260721000000_add_sso_identity_assertion/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260721000000_add_sso_identity_assertion/migration.sql new file mode 100644 index 00000000000..95412df0a96 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260721000000_add_sso_identity_assertion/migration.sql @@ -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") +); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index b27ddea010b..23a9c086c73 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py new file mode 100644 index 00000000000..e0927cc4f64 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py @@ -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 + ) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 01f4e040e58..ac6a2a4a7db 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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() diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 6c2e06a418c..3c8444ecf26 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -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, diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index b27ddea010b..23a9c086c73 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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 diff --git a/schema.prisma b/schema.prisma index b27ddea010b..23a9c086c73 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py new file mode 100644 index 00000000000..a3f46a49ba9 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index dffca3093fa..51f72f91dc3 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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", + ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 5631aa69102..e1856860c8a 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -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