Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -448,6 +448,12 @@ async def _store_per_user_token_server_side(
)
return # Don't warm Redis if DB write failed

from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)

await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server.server_id)

# Warm the Redis cache so the first subsequent MCP call is a cache hit
ttl = _compute_per_user_token_ttl(server, expires_in)
await mcp_per_user_token_cache.set(
Expand Down
27 changes: 25 additions & 2 deletions litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,9 @@
to_server_spec,
to_subject,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
InvalidatableOAuthTokenStore,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
LazyPerUserOAuthTokenStore,
)
Expand Down Expand Up @@ -552,9 +555,16 @@ def _obo_needs_endpoint_discovery(
"""
return auth_type == MCPAuth.oauth2_token_exchange and not (token_exchange_endpoint or token_url)

def __init__(self, cred_provider: Optional[UpstreamCredentialProvider] = None):
def __init__(
self,
cred_provider: Optional[UpstreamCredentialProvider] = None,
per_user_oauth_token_store: Optional[InvalidatableOAuthTokenStore] = None,
):
self._per_user_oauth_token_store = per_user_oauth_token_store or LazyPerUserOAuthTokenStore(
self.get_mcp_server_by_id
)
self._cred_provider = cred_provider or UpstreamCredentialProvider(
oauth_token_store=LazyPerUserOAuthTokenStore(self.get_mcp_server_by_id),
oauth_token_store=self._per_user_oauth_token_store,
token_exchanger=build_token_exchanger(),
)
self.registry: dict[str, MCPServer] = {}
Expand Down Expand Up @@ -3771,6 +3781,19 @@ async def has_user_oauth_token(self, server: MCPServer, user_api_key_auth: Optio
return False
return await self._cred_provider.has_user_token(to_subject(user_api_key_auth, None), spec)

async def invalidate_user_oauth_token_cache(self, user_id: str, server_id: str) -> None:
"""Drop the v2 chain's cached token for ``(user_id, server_id)`` after the credential row
changes (re-auth, revoke), so the next resolve reads the new row instead of serving the
replaced token until its cache TTL. Best-effort: a cache-drop failure is logged, never
raised, because the DB write already succeeded and the TTL remains the backstop.
"""
try:
await self._per_user_oauth_token_store.invalidate(user_id, server_id)
except Exception as exc: # noqa: BLE001 - cache drop is best-effort; TTL is the backstop
verbose_logger.warning(
"Failed to invalidate cached MCP OAuth token for user=%s server=%s: %s", user_id, server_id, exc
)

async def _resolve_oauth2_headers_for_tool_call(
self,
mcp_server: MCPServer,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,17 @@ class OAuthTokenStore(Protocol):
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: ...


class InvalidatableOAuthTokenStore(OAuthTokenStore, Protocol):
"""An ``OAuthTokenStore`` whose cached entry for a ``(user, server)`` pair can be dropped.

The write side calls ``invalidate`` after a (re)authorization or revocation changes the
credential row, so reads stop serving the replaced token immediately instead of until its
cache TTL. ``CachedOAuthTokenStore`` (the top of the per-user chain) satisfies this.
"""

async def invalidate(self, user_id: str, server_id: str) -> None: ...


class TokenRefresher(Protocol):
"""Mints a fresh token from an expired one and persists it, returning the new token.

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -24,8 +24,8 @@
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
CachedOAuthTokenStore,
InvalidatableOAuthTokenStore,
OAuthToken,
OAuthTokenStore,
RefreshCoordinator,
RefreshingTokenStore,
TokenCacheBackend,
Expand All @@ -51,7 +51,7 @@
_DEFAULT_TTL_SECONDS = 300.0

ServerLookup = Callable[[str], "MCPServer | None"]
StoreBuilder = Callable[[ServerLookup], tuple[OAuthTokenStore, bool]]
StoreBuilder = Callable[[ServerLookup], tuple[InvalidatableOAuthTokenStore, bool]]


async def _read_credential(user_id: str, server_id: str) -> dict[str, object] | None:
Expand Down Expand Up @@ -185,7 +185,7 @@ def __init__(
self._server_lookup = server_lookup
self._store_builder = store_builder
self._redis_available = redis_available
self._store: OAuthTokenStore | None = None
self._store: InvalidatableOAuthTokenStore | None = None
self._uses_redis = False
self._fetch_lock = asyncio.Condition()
self._local_fetches = 0
Expand All @@ -203,7 +203,26 @@ async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
if not uses_redis:
await self._finish_local_fetch()

async def _store_for_fetch(self) -> tuple[OAuthTokenStore, bool]:
async def invalidate(self, user_id: str, server_id: str) -> None:
"""Drop the chain's cached entry for ``(user_id, server_id)`` after the credential row
changes (re-auth, revoke). Builds the chain if no fetch has run yet, so a shared (Redis)
cache entry written by another worker is dropped too; the in-process case is then a no-op
on an empty cache.
"""
if self._uses_redis:
store = self._store
if store is not None:
await store.invalidate(user_id, server_id)
return

store, uses_redis = await self._store_for_fetch()
try:
await store.invalidate(user_id, server_id)
finally:
if not uses_redis:
await self._finish_local_fetch()

async def _store_for_fetch(self) -> tuple[InvalidatableOAuthTokenStore, bool]:
async with self._fetch_lock:
while (
self._store is not None and not self._uses_redis and self._redis_available() and self._local_fetches > 0
Expand Down
10 changes: 10 additions & 0 deletions litellm/proxy/management_endpoints/mcp_management_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -1874,6 +1874,11 @@ async def store_mcp_oauth_user_credential(
expires_in=payload.expires_in,
scopes=payload.scopes,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)

await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server_id)
# Read back the persisted record so the response reflects the stored
# expires_at rather than recomputing it here (which could diverge by
# milliseconds or if the storage logic ever adds a grace period).
Expand Down Expand Up @@ -1914,6 +1919,11 @@ async def delete_mcp_oauth_user_credential(
await delete_user_credential(prisma_client, user_id, server_id)
except RecordNotFoundError:
pass # Already gone — treat as a successful delete
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)

await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server_id)
return MCPOAuthUserCredentialStatus(
server_id=server_id,
has_credential=False,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,8 @@
import pytest

from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
InvalidatableOAuthTokenStore,
OAuthToken,
OAuthTokenStore,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
LazyPerUserOAuthTokenStore,
Expand All @@ -16,25 +16,33 @@ class _RecordingStore:
def __init__(self, access_token: str) -> None:
self._access_token = access_token
self.calls: list[tuple[str, str]] = []
self.invalidations: list[tuple[str, str]] = []

async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
self.calls.append((user_id, server_id))
return OAuthToken(access_token=self._access_token)

async def invalidate(self, user_id: str, server_id: str) -> None:
self.invalidations.append((user_id, server_id))


class _BlockingStore:
def __init__(self, access_token: str) -> None:
self._access_token = access_token
self.started = asyncio.Event()
self.release = asyncio.Event()
self.calls: list[tuple[str, str]] = []
self.invalidations: list[tuple[str, str]] = []

async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
self.calls.append((user_id, server_id))
self.started.set()
await self.release.wait()
return OAuthToken(access_token=self._access_token)

async def invalidate(self, user_id: str, server_id: str) -> None:
self.invalidations.append((user_id, server_id))


class _RedisAvailability:
def __init__(self) -> None:
Expand All @@ -59,7 +67,7 @@ async def test_lazy_store_rebuilds_when_redis_becomes_available() -> None:
redis_available = _RedisAvailability()
build_calls = 0

def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]:
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
if redis_available.available:
Expand Down Expand Up @@ -94,7 +102,7 @@ async def test_lazy_store_allows_concurrent_local_fetches_without_redis() -> Non
redis_available = _RedisAvailability()
build_calls = 0

def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]:
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
return local_store, False
Expand Down Expand Up @@ -127,7 +135,7 @@ async def test_lazy_store_waits_for_in_flight_local_fetch_before_redis_rebuild()
redis_store = _RecordingStore("redis")
redis_available = _RedisAvailability()

