feat(mcp): implement the authorization_code arm (per-user OAuth token)

Per-user 3LO: resolve() reads the stored upstream token for (tenant, subject, server,
resource) from an injected async TokenStore, returns it as a Bearer if valid, refreshes it
proactively within a 60s window via an injected TokenRefresher (the real body is SDK-backed;
a fake is used in tests), and fails closed with unauthorized when there is no token or it is
expired without a refresh token. Refresh failures surface the refresher's own CredError
(unauthorized when the grant is rejected, upstream_unavailable when the endpoint is
unreachable). The OAuth dance that populates the store is the separate AS surface (S7); this
arm only reads and refreshes, and never replays the inbound caller bearer.

Design A (explicit lookup + proactive refresh returning a Bearer snapshot, refresher backed
by the SDK) over returning the SDK OAuthClientProvider directly: testable, explicit lifecycle,
clean fail-closed Result, with reactive-401 left to the transport. Adds StoredToken/TokenKey +
TokenStore (InMemoryTokenStore), Clock (SystemClock), and the TokenRefresher port; the provider
now takes these by DI. Access and refresh tokens are SecretStr. 29 tests, gates green.
This commit is contained in:
Tin Chi Lo 2026-06-17 17:49:32 -07:00
parent 5cde452223
commit 0d6b9ab86f
5 changed files with 339 additions and 24 deletions

View file

@ -0,0 +1,19 @@
"""Injectable clock, so expiry / proactive-refresh decisions are deterministic in tests.
`resolve()` never calls `datetime.now()` directly; it reads `Clock.now()`, and tests inject a
fixed clock. Part of the leaf-adapter surface the S0 chassis later bundles into `GatewayDeps`.
"""
from __future__ import annotations
from datetime import datetime, timezone
from typing import Protocol
class Clock(Protocol):
def now(self) -> datetime: ...
class SystemClock:
def now(self) -> datetime:
return datetime.now(timezone.utc)

View file

