mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
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:
parent
fccd04a494
commit
a535f8c6c3
3 changed files with 159 additions and 9 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue