diff --git a/litellm/constants.py b/litellm/constants.py index cd72adc3db5..da731cb5eb2 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -151,6 +151,7 @@ DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD = float(os.getenv("DEFAULT_SEMANTIC_ MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS: Final = int(os.getenv("MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS", "60")) MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE", "200")) MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL", "3600")) +MCP_SSO_ASSERTION_CACHE_TTL_SECONDS: Final = int(os.getenv("MCP_SSO_ASSERTION_CACHE_TTL_SECONDS", "60")) # Default npm cache directory for STDIO MCP servers. # npm/npx needs a writable cache dir; in containers the default (~/.npm) 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 8a0d41584f9..6552008ca54 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 @@ -11,7 +11,8 @@ 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. +truth, the same contract as the per-user OAuth credential store. Reads use a per-process cache with +TTL ``MCP_SSO_ASSERTION_CACHE_TTL_SECONDS``; invalidation also guards against stale in-flight reads. """ from __future__ import annotations @@ -24,6 +25,8 @@ import jwt from pydantic import BaseModel, ConfigDict, SecretStr, TypeAdapter, ValidationError from litellm._logging import verbose_proxy_logger +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.constants import MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE, MCP_SSO_ASSERTION_CACHE_TTL_SECONDS if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient @@ -45,6 +48,46 @@ class SSOIdentityAssertion(BaseModel): expires_at: datetime | None = None +class SSOAssertionCache: + """Process-local read cache. ``invalidate`` bumps a process-wide epoch so a fetch that started + before a login cannot repopulate the old assertion after it.""" + + def __init__(self, ttl_seconds: int = MCP_SSO_ASSERTION_CACHE_TTL_SECONDS) -> None: + self._entries = InMemoryCache( + max_size_in_memory=MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE, + default_ttl=ttl_seconds, + ) + self._epoch: int = 0 + + def epoch(self) -> int: + return self._epoch + + def get(self, user_id: str) -> SSOIdentityAssertion | None: + cached: Final = self._entries.get_cache( # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped + user_id + ) + return cached if isinstance(cached, SSOIdentityAssertion) else None + + def set_if_unchanged(self, user_id: str, assertion: SSOIdentityAssertion, seen_epoch: int) -> None: + if self._epoch != seen_epoch: + return + self._entries.set_cache( # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped + user_id, assertion + ) + + def invalidate(self, user_id: str) -> None: + self._epoch += 1 + self._entries.delete_cache( # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped + user_id + ) + + def flush(self) -> None: + self._entries.flush_cache() # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped + + +_ASSERTION_CACHE: Final = SSOAssertionCache() + + class _IdTokenClaims(BaseModel): exp: float | None = None iss: str | None = None @@ -107,7 +150,9 @@ async def ema_assertion_retention_enabled() -> bool: return row is not None -async def persist_sso_identity_assertion(user_id: str, assertion: SSOIdentityAssertion) -> None: +async def persist_sso_identity_assertion( + user_id: str, assertion: SSOIdentityAssertion, cache: SSOAssertionCache = _ASSERTION_CACHE +) -> 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 @@ -127,11 +172,10 @@ async def persist_sso_identity_assertion(user_id: str, assertion: SSOIdentityAss "update": {"assertion_b64": encoded}, }, ) + cache.invalidate(user_id) -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.""" +async def _read_assertion_from_db(user_id: str) -> SSOIdentityAssertion | None: 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 @@ -160,6 +204,21 @@ async def fetch_sso_identity_assertion(user_id: str) -> SSOIdentityAssertion | N ) +async def fetch_sso_identity_assertion( + user_id: str, cache: SSOAssertionCache = _ASSERTION_CACHE +) -> 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.""" + cached: Final = cache.get(user_id) + if cached is not None: + return cached + seen_epoch: Final = cache.epoch() + assertion: Final = await _read_assertion_from_db(user_id) + if assertion is not None: + cache.set_if_unchanged(user_id, assertion, seen_epoch) + return assertion + + class AssertionStoreUnavailable(Exception): """Raised by ``fetch`` when the backing store is unreachable (e.g. the DB is down). @@ -189,9 +248,12 @@ class DbSSOAssertionStore: from credential resolution and from the upstream-401 retry. """ + def __init__(self, cache: SSOAssertionCache = _ASSERTION_CACHE) -> None: + self._cache = cache + async def fetch(self, user_id: str) -> SSOIdentityAssertion | None: try: - return await fetch_sso_identity_assertion(user_id) + return await fetch_sso_identity_assertion(user_id, cache=self._cache) except Exception as exc: # noqa: BLE001 # any driver/storage failure is an outage, not an absence raise AssertionStoreUnavailable(str(exc)) from exc 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 7b82e004f37..5d6d47b8c38 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 @@ -7,6 +7,7 @@ round-trips exactly, a store failure never escapes into the login path, and a sa rotation re-encrypts stored rows like the sibling per-user credential tables. """ +import asyncio import json import os import time @@ -16,8 +17,10 @@ import jwt as pyjwt import pytest from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import ( + _ASSERTION_CACHE, AssertionStoreUnavailable, DbSSOAssertionStore, + SSOAssertionCache, assertion_from_sso_login, ema_assertion_retention_enabled, fetch_sso_identity_assertion, @@ -25,7 +28,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_s 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.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper from litellm.types.mcp import MCPAuth SALT_KEY = "test-salt-key-for-sso-assertion-tests-1234" @@ -38,6 +41,11 @@ def _set_salt_key(monkeypatch): monkeypatch.setenv("LITELLM_SALT_KEY", SALT_KEY) +@pytest.fixture(autouse=True) +def _flush_assertion_cache(): + _ASSERTION_CACHE.flush() + + 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}, @@ -52,9 +60,7 @@ def _make_prisma(stored: dict, db_has_id_jag_server: bool = False): ``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 - ) + 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"] @@ -235,6 +241,78 @@ async def test_persist_overwrites_previous_login(): assert fetched.refresh_token is not None +@pytest.mark.asyncio +async def test_fetch_serves_second_read_from_cache_without_db_read(): + stored = {} + prisma = _make_prisma(stored) + cache = SSOAssertionCache() + token = _make_id_token() + assertion = assertion_from_sso_login(token, "rt_1") + with patch("litellm.proxy.proxy_server.prisma_client", prisma): # test-quality-ok: fake DB seam has no injection boundary + await persist_sso_identity_assertion("user-a", assertion, cache=cache) + first = await fetch_sso_identity_assertion("user-a", cache=cache) + second = await fetch_sso_identity_assertion("user-a", cache=cache) + assert first is not None + assert second is not None + assert first.id_token.get_secret_value() == token + assert second.id_token.get_secret_value() == token + prisma.db.litellm_ssoidentityassertion.find_unique.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_persist_busts_cache_so_relogin_is_visible_immediately(): + stored = {} + prisma = _make_prisma(stored) + cache = SSOAssertionCache() + first_token = _make_id_token(exp_offset=100) + second_token = _make_id_token(exp_offset=7200) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): # test-quality-ok: fake DB seam has no injection boundary + await persist_sso_identity_assertion("user-a", assertion_from_sso_login(first_token, None), cache=cache) + first = await fetch_sso_identity_assertion("user-a", cache=cache) + await persist_sso_identity_assertion("user-a", assertion_from_sso_login(second_token, "rt_new"), cache=cache) + second = await fetch_sso_identity_assertion("user-a", cache=cache) + assert first is not None + assert second is not None + assert first.id_token.get_secret_value() == first_token + assert second.id_token.get_secret_value() == second_token + + +@pytest.mark.asyncio +async def test_fetch_racing_a_relogin_does_not_cache_the_previous_assertion(): + stored = {} + prisma = _make_prisma(stored) + cache = SSOAssertionCache() + first_token = _make_id_token(exp_offset=100) + second_token = _make_id_token(exp_offset=7200) + db_read_started = asyncio.Event() + relogin_done = asyncio.Event() + unpaused_find_unique = prisma.db.litellm_ssoidentityassertion.find_unique + + async def _paused_find_unique(where): + row = await unpaused_find_unique(where=where) + db_read_started.set() + await relogin_done.wait() + return row + + prisma.db.litellm_ssoidentityassertion.find_unique = _paused_find_unique + with patch("litellm.proxy.proxy_server.prisma_client", prisma): # test-quality-ok: fake DB seam has no injection boundary + await persist_sso_identity_assertion("user-a", assertion_from_sso_login(first_token, None), cache=cache) + racing_fetch = asyncio.create_task(fetch_sso_identity_assertion("user-a", cache=cache)) + await db_read_started.wait() + await persist_sso_identity_assertion( + "user-a", + assertion_from_sso_login(second_token, "rt_new"), + cache=cache, + ) + relogin_done.set() + raced = await racing_fetch + after = await fetch_sso_identity_assertion("user-a", cache=cache) + assert raced is not None + assert after is not None + assert raced.id_token.get_secret_value() == first_token + assert after.id_token.get_secret_value() == second_token + + @pytest.mark.asyncio async def test_fetch_missing_row_returns_none(): prisma = _make_prisma({}) @@ -242,6 +320,18 @@ async def test_fetch_missing_row_returns_none(): assert await fetch_sso_identity_assertion("nobody") is None +@pytest.mark.asyncio +async def test_fetch_does_not_cache_a_missing_row(): + prisma = _make_prisma({}) + cache = SSOAssertionCache() + with patch("litellm.proxy.proxy_server.prisma_client", prisma): # test-quality-ok: fake DB seam has no injection boundary + first = await fetch_sso_identity_assertion("nobody", cache=cache) + second = await fetch_sso_identity_assertion("nobody", cache=cache) + assert first is None + assert second is None + assert prisma.db.litellm_ssoidentityassertion.find_unique.await_count == 2 + + @pytest.mark.asyncio async def test_fetch_undecryptable_row_returns_none(): prisma = _make_prisma({"user-a": "not-an-encrypted-blob"}) @@ -251,13 +341,37 @@ async def test_fetch_undecryptable_row_returns_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_cached_assertion_expires_after_ttl(): + stored = {} + prisma = _make_prisma(stored) + cache = SSOAssertionCache(ttl_seconds=1) + first_token = _make_id_token(exp_offset=100) + second_token = _make_id_token(exp_offset=7200) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): # test-quality-ok: fake DB seam has no injection boundary + await persist_sso_identity_assertion("user-a", assertion_from_sso_login(first_token, None), cache=cache) + first = await fetch_sso_identity_assertion("user-a", cache=cache) + await persist_sso_identity_assertion( + "user-a", + assertion_from_sso_login(second_token, "rt_new"), + cache=SSOAssertionCache(), + ) + cached = await fetch_sso_identity_assertion("user-a", cache=cache) + time.sleep(1.1) + expired = await fetch_sso_identity_assertion("user-a", cache=cache) + assert first is not None + assert cached is not None + assert expired is not None + assert first.id_token.get_secret_value() == first_token + assert cached.id_token.get_secret_value() == first_token + assert expired.id_token.get_secret_value() == second_token + + @pytest.mark.asyncio async def test_retain_noop_when_no_id_jag_server(): stored = {} @@ -357,6 +471,24 @@ async def test_db_store_converts_a_driver_failure_into_assertion_store_unavailab await DbSSOAssertionStore().fetch("alice") +@pytest.mark.asyncio +async def test_db_store_uses_injected_cache(): + stored = {} + prisma = _make_prisma(stored) + cache = SSOAssertionCache() + token = _make_id_token() + with patch("litellm.proxy.proxy_server.prisma_client", prisma): # test-quality-ok: fake DB seam has no injection boundary + await persist_sso_identity_assertion("alice", assertion_from_sso_login(token, None), cache=cache) + store = DbSSOAssertionStore(cache=cache) + first = await store.fetch("alice") + second = await store.fetch("alice") + assert first is not None + assert second is not None + assert first.id_token.get_secret_value() == token + assert second.id_token.get_secret_value() == token + prisma.db.litellm_ssoidentityassertion.find_unique.assert_awaited_once() + + @pytest.mark.asyncio async def test_db_store_returns_none_for_a_user_with_no_stored_assertion(): """An absent row stays an absence, not an outage, so a user who never signed in still gets the