mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
ea1c6135c2
commit
cec32b4574
2 changed files with 275 additions and 0 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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"
|
||||
Loading…
Add table
Reference in a new issue