perf(mcp): cache SSO identity assertion reads on the ID-JAG path (#39348)

* perf(mcp): cache SSO identity assertion reads on the ID-JAG path

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): guard sso assertion cache against stale relogin reads

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): keep sso assertion cache entries and generation markers in separate namespaces

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(router): rename duplicate get_configured_mode test so ruff F811 passes

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): use a process-wide epoch for sso assertion cache invalidation

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-03 17:25:37 -07:00 • committed by GitHub
parent e26d607f5d
commit 1add1b4655
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 207 additions and 12 deletions

View file

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

View file

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

View file

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