diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py new file mode 100644 index 00000000000..458f45367de --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py @@ -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) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_v2_token_store.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_v2_token_store.py new file mode 100644 index 00000000000..97d14b65f5a --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_v2_token_store.py @@ -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 == []