mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(mcp): preserve recorded OAuth scopes across authorization_code refresh
When a refresh response omits `scope` (RFC 6749 §5.1, where omission means unchanged), the v2 refresher persisted scopes=None and overwrote the user's recorded grant. v1 carried the prior scopes forward via `or cred.get("scopes")`; the v2 path lost that because OAuthToken did not model scopes
OAuthToken now carries scopes, V2PerUserTokenStore populates them on read, and AuthorizationCodeRefresher carries them forward for both the persisted write and the returned/cached token, so repeated refreshes do not erode them. A present `scope` in the response still replaces the prior grant
Adds regression tests: a refresh omitting `scope` preserves the prior scopes, and a present `scope` overrides them
This commit is contained in:
parent
596abecb09
commit
ea11a97a1e
4 changed files with 46 additions and 4 deletions
|
|
@ -64,7 +64,8 @@ class AuthorizationCodeRefresher:
|
|||
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.
|
||||
the response replaces the old one; an omitted one is carried forward, as are the recorded scopes
|
||||
when the response omits ``scope``.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
|
|
@ -107,13 +108,14 @@ class AuthorizationCodeRefresher:
|
|||
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"))
|
||||
scopes = _parse_scopes(body.get("scope")) or token.scopes
|
||||
|
||||
await self._persist(
|
||||
user_id, server_id, access_token, new_refresh, expires_in, scopes
|
||||
user_id, server_id, access_token, new_refresh, expires_in, scopes or None
|
||||
)
|
||||
return OAuthToken(
|
||||
access_token=access_token,
|
||||
expires_at=self._clock() + expires_in if expires_in is not None else None,
|
||||
refresh_token=new_refresh,
|
||||
scopes=scopes,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -31,15 +31,20 @@ class OAuthToken:
|
|||
with each mode); it is never minted into a header directly. ``repr`` masks both secrets so a
|
||||
stray log line cannot leak them (the values are still plain ``str`` for the header path, since
|
||||
``SecretStr`` resolves as unknown under this repo's basedpyright).
|
||||
|
||||
``scopes`` is the recorded grant. A refresh response that omits ``scope`` (RFC 6749 §5.1: an
|
||||
omitted ``scope`` means unchanged) carries the prior value forward, so a refresh never silently
|
||||
drops it; the resolver itself does not read it.
|
||||
"""
|
||||
|
||||
access_token: str
|
||||
expires_at: float | None = None
|
||||
refresh_token: str | None = None
|
||||
scopes: tuple[str, ...] = ()
|
||||
|
||||
def __repr__(self) -> str:
|
||||
has_refresh = self.refresh_token is not None
|
||||
return f"OAuthToken(access_token=***, expires_at={self.expires_at!r}, has_refresh_token={has_refresh})"
|
||||
return f"OAuthToken(access_token=***, expires_at={self.expires_at!r}, has_refresh_token={has_refresh}, scopes={self.scopes!r})"
|
||||
|
||||
|
||||
class TokenStoreUnavailable(Exception):
|
||||
|
|
|
|||
|
|
@ -33,6 +33,12 @@ def _iso_to_epoch(expires_at: str) -> float | None:
|
|||
return dt.timestamp()
|
||||
|
||||
|
||||
def _to_scopes(raw: object) -> tuple[str, ...]:
|
||||
if isinstance(raw, (list, tuple)):
|
||||
return tuple(s for s in raw if isinstance(s, str))
|
||||
return ()
|
||||
|
||||
|
||||
def _to_oauth_token(payload: dict[str, object]) -> OAuthToken | None:
|
||||
access_token = payload.get("access_token")
|
||||
if not isinstance(access_token, str):
|
||||
|
|
@ -43,6 +49,7 @@ def _to_oauth_token(payload: dict[str, object]) -> OAuthToken | None:
|
|||
access_token=access_token,
|
||||
expires_at=_iso_to_epoch(expires_at) if isinstance(expires_at, str) else None,
|
||||
refresh_token=refresh_token if isinstance(refresh_token, str) else None,
|
||||
scopes=_to_scopes(payload.get("scopes")),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -154,3 +154,31 @@ async def test_unrotated_refresh_token_is_carried_forward():
|
|||
) # 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"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unrecorded_scope_is_carried_forward():
|
||||
persisted = []
|
||||
refresher = _refresher(body={"access_token": "new-at"}, persist_sink=persisted)
|
||||
token = await refresher.refresh(
|
||||
"a", "s", OAuthToken("old", refresh_token="rt", scopes=("read", "write"))
|
||||
)
|
||||
assert token is not None
|
||||
# response omitted "scope" -> the user's recorded grant is preserved, not dropped
|
||||
assert token.scopes == ("read", "write")
|
||||
assert persisted[0][5] == ("read", "write")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returned_scope_overrides_prior_when_present():
|
||||
persisted = []
|
||||
refresher = _refresher(
|
||||
body={"access_token": "new-at", "scope": "read"},
|
||||
persist_sink=persisted,
|
||||
)
|
||||
token = await refresher.refresh(
|
||||
"a", "s", OAuthToken("old", refresh_token="rt", scopes=("read", "write"))
|
||||
)
|
||||
assert token is not None
|
||||
assert token.scopes == ("read",) # a present scope replaces the prior grant
|
||||
assert persisted[0][5] == ("read",)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue