mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
e26d607f5d
commit
1add1b4655
3 changed files with 207 additions and 12 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue