diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py new file mode 100644 index 00000000000..60b25c53fc9 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py @@ -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, + ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py new file mode 100644 index 00000000000..c42288f2a8e --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py @@ -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"