feat(mcp): v2-native authorization_code token refresher (step 1b)

The refresh_token grant for the authorization_code mode: POSTs the RFC 6749 refresh_token grant to
the server's token endpoint, persists the rotated triple, and returns the new typed OAuthToken for
RefreshingTokenStore to cache. HTTP post and persist are injected so the grant + response parsing
are testable without a live IdP/DB. Also extends the TokenRefresher seam with (user_id, server_id),
which the foundation's refresh(token) lacked but the grant (server config) and persist (key) need.
This commit is contained in:
Tin Chi Lo 2026-06-25 19:38:37 -07:00
parent ea1c6135c2
commit cec32b4574
2 changed files with 275 additions and 0 deletions

View file

@ -0,0 +1,119 @@
"""v2-native refresher for the ``authorization_code`` mode: the refresh_token grant, then persist.
Mints a fresh access token from a stored refresh_token by POSTing the RFC 6749 refresh_token grant to
the server's token endpoint, persists the rotated triple, and returns the new typed ``OAuthToken`` for
``RefreshingTokenStore`` to cache. The HTTP post and the persist are injected, so the orchestration
and the (untyped) response parsing stay testable without a live IdP or DB. Replaces v1's
``refresh_user_oauth_token`` as part of step 1b; rotation safety - one refresh per (user, server)
across replicas - is the wrapping store's distributed single-flight, not this refresher's concern.
"""
from __future__ import annotations
import time
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Protocol
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
if TYPE_CHECKING:
from litellm.types.mcp_server.mcp_server_manager import MCPServer
ServerLookup = Callable[[str], "MCPServer | None"]
TokenEndpointPost = Callable[
[str, dict[str, str]], Awaitable["dict[str, object] | None"]
]
class CredentialPersist(Protocol):
async def __call__(
self,
user_id: str,
server_id: str,
access_token: str,
refresh_token: str | None,
expires_in: int | None,
scopes: tuple[str, ...] | None,
) -> None: ...
def _parse_expires_in(raw: object) -> int | None:
if isinstance(raw, bool):
return None
if isinstance(raw, int):
return raw
if isinstance(raw, str):
try:
return int(raw)
except ValueError:
return None
return None
def _parse_scopes(raw: object) -> tuple[str, ...] | None:
return tuple(raw.split()) if isinstance(raw, str) and raw else None
class AuthorizationCodeRefresher:
"""``TokenRefresher`` for authorization_code: refresh_token grant against the server, then persist.
``token_endpoint`` POSTs the OAuth form and returns the parsed JSON body (``None`` on any
transport/HTTP failure, mirroring v1: a failed refresh is a miss, not a 500). ``persist`` writes
the rotated triple for ``(user, server)`` - the v1 ``store_user_oauth_credential`` write, which
stays. Returns ``None`` (the arm challenges) when there is no refresh_token, the server lacks a
token endpoint, or the grant fails; never a stale or partial token. A rotated refresh_token from
the response replaces the old one; an omitted one is carried forward.
"""
def __init__(
self,
server_lookup: ServerLookup,
token_endpoint: TokenEndpointPost,
persist: CredentialPersist,
*,
clock: Callable[[], float] = time.time,
) -> None:
self._server_lookup = server_lookup
self._token_endpoint = token_endpoint
self._persist = persist
self._clock = clock
async def refresh(
self, user_id: str, server_id: str, token: OAuthToken
) -> OAuthToken | None:
if token.refresh_token is None:
return None
server = self._server_lookup(server_id)
if server is None or not server.token_url:
return None
form = {
"grant_type": "refresh_token",
"refresh_token": token.refresh_token,
**({"client_id": server.client_id} if server.client_id else {}),
**({"client_secret": server.client_secret} if server.client_secret else {}),
}
body = await self._token_endpoint(server.token_url, form)
if body is None:
return None
access_token = body.get("access_token")
if not isinstance(access_token, str) or not access_token:
return None
rotated = body.get("refresh_token")
new_refresh = (
rotated if isinstance(rotated, str) and rotated else token.refresh_token
)
expires_in = _parse_expires_in(body.get("expires_in"))
scopes = _parse_scopes(body.get("scope"))
await self._persist(
user_id, server_id, access_token, new_refresh, expires_in, scopes
)
return OAuthToken(
access_token=access_token,
expires_at=self._clock() + expires_in if expires_in is not None else None,
refresh_token=new_refresh,
)

