mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
b9a2589124
commit
ea1c6135c2
2 changed files with 135 additions and 0 deletions
|
|
@ -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)
|
||||
|
|
@ -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 == []
|
||||
Loading…
Add table
Reference in a new issue