mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
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:
parent
cec32b4574
commit
9de7157f5d
2 changed files with 91 additions and 0 deletions
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue