feat(mcp/v2): wire authorization_code TokenStore bridge (read path)

V1OAuthTokenStore backs the authorization_code arm's TokenStore by reading the user's
currently-valid upstream token via v1's resolve_valid_user_oauth_token (read + refresh-on-read).
Because v1 refreshes on every read, the returned token is always valid, so expires_at is set beyond
the resolver's refresh buffer to keep v2's proactive-refresh path inert; the OAuth dance and refresh
stay on v1 until it retires (token_refresher remains unwired). A missing/invalid token is Ok(None)
(the arm 401s to start the OAuth flow); a DB outage is upstream_unavailable. The reader is injected
for unit tests.

Wired into _provider (replaces the in-memory stub). With the _to_server_spec authorization_code
mapping, an interactive oauth2 (delegate=false) server resolves to a stored per-user token or a 401,
never forwarding the caller JWT (the v2 form of LIT-3795). Fires live at the egress cutover.
This commit is contained in:
Tin Chi Lo 2026-06-19 10:44:54 -07:00
parent fccd04a494
commit a535f8c6c3
3 changed files with 159 additions and 9 deletions

View file

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

View file

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

View file

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