def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]:
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
if redis_available.available:
return redis_store, True
return local_store, False
Expand Down Expand Up @@ -158,3 +166,83 @@ def server_lookup(_server_id: str) -> None:
assert second is not None and second.access_token == "redis"
assert local_store.calls == [("u", "s")]
assert redis_store.calls == [("u", "s")]


@pytest.mark.asyncio
async def test_lazy_store_invalidate_builds_chain_and_delegates() -> None:
local_store = _RecordingStore("local")
build_calls = 0

def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
return local_store, False

def server_lookup(_server_id: str) -> None:
return None

store = LazyPerUserOAuthTokenStore(
server_lookup,
store_builder=build_store,
redis_available=_RedisAvailability(),
)

await store.invalidate("u", "s")

assert build_calls == 1
assert local_store.invalidations == [("u", "s")]


@pytest.mark.asyncio
async def test_lazy_store_invalidate_reaches_the_store_fetch_reads() -> None:
local_store = _RecordingStore("local")
build_calls = 0

def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
return local_store, False

def server_lookup(_server_id: str) -> None:
return None

store = LazyPerUserOAuthTokenStore(
server_lookup,
store_builder=build_store,
redis_available=_RedisAvailability(),
)

await store.fetch("u", "s")
await store.invalidate("u", "s")

assert build_calls == 1
assert local_store.calls == [("u", "s")]
assert local_store.invalidations == [("u", "s")]


@pytest.mark.asyncio
async def test_lazy_store_invalidate_works_after_redis_chain_is_built() -> None:
redis_store = _RecordingStore("redis")
redis_available = _RedisAvailability()
redis_available.available = True
build_calls = 0

def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
return redis_store, True

def server_lookup(_server_id: str) -> None:
return None

store = LazyPerUserOAuthTokenStore(
server_lookup,
store_builder=build_store,
redis_available=redis_available,
)

await store.fetch("u", "s")
await store.invalidate("u", "s")

assert build_calls == 1
assert redis_store.invalidations == [("u", "s")]
Loading
Loading