View file

@ -0,0 +1,156 @@
"""Tests for the authorization_code refresher: the refresh_token grant, parsing, and persist."""
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.authz_code_refresher import (
AuthorizationCodeRefresher,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
class _Server:
def __init__(
self,
token_url="https://idp.example.com/token",
client_id="cid",
client_secret="sec",
):
self.token_url = token_url
self.client_id = client_id
self.client_secret = client_secret
def _lookup(server):
return lambda server_id: server
def _endpoint(body, sink=None):
async def post(url, form):
if sink is not None:
sink.append((url, form))
return body
return post
def _recording_persist(sink):
async def persist(
user_id, server_id, access_token, refresh_token, expires_in, scopes
):
sink.append(
(user_id, server_id, access_token, refresh_token, expires_in, scopes)
)
return persist
def _refresher(
server=None, body=None, *, post_sink=None, persist_sink=None, clock=lambda: 1000.0
):
return AuthorizationCodeRefresher(
_lookup(server if server is not None else _Server()),
_endpoint(body, post_sink),
_recording_persist(persist_sink if persist_sink is not None else []),
clock=clock,
)
@pytest.mark.asyncio
async def test_refreshes_persists_and_returns_typed_token():
persisted = []
posted = []
refresher = _refresher(
body={
"access_token": "new-at",
"expires_in": 3600,
"refresh_token": "new-rt",
"scope": "a b",
},
post_sink=posted,
persist_sink=persisted,
)
token = await refresher.refresh(
"alice", "srv", OAuthToken(access_token="old", refresh_token="old-rt")
)
assert token is not None
assert token.access_token == "new-at"
assert token.refresh_token == "new-rt"
assert token.expires_at == 1000.0 + 3600 # clock + expires_in -> epoch
# the rotated triple is persisted for (user, server) with parsed scopes
assert persisted == [("alice", "srv", "new-at", "new-rt", 3600, ("a", "b"))]
# the grant carried the refresh_token + client credentials
url, form = posted[0]
assert url == "https://idp.example.com/token"
assert form == {
"grant_type": "refresh_token",
"refresh_token": "old-rt",
"client_id": "cid",
"client_secret": "sec",
}
@pytest.mark.asyncio
async def test_no_refresh_token_is_not_refreshable():
posted = []
refresher = _refresher(body={"access_token": "x"}, post_sink=posted)
assert (
await refresher.refresh("alice", "srv", OAuthToken(access_token="old")) is None
)
assert posted == [] # never hit the IdP
@pytest.mark.asyncio
async def test_unknown_server_or_no_token_url_yields_none():
assert (
await _refresher(server=None).refresh(
"a", "s", OAuthToken("old", refresh_token="rt")
)
is None
)
no_url = _Server(token_url=None)
assert (
await _refresher(server=no_url).refresh(
"a", "s", OAuthToken("old", refresh_token="rt")
)
is None
)
@pytest.mark.asyncio
async def test_grant_failure_does_not_persist():
persisted = []
refresher = _refresher(
body=None, persist_sink=persisted
) # token_endpoint signals failure
assert (
await refresher.refresh("a", "s", OAuthToken("old", refresh_token="rt")) is None
)
assert persisted == []
@pytest.mark.asyncio
async def test_response_without_access_token_does_not_persist():
persisted = []
refresher = _refresher(body={"expires_in": 60}, persist_sink=persisted)
assert (
await refresher.refresh("a", "s", OAuthToken("old", refresh_token="rt")) is None
)
assert persisted == []
@pytest.mark.asyncio
async def test_unrotated_refresh_token_is_carried_forward():
persisted = []
refresher = _refresher(body={"access_token": "new-at"}, persist_sink=persisted)
token = await refresher.refresh(
"a", "s", OAuthToken("old", refresh_token="keep-rt")
)
assert token is not None
assert (
token.refresh_token == "keep-rt"
) # response omitted refresh_token -> reuse the old one
assert token.expires_at is None # no expires_in -> no known expiry
assert persisted[0][3] == "keep-rt"