feat(mcp): encrypt+serialize codec for caching OAuth tokens in Redis (step 1b §1.5)

The serialize+encrypt boundary a cross-replica cache needs: a plaintext bearer in Redis is a leak, so
encode() encrypts (NaCl in prod via the injected encrypt, identity in tests). Caches only access_token
and expires_at, never the refresh_token - the hot path needs just the bearer, and the long-lived
refresh_token stays in the DB (the refresh path is always a cache miss), matching v1. A decoded token
always has refresh_token=None. Undecryptable (key rotation) or corrupt entries read as a miss.
This commit is contained in:
Tin Chi Lo 2026-06-25 21:54:26 -07:00
parent cec32b4574
commit 9de7157f5d
2 changed files with 91 additions and 0 deletions

View file

@ -0,0 +1,37 @@
"""Serialize + encrypt boundary for caching an OAuth token in a shared (Redis) cache.
A cross-replica cache must serialize the token, and a plaintext bearer in Redis is a leak, so this
encrypts the value (NaCl in production via the injected ``encrypt``, identity in tests). It caches
**only** the ``access_token``: the hot path needs just the bearer, expiry is carried by the cache
entry's TTL (set from the token's ``expires_at`` by the cache), and the long-lived refresh_token stays
in the DB - the refresh path is always a cache miss that re-reads it - so it never reaches Redis. A
decoded token therefore carries only the bearer (``expires_at`` and ``refresh_token`` both None); the
TTL, not the value, bounds its life. An empty/undecryptable blob (e.g. master-key rotation) is a miss.
"""
from __future__ import annotations
from collections.abc import Callable
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
class OAuthTokenCacheCodec:
def __init__(
self,
encrypt: Callable[[str], str],
decrypt: Callable[[str], str | None],
) -> None:
self._encrypt = encrypt
self._decrypt = decrypt
def encode(self, token: OAuthToken) -> str:
return self._encrypt(token.access_token)
def decode(self, blob: str) -> OAuthToken | None:
access_token = self._decrypt(blob)
if not access_token:
return None
return OAuthToken(access_token=access_token, refresh_token=None)

View file

@ -0,0 +1,54 @@
"""Tests for the cache codec: encrypt on encode, drop the refresh_token, round-trip the bearer."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_cache_codec import (
OAuthTokenCacheCodec,
)
def _wrapping_codec():
# A reversible stand-in for NaCl: proves encode() encrypts (output is wrapped) and decode()
# decrypts, without needing a salt key.
return OAuthTokenCacheCodec(
encrypt=lambda s: f"enc:{s}",
decrypt=lambda b: b[4:] if b.startswith("enc:") else None,
)
def test_round_trips_the_access_token():
codec = _wrapping_codec()
token = codec.decode(
codec.encode(OAuthToken(access_token="at-123", expires_at=1234.5))
)
assert token is not None
assert token.access_token == "at-123"
def test_encode_encrypts_and_omits_the_refresh_token():
codec = _wrapping_codec()
blob = codec.encode(OAuthToken(access_token="at", refresh_token="super-secret-rt"))
assert blob.startswith("enc:") # encryption was applied
assert (
"super-secret-rt" not in blob
) # the long-lived secret never reaches the cache
decoded = codec.decode(blob)
assert decoded is not None and decoded.refresh_token is None
def test_decoded_token_defers_expiry_to_the_cache_ttl():
codec = _wrapping_codec()
token = codec.decode(codec.encode(OAuthToken(access_token="at", expires_at=999.0)))
# The value carries no expiry; the cache entry's TTL bounds its life instead.
assert token is not None and token.expires_at is None
def test_undecryptable_blob_is_a_miss():
# e.g. master-key rotation makes an old entry unreadable -> treat as a miss, not a crash.
assert _wrapping_codec().decode("not-our-prefix") is None
def test_empty_plaintext_is_a_miss():
codec = OAuthTokenCacheCodec(encrypt=lambda s: s, decrypt=lambda b: b)
assert codec.decode("") is None