@ -9,20 +9,26 @@ own fully-typed config with every field guaranteed present — no `None`-checks.
wildcard-free with an `assert_never` tail, so adding a mode without an arm fails the type
gate, and a bypassed gate fails loudly at runtime instead of returning `None`.
Implemented: `none`, `passthrough`, and `api_key` (shared from config, or per-user / BYOK
pulled from the injected `CredentialStore`). The OAuth-flow and signing modes are typed
stubs that fail closed until their collaborators (token store, OAuth providers, RFC 8693
exchanger, SigV4 signer) are injected.
Implemented: `none`, `passthrough`, `api_key` (shared from config, or per-user / BYOK pulled
from the injected `CredentialStore`), and `authorization_code` (per-user token read from the
injected `TokenStore`, refreshed proactively via the `TokenRefresher`). The `client_credentials`,
`token_exchange`, and `aws_sigv4` modes are typed stubs that fail closed until their
collaborators are injected.
"""
from __future__ import annotations
from datetime import timedelta
import httpx
from typing_extensions import assert_never
from ..result import Error, Ok, Result
from .clock import Clock
from .credential_store import CredentialKey, CredentialStore
from .httpx_auth import NoOpAuth, StaticHeaderAuth
from .token_refresher import TokenRefresher
from .token_store import StoredToken, TokenKey, TokenStore
from .types import (
ApiKeyConfig,
AuthorizationCodeConfig,
@ -40,12 +46,24 @@ from .types import (
TokenExchangeConfig,
)
# Refresh proactively once the token is within this window of expiry.
_REFRESH_BUFFER = timedelta(seconds=60)
class UpstreamCredentialProvider:
"""Produces the one `httpx.Auth` for a `(subject, upstream)` pair, per declared mode."""
def __init__(self, credential_store: CredentialStore) -> None:
def __init__(
self,
credential_store: CredentialStore,
token_store: TokenStore,
token_refresher: TokenRefresher,
clock: Clock,
) -> None:
self._credential_store = credential_store
self._token_store = token_store
self._token_refresher = token_refresher
self._clock = clock
async def resolve(
self, subject: Subject, server: ServerSpec
@ -125,12 +143,46 @@ class UpstreamCredentialProvider:
StaticHeaderAuth(f"Bearer {subject.inbound_token.get_secret_value()}")
)
# --- arms awaiting their collaborators (typed stubs, fail closed) ----------------------
async def _authorization_code(
self, subject: Subject, server: ServerSpec, config: AuthorizationCodeConfig
) -> Result[httpx.Auth, CredError]:
return _todo(AuthSpecKind.authorization_code)
# Per-user 3LO: read the stored token, refresh proactively near expiry, or fail closed
# so the edge returns a 401 that starts the OAuth dance (the AS surface writes the token
# this reads). The inbound caller bearer is never sent upstream.
key = TokenKey(
tenant_id=subject.tenant_id,
subject_id=subject.subject_id,
server_id=server.server_id,
resource=server.resource,
)
token = await self._token_store.get(key)
if token is None:
return Error(
CredError.of_unauthorized(
"authorization_code: no stored token; start the OAuth flow"
)
)
if not self._is_near_expiry(token):
return Ok(_bearer(token))
if token.refresh_token is None:
return Error(
CredError.of_unauthorized(
"authorization_code: token expired with no refresh token; re-authenticate"
)
)
refreshed = await self._token_refresher.refresh(config, token.refresh_token)
match refreshed:
case Ok(new_token):
await self._token_store.put(key, new_token)
return Ok(_bearer(new_token))
case Error(err):
return Error(err)
assert_never(refreshed)
def _is_near_expiry(self, token: StoredToken) -> bool:
return self._clock.now() >= token.expires_at - _REFRESH_BUFFER
# --- arms awaiting their collaborators (typed stubs, fail closed) ----------------------
async def _client_credentials(
self, subject: Subject, server: ServerSpec, config: ClientCredentialsConfig
) -> Result[httpx.Auth, CredError]:
@ -147,6 +199,10 @@ class UpstreamCredentialProvider:
return _todo(AuthSpecKind.aws_sigv4)
def _bearer(token: StoredToken) -> StaticHeaderAuth:
return StaticHeaderAuth(f"Bearer {token.access_token.get_secret_value()}")
def _todo(kind: AuthSpecKind) -> Result[httpx.Auth, CredError]:
return Error(
CredError.of_misconfigured(f"{kind.value}: resolver arm not implemented yet")

View file

@ -0,0 +1,29 @@
"""The OAuth refresh-token grant, isolated behind a port.
`resolve()` owns the orchestration (which token, expiry decision, persist, fail closed); the
actual RFC 6749 refresh exchange is delegated here so it is not hand-rolled. The production
body is backed by the MCP SDK / a standard OAuth client; a fake is injected in tests. Async
network I/O, so it lands behind this Protocol with its real body wired later.
"""
from __future__ import annotations
from typing import Protocol
from pydantic import SecretStr
from ..result import Result
from .token_store import StoredToken
from .types import AuthorizationCodeConfig, CredError
class TokenRefresher(Protocol):
"""Exchanges a refresh token for a fresh `StoredToken`, or fails closed.
Returns `unauthorized` when the grant is rejected (refresh token revoked/expired, the user
must re-authenticate) and `upstream_unavailable` when the token endpoint cannot be reached.
"""
async def refresh(
self, config: AuthorizationCodeConfig, refresh_token: SecretStr
) -> Result[StoredToken, CredError]: ...

View file

@ -0,0 +1,57 @@
"""The per-subject OAuth token store for the `authorization_code` arm.
Holds the user's stored upstream token (access + optional refresh + expiry), keyed by
`(tenant, subject, server, resource)` so tokens are per-user and audience-bound (RFC 8707).
The AS surface writes it during the OAuth dance; `resolve()` reads it. Async because the
durable body queries Prisma / Redis on LiteLLM's async stack; `InMemoryTokenStore` is a
working body for tests and local wiring.
"""
from __future__ import annotations
from datetime import datetime
from typing import Protocol
from pydantic import BaseModel, ConfigDict, SecretStr
class TokenKey(BaseModel):
"""Identifies a per-user upstream token. Per-tenant / per-user / per-audience isolation."""
model_config = ConfigDict(frozen=True)
tenant_id: str
subject_id: str
server_id: str
resource: str # RFC 8707 audience the token is bound to
class StoredToken(BaseModel):
"""A user's upstream OAuth token as persisted. Secrets are `SecretStr` so they never log."""
model_config = ConfigDict(frozen=True)
access_token: SecretStr
expires_at: datetime
refresh_token: SecretStr | None = None
class TokenStore(Protocol):
"""Persists and retrieves per-`(subject, server, resource)` OAuth tokens."""
async def get(self, key: TokenKey) -> StoredToken | None: ...
async def put(self, key: TokenKey, token: StoredToken) -> None: ...
class InMemoryTokenStore:
"""A working in-memory `TokenStore` for tests and local wiring."""
def __init__(self, seeded: dict[TokenKey, StoredToken] | None = None) -> None:
self._tokens: dict[TokenKey, StoredToken] = dict(seeded or {})
async def get(self, key: TokenKey) -> StoredToken | None:
return self._tokens.get(key)
async def put(self, key: TokenKey, token: StoredToken) -> None:
self._tokens[key] = (
token # mutable-ok: an in-memory store's backing must be mutable
)

View file

@ -4,10 +4,13 @@ Clean-room litmus: every case constructs `Subject` / `ServerSpec` directly, with
fixtures. If an arm could not be exercised without a v1 request object, the seam has leaked.
"""
from datetime import datetime, timedelta, timezone
import httpx
import pytest
from pydantic import ValidationError
from pydantic import SecretStr, ValidationError
from litellm.proxy.gateway.mcp._spike_exhaustiveness import http_status
from litellm.proxy.gateway.mcp.outbound_credentials.credential_store import (
CredentialKey,
InMemoryCredentialStore,
@ -16,9 +19,17 @@ from litellm.proxy.gateway.mcp.outbound_credentials.httpx_auth import (
NoOpAuth,
StaticHeaderAuth,
)
from litellm.proxy.gateway.mcp._spike_exhaustiveness import http_status
from litellm.proxy.gateway.mcp.outbound_credentials.resolver import (
UpstreamCredentialProvider,
)
from litellm.proxy.gateway.mcp.outbound_credentials.token_store import (
InMemoryTokenStore,
StoredToken,
TokenKey,
)
from litellm.proxy.gateway.mcp.outbound_credentials.types import (
ApiKeyConfig,
AuthorizationCodeConfig,
AuthSpecKind,
Byok,
CredError,
@ -29,17 +40,62 @@ from litellm.proxy.gateway.mcp.outbound_credentials.types import (
SharedKey,
Subject,
)
from litellm.proxy.gateway.mcp.outbound_credentials.resolver import (
UpstreamCredentialProvider,
)
from litellm.proxy.gateway.mcp.result import Error, Ok
from litellm.proxy.gateway.mcp.result import Error, Ok, Result
PROVIDER = UpstreamCredentialProvider(InMemoryCredentialStore())
RESOURCE = "https://up.example/mcp"
NOW = datetime(2026, 6, 17, 12, 0, 0, tzinfo=timezone.utc)
SUBJECT = Subject(tenant_id="t1", subject_id="u1")
class FixedClock:
def now(self) -> datetime:
return NOW
class FakeRefresher:
def __init__(self, result: Result[StoredToken, CredError]) -> None:
self._result = result
async def refresh(
self, config: AuthorizationCodeConfig, refresh_token: SecretStr
) -> Result[StoredToken, CredError]:
return self._result
def _provider(
*,
credential_store: InMemoryCredentialStore | None = None,
token_store: InMemoryTokenStore | None = None,
refresher: FakeRefresher | None = None,
clock: FixedClock | None = None,
) -> UpstreamCredentialProvider:
return UpstreamCredentialProvider(
credential_store=credential_store or InMemoryCredentialStore(),
token_store=token_store or InMemoryTokenStore(),
token_refresher=refresher
or FakeRefresher(Error(CredError.of_upstream_unavailable("unused"))),
clock=clock or FixedClock(),
)
PROVIDER = _provider()
def _spec(config: object) -> ServerSpec:
return ServerSpec(server_id="s1", resource="https://up.example/mcp", config=config) # type: ignore[arg-type]
return ServerSpec(server_id="s1", resource=RESOURCE, config=config) # type: ignore[arg-type]
def _token_key() -> TokenKey:
return TokenKey(tenant_id="t1", subject_id="u1", server_id="s1", resource=RESOURCE)
def _authz_cfg() -> AuthorizationCodeConfig:
return AuthorizationCodeConfig(
client_id="c",
client_secret="s",
authorization_url="https://idp/auth",
token_url="https://idp/token",
)
def _applied_headers(auth: httpx.Auth) -> httpx.Headers:
@ -102,7 +158,7 @@ async def test_api_key_per_user_pulls_the_subject_credential(source: object):
store = InMemoryCredentialStore(
{CredentialKey(tenant_id="t1", subject_id="u1", server_id="s1"): "user-secret"}
)
provider = UpstreamCredentialProvider(store)
provider = _provider(credential_store=store)
result = await provider.resolve(SUBJECT, _spec(ApiKeyConfig(key_source=source))) # type: ignore[arg-type]
assert isinstance(result, Ok)
assert _applied_headers(result.ok)["Authorization"] == "Bearer user-secret"
@ -129,7 +185,7 @@ async def test_api_key_per_user_isolated_by_subject():
store = InMemoryCredentialStore(
{CredentialKey(tenant_id="t1", subject_id="u1", server_id="s1"): "u1-secret"}
)
provider = UpstreamCredentialProvider(store)
provider = _provider(credential_store=store)
other = Subject(tenant_id="t1", subject_id="u2")
result = await provider.resolve(other, _spec(ApiKeyConfig(key_source=Byok())))
assert isinstance(result, Error)
@ -171,13 +227,6 @@ async def test_self_contained_arms_never_read_the_inbound_token():
@pytest.mark.parametrize(
"config",
[
{
"kind": "authorization_code",
"client_id": "c",
"client_secret": "s",
"authorization_url": "https://idp/auth",
"token_url": "https://idp/token",
},
{
"kind": "client_credentials",
"client_id": "c",
@ -196,3 +245,108 @@ async def test_unimplemented_arms_fail_closed(config: dict):
result = await PROVIDER.resolve(SUBJECT, _spec(config))
assert isinstance(result, Error)
assert result.error.tag == "misconfigured"
async def test_authorization_code_returns_a_valid_stored_token():
store = InMemoryTokenStore(
{
_token_key(): StoredToken(
access_token="valid", expires_at=NOW + timedelta(hours=1)
)
}
)
result = await _provider(token_store=store).resolve(SUBJECT, _spec(_authz_cfg()))
assert isinstance(result, Ok)
assert _applied_headers(result.ok)["Authorization"] == "Bearer valid"
async def test_authorization_code_without_a_token_fails_closed():
# No stored token -> unauthorized, which the edge turns into the 401 that starts the dance.
result = await _provider().resolve(SUBJECT, _spec(_authz_cfg()))
assert isinstance(result, Error)
assert result.error.tag == "unauthorized"
async def test_authorization_code_expired_without_refresh_fails_closed():
store = InMemoryTokenStore(
{
_token_key(): StoredToken(
access_token="old", expires_at=NOW - timedelta(minutes=1)
)
}
)
result = await _provider(token_store=store).resolve(SUBJECT, _spec(_authz_cfg()))
assert isinstance(result, Error)
assert result.error.tag == "unauthorized"
async def test_authorization_code_refreshes_proactively_near_expiry():
store = InMemoryTokenStore(
{
_token_key(): StoredToken(
access_token="old",
expires_at=NOW + timedelta(seconds=30), # within the 60s refresh buffer
refresh_token="r",
)
}
)
fresh = StoredToken(
access_token="new", expires_at=NOW + timedelta(hours=1), refresh_token="r2"
)
provider = _provider(token_store=store, refresher=FakeRefresher(Ok(fresh)))
result = await provider.resolve(SUBJECT, _spec(_authz_cfg()))
assert isinstance(result, Ok)
assert _applied_headers(result.ok)["Authorization"] == "Bearer new"
persisted = await store.get(_token_key())
assert persisted is not None
assert persisted.access_token.get_secret_value() == "new"
async def test_authorization_code_refresh_rejected_fails_closed():
store = InMemoryTokenStore(
{
_token_key(): StoredToken(
access_token="old",
expires_at=NOW - timedelta(minutes=1),
refresh_token="r",
)
}
)
provider = _provider(
token_store=store,
refresher=FakeRefresher(Error(CredError.of_unauthorized("refresh revoked"))),
)
result = await provider.resolve(SUBJECT, _spec(_authz_cfg()))
assert isinstance(result, Error)
assert result.error.tag == "unauthorized"
async def test_authorization_code_refresh_unreachable_is_upstream_unavailable():
store = InMemoryTokenStore(
{
_token_key(): StoredToken(
access_token="old",
expires_at=NOW - timedelta(minutes=1),
refresh_token="r",
)
}
)
provider = _provider(
token_store=store,
refresher=FakeRefresher(
Error(CredError.of_upstream_unavailable("token endpoint timeout"))
),
)
result = await provider.resolve(SUBJECT, _spec(_authz_cfg()))
assert isinstance(result, Error)
assert result.error.tag == "upstream_unavailable"
def test_stored_token_secrets_are_masked():
token = StoredToken(
access_token="ACCESS-SECRET", expires_at=NOW, refresh_token="REFRESH-SECRET"
)
dumped = token.model_dump_json()
assert "ACCESS-SECRET" not in dumped
assert "REFRESH-SECRET" not in dumped
assert "**********" in dumped