diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 55762813e2d..23a9c086c73 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -404,7 +404,7 @@ model LiteLLM_MCPServerOAuthClient { } // 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?, connected_at}. +// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}. model LiteLLM_SSOIdentityAssertion { user_id String @id assertion_b64 String 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 index 4b5183327a1..73bb87dcc33 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py @@ -71,8 +71,10 @@ def assertion_from_sso_login(id_token: object, refresh_token: object) -> SSOIden 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 # any decode failure means the token is not EMA-exchangeable; never raise into login - verbose_proxy_logger.warning("SSO id_token could not be decoded as a JWT; not retaining it for EMA egress.") + 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), @@ -82,21 +84,31 @@ def assertion_from_sso_login(id_token: object, refresh_token: object) -> SSOIden ) -def ema_assertion_retention_enabled() -> bool: - """Whether any registered 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.""" - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids a module-load circular import +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. The local registry is + the fast path; when it has none, the DB is consulted as the authoritative source, because + the registry is a per-process snapshot while the assertion write targets the shared DB. A + server added on another pod (or not yet loaded during startup) must still enable retention + here, or the drop only surfaces later as an unexplained challenge at the EMA upstream.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids import cycle global_mcp_server_manager, ) - from litellm.types.mcp import MCPAuth # noqa: PLC0415 # runtime global, not wired at import time + from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global + from litellm.types.mcp import MCPAuth # noqa: PLC0415 # runtime global servers = global_mcp_server_manager.get_registry().values() - return any(server.auth_type == MCPAuth.oauth2_id_jag for server in servers) + if any(server.auth_type == MCPAuth.oauth2_id_jag for server in 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, not wired at import time - from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global, not wired at import time + 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 @@ -119,8 +131,8 @@ async def persist_sso_identity_assertion(user_id: str, assertion: SSOIdentityAss 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, not wired at import time - from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global, not wired at import time + 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 @@ -150,39 +162,38 @@ async def fetch_sso_identity_assertion(user_id: str) -> SSOIdentityAssertion | N 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.""" - from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: PLC0415 # runtime global, not wired at import time + 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, ) - rows = await prisma_client.db.litellm_ssoidentityassertion.find_many() - decoded = [ - ( - row, - _MAYBE_STR_ADAPTER.validate_python( - decrypt_value_helper(row.assertion_b64, _ASSERTION_DECRYPT_LOG_KEY, exception_type="debug") - ), + 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") ) - for row in rows - ] - for row, plaintext in decoded: 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, ) - continue + 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}, ) - rotated = sum(1 for _, plaintext in decoded if plaintext is not None) + 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", - rotated, - len(decoded) - rotated, + sum(outcomes), + len(outcomes) - sum(outcomes), ) @@ -193,7 +204,7 @@ async def retain_sso_identity_assertion_for_ema(user_id: str, assertion: SSOIden if assertion is None: return try: - if not ema_assertion_retention_enabled(): + 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 diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 55762813e2d..23a9c086c73 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -404,7 +404,7 @@ model LiteLLM_MCPServerOAuthClient { } // 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?, connected_at}. +// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}. model LiteLLM_SSOIdentityAssertion { user_id String @id assertion_b64 String diff --git a/schema.prisma b/schema.prisma index 55762813e2d..23a9c086c73 100644 --- a/schema.prisma +++ b/schema.prisma @@ -404,7 +404,7 @@ model LiteLLM_MCPServerOAuthClient { } // 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?, connected_at}. +// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}. model LiteLLM_SSOIdentityAssertion { user_id String @id assertion_b64 String 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 index df4cf50bcec..040bcbc1594 100644 --- 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 @@ -43,10 +43,15 @@ def _make_id_token(exp_offset: int = 3600, iss: str = ISSUER) -> str: ) -def _make_prisma(stored: dict): +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.""" + (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"] @@ -125,18 +130,54 @@ def test_assertion_without_exp_or_iss_still_retained(): assert assertion.issuer is None -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: +@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.get_registry.return_value = { "s1": _server_with_auth(MCPAuth.oauth2), "s2": _server_with_auth(None), } - assert ema_assertion_retention_enabled() is False + assert await ema_assertion_retention_enabled() is False manager.get_registry.return_value = { "s1": _server_with_auth(MCPAuth.oauth2), "s2": _server_with_auth(MCPAuth.oauth2_id_jag), } - assert ema_assertion_retention_enabled() is True + assert await ema_assertion_retention_enabled() is True + + +@pytest.mark.asyncio +async def test_retention_gate_falls_back_to_db_when_registry_is_cold(): + """A server added on another pod (or before this pod's DB load) is invisible to the local + registry; the gate must still enable retention off the authoritative DB row, and read + False only when neither the registry nor the DB knows an id_jag server.""" + with patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager: + manager.get_registry.return_value = {"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_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.get_registry.return_value = {} + 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 @@ -274,11 +315,13 @@ async def test_rotation_reencrypts_under_new_key(monkeypatch): @pytest.mark.asyncio -async def test_rotation_skips_unreadable_rows(): +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