Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py: 27%
76 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""v2-native refresher for the ``authorization_code`` mode: the refresh_token grant, then persist.
3Mints a fresh access token from a stored refresh_token by POSTing the RFC 6749 refresh_token grant to
4the server's token endpoint, persists the rotated triple, and returns the new typed ``OAuthToken`` for
5``RefreshingTokenStore`` to cache. The HTTP post and the persist are injected, so the orchestration
6and the (untyped) response parsing stay testable without a live IdP or DB. Replaces v1's
7``refresh_user_oauth_token`` as part of step 1b; rotation safety - one refresh per (user, server)
8across replicas - is the wrapping store's distributed single-flight, not this refresher's concern.
9"""
11from __future__ import annotations
13import time
14from collections.abc import Awaitable, Callable
15from typing import TYPE_CHECKING, Final, Protocol
17from fastapi import HTTPException
19from litellm._logging import verbose_logger
20from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
21 TokenEndpointAuthConfigError,
22)
23from litellm.proxy._experimental.mcp_server.oauth_identity_binding import (
24 BindingValidator,
25 RefreshTokenPresented,
26 enforce_oauth_identity_binding,
27)
28from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request
29from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
30 OAuthToken,
31)
33if TYPE_CHECKING: 33 ↛ 34line 33 didn't jump to line 34 because the condition on line 33 was never true
34 from litellm.types.mcp_server.mcp_server_manager import MCPServer
36ServerLookup = Callable[[str], "MCPServer | None"]
37TokenEndpointPost = Callable[[str, dict[str, str], dict[str, str]], Awaitable["dict[str, object] | None"]]
40class CredentialPersist(Protocol):
41 async def __call__( 41 ↛ exitline 41 didn't return from function '__call__' because
42 self,
43 user_id: str,
44 server_id: str,
45 access_token: str,
46 refresh_token: str | None,
47 expires_in: int | None,
48 scopes: tuple[str, ...] | None,
49 identity_binding_proof: str | None = None,
50 ) -> None: ...
53def _parse_expires_in(raw: object) -> int | None:
54 if isinstance(raw, bool):
55 return None
56 if isinstance(raw, int):
57 return raw
58 if isinstance(raw, str):
59 try:
60 return int(raw)
61 except ValueError:
62 return None
63 return None
66def _parse_scopes(raw: object) -> tuple[str, ...] | None:
67 return tuple(raw.split()) if isinstance(raw, str) and raw else None
70class AuthorizationCodeRefresher:
71 """``TokenRefresher`` for authorization_code: refresh_token grant against the server, then persist.
73 ``token_endpoint`` POSTs the OAuth form and returns the parsed JSON body (``None`` on any
74 transport/HTTP failure, mirroring v1: a failed refresh is a miss, not a 500). ``persist`` writes
75 the rotated triple for ``(user, server)`` - the v1 ``store_user_oauth_credential`` write, which
76 stays. Returns ``None`` (the arm challenges) when there is no refresh_token, the server lacks a
77 token endpoint, or the grant fails; never a stale or partial token. A rotated refresh_token from
78 the response replaces the old one; an omitted one is carried forward, as are the recorded scopes
79 when the response omits ``scope``.
80 """
82 def __init__(
83 self,
84 server_lookup: ServerLookup,
85 token_endpoint: TokenEndpointPost,
86 persist: CredentialPersist,
87 *,
88 clock: Callable[[], float] = time.time,
89 identity_validator: BindingValidator = enforce_oauth_identity_binding,
90 ) -> None:
91 self._server_lookup = server_lookup
92 self._token_endpoint = token_endpoint
93 self._persist = persist
94 self._clock = clock
95 self._identity_validator = identity_validator
97 async def refresh(self, user_id: str, server_id: str, token: OAuthToken) -> OAuthToken | None:
98 try:
99 return await self._refresh(user_id, server_id, token)
100 except HTTPException as exc:
101 if exc.status_code != 403:
102 raise
103 return None
105 async def _refresh(self, user_id: str, server_id: str, token: OAuthToken) -> OAuthToken | None:
106 if token.refresh_token is None:
107 return None
108 server: Final = self._server_lookup(server_id)
109 if server is None:
110 return None
111 token_url: Final = server.effective_token_url
112 if not token_url:
113 return None
115 try:
116 token_request: Final = build_upstream_oauth2_token_request(
117 server,
118 auth_method=server.token_endpoint_auth_method,
119 client_id=server.client_id,
120 client_secret=server.client_secret,
121 )
122 except TokenEndpointAuthConfigError as exc:
123 verbose_logger.warning("MCP OAuth refresh misconfigured for server %s: %s", server_id, exc)
124 return None
125 binding: Final = server.oauth_identity_binding
126 if binding is not None and binding.mode == "enforce":
127 await self._identity_validator(
128 server=server,
129 token_response={},
130 litellm_user_id=user_id,
131 grant_type="refresh_token",
132 refresh_ownership=RefreshTokenPresented(token.refresh_token),
133 )
134 form: Final = {
135 "grant_type": "refresh_token",
136 "refresh_token": token.refresh_token,
137 **token_request.body,
138 }
139 body: Final = await self._token_endpoint(token_url, form, token_request.headers)
140 if body is None:
141 return None
142 access_token: Final = body.get("access_token")
143 if not isinstance(access_token, str) or not access_token:
144 return None
146 binding_proof: Final = await self._identity_validator(
147 server=server,
148 token_response=body,
149 litellm_user_id=user_id,
150 grant_type="refresh_token",
151 refresh_ownership=RefreshTokenPresented(token.refresh_token),
152 )
153 rotated: Final = body.get("refresh_token")
154 new_refresh: Final = rotated if isinstance(rotated, str) and rotated else token.refresh_token
155 expires_in: Final = _parse_expires_in(body.get("expires_in"))
156 scopes: Final = _parse_scopes(body.get("scope")) or token.scopes
158 if binding_proof is not None:
159 await self._persist(
160 user_id,
161 server_id,
162 access_token,
163 new_refresh,
164 expires_in,
165 scopes or None,
166 identity_binding_proof=binding_proof,
167 )
168 else:
169 await self._persist(user_id, server_id, access_token, new_refresh, expires_in, scopes or None)
170 return OAuthToken(
171 access_token=access_token,
172 expires_at=self._clock() + expires_in if expires_in is not None else None,
173 refresh_token=new_refresh,
174 scopes=scopes,
175 identity_binding_proof=binding_proof,
176 )