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

1"""v2-native refresher for the ``authorization_code`` mode: the refresh_token grant, then persist. 

2 

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""" 

10 

11from __future__ import annotations 

12 

13import time 

14from collections.abc import Awaitable, Callable 

15from typing import TYPE_CHECKING, Final, Protocol 

16 

17from fastapi import HTTPException 

18 

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) 

32 

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 

35 

36ServerLookup = Callable[[str], "MCPServer | None"] 

37TokenEndpointPost = Callable[[str, dict[str, str], dict[str, str]], Awaitable["dict[str, object] | None"]] 

38 

39 

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: ... 

51 

52 

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 

64 

65 

66def _parse_scopes(raw: object) -> tuple[str, ...] | None: 

67 return tuple(raw.split()) if isinstance(raw, str) and raw else None 

68 

69 

70class AuthorizationCodeRefresher: 

71 """``TokenRefresher`` for authorization_code: refresh_token grant against the server, then persist. 

72 

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 """ 

81 

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 

96 

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 

104 

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 

114 

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 

145 

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 

157 

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 )