diff --git a/litellm/proxy/_experimental/mcp_server/v2_port_bodies.py b/litellm/proxy/_experimental/mcp_server/v2_port_bodies.py index 8330a0a4530..c00051e7cc8 100644 --- a/litellm/proxy/_experimental/mcp_server/v2_port_bodies.py +++ b/litellm/proxy/_experimental/mcp_server/v2_port_bodies.py @@ -9,7 +9,7 @@ from __future__ import annotations import asyncio from datetime import timedelta -from typing import Awaitable, Callable, Optional +from typing import Awaitable, Callable, Mapping, Optional import httpx from pydantic import BaseModel, ConfigDict, SecretStr, ValidationError @@ -20,7 +20,10 @@ from litellm.proxy.gateway.mcp.outbound_credentials.clock import Clock, SystemCl from litellm.proxy.gateway.mcp.outbound_credentials.credential_store import ( CredentialKey, ) -from litellm.proxy.gateway.mcp.outbound_credentials.token_store import StoredToken +from litellm.proxy.gateway.mcp.outbound_credentials.token_store import ( + StoredToken, + TokenKey, +) from litellm.proxy.gateway.mcp.outbound_credentials.types import ( AssumeRole, AwsSigV4Config, @@ -256,3 +259,91 @@ def _classify_sigv4_error(error: Exception) -> CredError: return CredError.of_misconfigured( f"aws_sigv4 credentials could not be resolved: {error}" ) + + +# v1 refreshes the per-user token on every read, so the token this store returns is always +# currently valid; expiry is set beyond the resolver's refresh buffer to keep v2's proactive +# refresh inert (the OAuth dance and refresh stay on v1 until it retires). +_BRIDGE_TOKEN_TTL = timedelta(hours=1) + + +class _OAuthPayload(BaseModel): + """The slice of v1's stored OAuth2 payload we consume; extra fields are ignored.""" + + model_config = ConfigDict(extra="ignore") + access_token: str + + +async def _read_valid_user_oauth_token( + subject_id: str, server_id: str +) -> Optional[Mapping[str, object]]: + from litellm.proxy._experimental.mcp_server.db import ( + get_user_oauth_credential, + resolve_valid_user_oauth_token, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise RuntimeError("no DB client available for OAuth token lookup") + server = global_mcp_server_manager.get_mcp_server_by_id(server_id) + if server is None: + return None + cred = await get_user_oauth_credential(prisma_client, subject_id, server_id) + return await resolve_valid_user_oauth_token(subject_id, server, cred, prisma_client) + + +class V1OAuthTokenStore: + """Per-user authorization_code TokenStore backed by v1's stored OAuth tokens. + + Read-path graft: ``get`` returns the user's currently-valid upstream token, reusing v1's + ``resolve_valid_user_oauth_token`` (which reads the stored credential and refreshes it via the + server's token_url when near expiry). Because v1 refreshes on every read, the returned token is + always currently valid, so ``expires_at`` is set beyond the resolver's refresh buffer to keep + v2's proactive-refresh path inert; the dance and refresh stay on v1 until it retires, at which + point a v2-owned store returns real expiry and the resolver owns refresh. A missing/invalid + token is ``Ok(None)`` (the arm turns that into a 401 that starts the OAuth flow); a DB outage is + ``upstream_unavailable``. The reader is injected so the arm stays unit-testable without a DB. + """ + + def __init__( + self, + reader: Callable[ + [str, str], Awaitable[Optional[Mapping[str, object]]] + ] = _read_valid_user_oauth_token, + clock: Optional[Clock] = None, + ) -> None: + self._reader = reader + self._clock: Clock = clock or SystemClock() + + async def get(self, key: TokenKey) -> Result[Optional[StoredToken], CredError]: + if not key.subject_id: + # No authenticated identity -> no per-user token; never share one slot. + return Ok(None) + try: + payload = await self._reader(key.subject_id, key.server_id) + except Exception as e: + return Error( + CredError.of_upstream_unavailable(f"OAuth token lookup failed: {e}") + ) + if payload is None: + return Ok(None) + try: + parsed = _OAuthPayload.model_validate(payload) + except ValidationError: + return Ok(None) # no usable access_token -> drives the OAuth dance + return Ok( + StoredToken( + access_token=SecretStr(parsed.access_token), + expires_at=self._clock.now() + _BRIDGE_TOKEN_TTL, + refresh_token=None, + ) + ) + + async def put(self, key: TokenKey, token: StoredToken) -> Result[None, CredError]: + # Unreached during the bridge: get returns non-near-expiry tokens, so the resolver never + # refreshes/persists through this port; v1 persists inside resolve_valid_user_oauth_token. + _ = (key, token) + return Ok(None) diff --git a/litellm/proxy/_experimental/mcp_server/v2_resolver_bridge.py b/litellm/proxy/_experimental/mcp_server/v2_resolver_bridge.py index c3a14188c3b..c4733e05a84 100644 --- a/litellm/proxy/_experimental/mcp_server/v2_resolver_bridge.py +++ b/litellm/proxy/_experimental/mcp_server/v2_resolver_bridge.py @@ -25,6 +25,7 @@ from litellm.proxy._experimental.mcp_server.v2_port_bodies import ( HttpxSigV4Signer, HttpxTokenExchanger, V1ByokCredentialStore, + V1OAuthTokenStore, ) from litellm.proxy.gateway.mcp.outbound_credentials.clock import SystemClock from litellm.proxy.gateway.mcp.outbound_credentials.resolver import ( @@ -33,10 +34,7 @@ from litellm.proxy.gateway.mcp.outbound_credentials.resolver import ( from litellm.proxy.gateway.mcp.outbound_credentials.service_token_store import ( InMemoryServiceTokenStore, ) -from litellm.proxy.gateway.mcp.outbound_credentials.token_store import ( - InMemoryTokenStore, - StoredToken, -) +from litellm.proxy.gateway.mcp.outbound_credentials.token_store import StoredToken from litellm.proxy.gateway.mcp.outbound_credentials.types import ( Ambient, ApiKeyConfig, @@ -94,12 +92,13 @@ class _Unwired: @functools.lru_cache(maxsize=1) def _provider() -> UpstreamCredentialProvider: - # none + api_key resolve from config alone, so the stores/ports below are inert placeholders; - # real bodies get wired in as their modes are grafted. + # Real bodies are wired as their modes are grafted. token_refresher stays unwired: the bridge's + # V1OAuthTokenStore returns currently-valid tokens (v1 refreshes on read), so the resolver's + # proactive-refresh path is inert until v1 retires. unwired = _Unwired() return UpstreamCredentialProvider( credential_store=V1ByokCredentialStore(), - token_store=InMemoryTokenStore(), + token_store=V1OAuthTokenStore(), token_refresher=unwired, clock=SystemClock(), service_token_store=InMemoryServiceTokenStore(), diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_v2_port_bodies.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_v2_port_bodies.py index b28a1c54d66..2d56cbddc37 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_v2_port_bodies.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_v2_port_bodies.py @@ -13,11 +13,13 @@ from litellm.proxy._experimental.mcp_server.v2_port_bodies import ( HttpxSigV4Signer, HttpxTokenExchanger, V1ByokCredentialStore, + V1OAuthTokenStore, _classify_sigv4_error, ) from litellm.proxy.gateway.mcp.outbound_credentials.credential_store import ( CredentialKey, ) +from litellm.proxy.gateway.mcp.outbound_credentials.token_store import TokenKey from litellm.proxy.gateway.mcp.outbound_credentials.types import ( AwsSigV4Config, ClientCredentialsConfig, @@ -211,6 +213,64 @@ async def test_byok_store_db_error_is_upstream_unavailable(): assert result.error.tag == "upstream_unavailable" +def _token_key(subject_id="u1", server_id="s1"): + return TokenKey( + tenant_id="org1", + subject_id=subject_id, + server_id=server_id, + resource="https://up.example/mcp", + ) + + +async def test_oauth_store_returns_valid_token(): + async def reader(subject_id, server_id): + assert (subject_id, server_id) == ("u1", "s1") + return {"type": "oauth2", "access_token": "valid-access", "refresh_token": "r"} + + result = await V1OAuthTokenStore(reader=reader).get(_token_key()) + assert isinstance(result, Ok) + assert result.ok is not None + assert result.ok.access_token.get_secret_value() == "valid-access" + # v1 owns refresh, so v2's refresh path is kept inert: no refresh_token, expiry beyond the buffer + assert result.ok.refresh_token is None + + +async def test_oauth_store_missing_token_is_ok_none(): + async def reader(subject_id, server_id): + return None + + result = await V1OAuthTokenStore(reader=reader).get(_token_key()) + assert isinstance(result, Ok) + assert result.ok is None + + +async def test_oauth_store_payload_without_access_token_is_ok_none(): + async def reader(subject_id, server_id): + return {"type": "oauth2"} # no access_token -> drives the OAuth dance + + result = await V1OAuthTokenStore(reader=reader).get(_token_key()) + assert isinstance(result, Ok) + assert result.ok is None + + +async def test_oauth_store_empty_subject_skips_the_store(): + async def reader(subject_id, server_id): + raise AssertionError("store must not be queried for an empty subject") + + result = await V1OAuthTokenStore(reader=reader).get(_token_key(subject_id="")) + assert isinstance(result, Ok) + assert result.ok is None + + +async def test_oauth_store_db_error_is_upstream_unavailable(): + async def reader(subject_id, server_id): + raise RuntimeError("db down") + + result = await V1OAuthTokenStore(reader=reader).get(_token_key()) + assert isinstance(result, Error) + assert result.error.tag == "upstream_unavailable" + + def _te_config(endpoint="https://idp/exchange", scopes=(), client_secret="csecret"): return TokenExchangeConfig( token_exchange_endpoint=endpoint,