feat(mcp): v2-native per-user token read store (step 1b inner store)

Reads the user's persisted authorization_code credential and returns a typed OAuthToken (access
token, epoch expiry, refresh token), validating the decoded blob at this boundary so no Any leaks
past it. The raw inner store that RefreshingTokenStore/CachedOAuthTokenStore wrap; the DB read +
decode collaborator is injected so it stays testable. Not yet wired - V1PerUserTokenStore is still
the composition-root store until the refresher and cross-worker cache land.
This commit is contained in:
Tin Chi Lo 2026-06-25 19:23:24 -07:00
parent b9a2589124
commit ea1c6135c2
2 changed files with 135 additions and 0 deletions

View file

@ -0,0 +1,62 @@
"""v2-native per-user OAuth token read store for the ``authorization_code`` mode.
The raw "inner" store that ``RefreshingTokenStore`` and ``CachedOAuthTokenStore`` wrap: it reads the
user's persisted credential and returns a typed ``OAuthToken`` (access token, epoch expiry, refresh
token), validating the decoded credential blob at this boundary so no ``Any`` leaks past it. It does
not cache or refresh - those are the decorators. This replaces ``V1PerUserTokenStore`` (which handed
the whole read + cache + refresh to v1's core) as step 1b: the ``read_credential`` collaborator is
injected, so the DB/decoding plumbing stays testable and out of this seam.
"""
from __future__ import annotations
from collections.abc import Awaitable, Callable
from datetime import datetime
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
CredentialReader = Callable[[str, str], Awaitable["dict[str, object] | None"]]
def _iso_to_epoch(expires_at: str) -> float | None:
try:
return datetime.fromisoformat(expires_at).timestamp()
except ValueError:
return None
def _to_oauth_token(payload: dict[str, object]) -> OAuthToken | None:
access_token = payload.get("access_token")
if not isinstance(access_token, str):
return None
refresh_token = payload.get("refresh_token")
expires_at = payload.get("expires_at")
return OAuthToken(
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,
)
class V2PerUserTokenStore:
"""``OAuthTokenStore`` that reads the user's persisted authorization_code credential, typed.
The injected ``read_credential`` returns the decoded credential payload for a ``(user, server)``
pair, or ``None`` when the user has not completed OAuth. A backing-store outage surfaces as
``TokenStoreUnavailable`` from the reader, which the arm turns into a challenge rather than a
500, so ``fetch`` lets it propagate. Refresh is the wrapping ``RefreshingTokenStore``'s job, so
the returned token carries ``expires_at`` and ``refresh_token`` for it to act on.
"""
def __init__(self, read_credential: CredentialReader) -> None:
self._read_credential = read_credential
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
if not user_id:
return None
payload = await self._read_credential(user_id, server_id)
if payload is None:
return None
return _to_oauth_token(payload)

View file

@ -0,0 +1,73 @@
"""Tests for the v2-native per-user token read store: validate the decoded blob into a typed token."""
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.v2_token_store import (
V2PerUserTokenStore,
)
def _reader(payload):
async def read(user_id: str, server_id: str):
return payload
return read
@pytest.mark.asyncio
async def test_builds_typed_token_with_epoch_expiry():
store = V2PerUserTokenStore(
_reader(
{
"access_token": "at",
"expires_at": "2099-01-01T00:00:00+00:00",
"refresh_token": "rt",
}
)
)
token = await store.fetch("alice", "srv")
assert token is not None
assert token.access_token == "at"
assert token.refresh_token == "rt"
# ISO string is converted to epoch seconds for RefreshingTokenStore to compare against.
assert token.expires_at == 4070908800.0
@pytest.mark.asyncio
async def test_optional_fields_absent_yield_none():
store = V2PerUserTokenStore(_reader({"access_token": "at"}))
token = await store.fetch("alice", "srv")
assert token is not None
assert token.expires_at is None and token.refresh_token is None
@pytest.mark.asyncio
async def test_unparseable_expiry_is_dropped_not_raised():
store = V2PerUserTokenStore(
_reader({"access_token": "at", "expires_at": "not-a-date"})
)
token = await store.fetch("alice", "srv")
assert token is not None and token.expires_at is None
@pytest.mark.asyncio
async def test_missing_access_token_is_not_authorized():
store = V2PerUserTokenStore(_reader({"expires_at": "2099-01-01T00:00:00+00:00"}))
assert await store.fetch("alice", "srv") is None
@pytest.mark.asyncio
async def test_no_credential_is_not_authorized():
assert await V2PerUserTokenStore(_reader(None)).fetch("alice", "srv") is None
@pytest.mark.asyncio
async def test_empty_user_short_circuits_without_reading():
calls = []
async def read(user_id: str, server_id: str):
calls.append((user_id, server_id))
return {"access_token": "at"}
assert await V2PerUserTokenStore(read).fetch("", "srv") is None
assert calls == []