feat(mcp): implement the client_credentials (M2M) arm

Mints/serves one shared service-account token per server via the client_credentials grant.
Keyed by (server_id, resource) with no subject (ServiceTokenStore, a new port + in-memory
body), so every user shares the identity; the grant itself is delegated to a ClientCredentials
Fetcher port (SDK-backed body, faked in tests). The arm never reads the caller bearer (closes
the v1 M2M auth-bypass by construction), maps a rejected grant to misconfigured (500, the
operator's secret is wrong, no user to re-auth) and an unreachable endpoint to
upstream_unavailable (503).

The cache is treated as a pure optimization, not a source of truth: a read failure degrades to
a fresh mint and a write failure is best-effort (return the valid token, skip caching), because
an M2M token is always re-mintable from the secret - so a cache outage never fails a request
while the token endpoint is up. This is the explicit design choice starred on the Detailed
Phase 1 per-mode page, and it deliberately differs from authorization_code's propagate-on-write.

Tests cover cached-fresh (no fetch), mint+cache on miss, re-mint near expiry, rejected->500,
endpoint-down->503, never-reads-inbound, shared-across-subjects, and the read-degrade /
write-best-effort cache paths. 39 tests, gates green.
This commit is contained in:
Tin Chi Lo 2026-06-17 18:47:54 -07:00
parent a871c85bc7
commit 108e75c3cd
4 changed files with 299 additions and 12 deletions

View file

@ -0,0 +1,27 @@
"""The client_credentials (M2M) grant, isolated behind a port.
`resolve()` owns the orchestration (cache, expiry, fail closed); the RFC 6749 client-credentials
exchange is delegated here so it is not hand-rolled. The production body is backed by the MCP
SDK's `ClientCredentialsOAuthProvider`; a fake is injected in tests. Async network I/O.
"""
from __future__ import annotations
from typing import Protocol
from ..result import Result
from .token_store import StoredToken
from .types import ClientCredentialsConfig, CredError
class ClientCredentialsFetcher(Protocol):
"""Runs the client_credentials grant for a service account, or fails closed.
Returns `misconfigured` when the grant is rejected (bad client_id / secret / scope, an
operator error with no user to re-authenticate) and `upstream_unavailable` when the token
endpoint cannot be reached.
"""
async def fetch(
self, config: ClientCredentialsConfig
) -> Result[StoredToken, CredError]: ...

View file

@ -10,10 +10,11 @@ wildcard-free with an `assert_never` tail, so adding a mode without an arm fails
gate, and a bypassed gate fails loudly at runtime instead of returning `None`.
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 the injected `CredentialStore`), `authorization_code` (per-user token read from the
injected `TokenStore`, refreshed proactively via the `TokenRefresher`), and `client_credentials`
(shared service-account token cached in the `ServiceTokenStore`, minted by the
`ClientCredentialsFetcher`). The `token_exchange` and `aws_sigv4` modes are typed stubs that
fail closed until their collaborators are injected.
"""
from __future__ import annotations
@ -24,9 +25,11 @@ import httpx
from typing_extensions import assert_never
from ..result import Error, Ok, Result
from .client_credentials_fetcher import ClientCredentialsFetcher
from .clock import Clock
from .credential_store import CredentialKey, CredentialStore
from .httpx_auth import NoOpAuth, StaticHeaderAuth
from .service_token_store import ServiceTokenKey, ServiceTokenStore
from .token_refresher import TokenRefresher
from .token_store import StoredToken, TokenKey, TokenStore
from .types import (
@ -59,11 +62,15 @@ class UpstreamCredentialProvider:
token_store: TokenStore,
token_refresher: TokenRefresher,
clock: Clock,
service_token_store: ServiceTokenStore,
client_credentials_fetcher: ClientCredentialsFetcher,
) -> None:
self._credential_store = credential_store
self._token_store = token_store
self._token_refresher = token_refresher
self._clock = clock
self._service_token_store = service_token_store
self._client_credentials_fetcher = client_credentials_fetcher
async def resolve(
self, subject: Subject, server: ServerSpec
@ -195,12 +202,28 @@ class UpstreamCredentialProvider:
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]:
return _todo(AuthSpecKind.client_credentials)
# M2M: one shared service-account token per (server, resource), cached without a
# subject; the caller bearer is never read. The cache is best-effort - re-mint on a
# read failure, ignore a write failure - because the token is always re-mintable.
key = ServiceTokenKey(server_id=server.server_id, resource=server.resource)
cached = await self._service_token_store.get(key)
if isinstance(cached, Ok):
token = cached.ok
if token is not None and not self._is_near_expiry(token):
return Ok(_bearer(token))
fetched = await self._client_credentials_fetcher.fetch(config)
if isinstance(fetched, Error):
return Error(fetched.error) # rejected creds -> 500, endpoint down -> 503
minted = fetched.ok
await self._service_token_store.put(
key, minted
) # best-effort; token is valid anyway
return Ok(_bearer(minted))
# --- arms awaiting their collaborators (typed stubs, fail closed) ----------------------
async def _token_exchange(
self, subject: Subject, server: ServerSpec, config: TokenExchangeConfig
) -> Result[httpx.Auth, CredError]:

View file

@ -0,0 +1,58 @@
"""The shared service-account token cache for the `client_credentials` (M2M) arm.
The M2M token is the same for every user of a server, so it is keyed by `(server_id, resource)`
with no subject. It is a pure optimization cache, not a source of truth: the token is always
re-mintable from the client secret, so the resolver degrades gracefully on a cache outage
(re-mint on a read failure, best-effort on write) rather than failing. Redis-with-TTL later;
`InMemoryServiceTokenStore` is a working body for tests and local wiring.
"""
from __future__ import annotations
from typing import Protocol
from pydantic import BaseModel, ConfigDict
from ..result import Ok, Result
from .token_store import StoredToken
from .types import CredError
class ServiceTokenKey(BaseModel):
"""Identifies a shared service-account token. No subject: one identity for all users."""
model_config = ConfigDict(frozen=True)
server_id: str
resource: str # RFC 8707 audience the token is bound to
class ServiceTokenStore(Protocol):
"""Caches the shared M2M token per `(server_id, resource)`."""
async def get(
self, key: ServiceTokenKey
) -> Result[StoredToken | None, CredError]: ...
async def put(
self, key: ServiceTokenKey, token: StoredToken
) -> Result[None, CredError]: ...
class InMemoryServiceTokenStore:
"""A working in-memory `ServiceTokenStore` for tests and local wiring."""
def __init__(
self, seeded: dict[ServiceTokenKey, StoredToken] | None = None
) -> None:
self._tokens: dict[ServiceTokenKey, StoredToken] = dict(seeded or {})
async def get(self, key: ServiceTokenKey) -> Result[StoredToken | None, CredError]:
return Ok(self._tokens.get(key))
async def put(
self, key: ServiceTokenKey, token: StoredToken
) -> Result[None, CredError]:
self._tokens[key] = (
token # mutable-ok: an in-memory store's backing must be mutable
)
return Ok(None)

View file

@ -11,6 +11,9 @@ import pytest
from pydantic import SecretStr, ValidationError
from litellm.proxy.gateway.mcp._spike_exhaustiveness import http_status
from litellm.proxy.gateway.mcp.outbound_credentials.client_credentials_fetcher import (
ClientCredentialsFetcher,
)
from litellm.proxy.gateway.mcp.outbound_credentials.clock import Clock
from litellm.proxy.gateway.mcp.outbound_credentials.credential_store import (
CredentialKey,
@ -24,6 +27,11 @@ from litellm.proxy.gateway.mcp.outbound_credentials.httpx_auth import (
from litellm.proxy.gateway.mcp.outbound_credentials.resolver import (
UpstreamCredentialProvider,
)
from litellm.proxy.gateway.mcp.outbound_credentials.service_token_store import (
InMemoryServiceTokenStore,
ServiceTokenKey,
ServiceTokenStore,
)
from litellm.proxy.gateway.mcp.outbound_credentials.token_refresher import (
TokenRefresher,
)
@ -38,6 +46,7 @@ from litellm.proxy.gateway.mcp.outbound_credentials.types import (
AuthorizationCodeConfig,
AuthSpecKind,
Byok,
ClientCredentialsConfig,
CredError,
NoneConfig,
PassthroughConfig,
@ -81,12 +90,48 @@ class FailingTokenStore:
return Error(CredError.of_upstream_unavailable("token store down"))
class FakeFetcher:
def __init__(self, result: Result[StoredToken, CredError]) -> None:
self._result = result
self.calls = 0
async def fetch(
self, config: ClientCredentialsConfig
) -> Result[StoredToken, CredError]:
self.calls += 1
return self._result
class FlakyServiceTokenStore:
"""A ServiceTokenStore that can be told to fail reads and/or writes."""
def __init__(self, *, fail_get: bool = False, fail_put: bool = False) -> None:
self._fail_get = fail_get
self._fail_put = fail_put
self._tokens: dict[ServiceTokenKey, StoredToken] = {} # mutable-ok: test double
async def get(self, key: ServiceTokenKey) -> Result[StoredToken | None, CredError]:
if self._fail_get:
return Error(CredError.of_upstream_unavailable("cache read down"))
return Ok(self._tokens.get(key))
async def put(
self, key: ServiceTokenKey, token: StoredToken
) -> Result[None, CredError]:
if self._fail_put:
return Error(CredError.of_upstream_unavailable("cache write down"))
self._tokens[key] = token # mutable-ok: test double
return Ok(None)
def _provider(
*,
credential_store: CredentialStore | None = None,
token_store: TokenStore | None = None,
refresher: TokenRefresher | None = None,
clock: Clock | None = None,
service_token_store: ServiceTokenStore | None = None,
fetcher: ClientCredentialsFetcher | None = None,
) -> UpstreamCredentialProvider:
return UpstreamCredentialProvider(
credential_store=credential_store or InMemoryCredentialStore(),
@ -94,6 +139,9 @@ def _provider(
token_refresher=refresher
or FakeRefresher(Error(CredError.of_upstream_unavailable("unused"))),
clock=clock or FixedClock(),
service_token_store=service_token_store or InMemoryServiceTokenStore(),
client_credentials_fetcher=fetcher
or FakeFetcher(Error(CredError.of_upstream_unavailable("unused"))),
)
@ -117,6 +165,16 @@ def _authz_cfg() -> AuthorizationCodeConfig:
)
def _cc_cfg() -> ClientCredentialsConfig:
return ClientCredentialsConfig(
client_id="c", client_secret="s", token_url="https://idp/token"
)
def _svc_key() -> ServiceTokenKey:
return ServiceTokenKey(server_id="s1", resource=RESOURCE)
def _applied_headers(auth: httpx.Auth) -> httpx.Headers:
request = httpx.Request("POST", "https://up.example/mcp")
return next(auth.auth_flow(request)).headers
@ -247,12 +305,6 @@ async def test_self_contained_arms_never_read_the_inbound_token():
@pytest.mark.parametrize(
"config",
[
{
"kind": "client_credentials",
"client_id": "c",
"client_secret": "s",
"token_url": "https://idp/token",
},
{
"kind": "token_exchange",
"token_exchange_endpoint": "https://idp/token",
@ -374,6 +426,133 @@ def test_stored_token_secrets_are_masked():
assert "**********" in dumped
async def test_client_credentials_uses_cached_fresh_token():
store = InMemoryServiceTokenStore(
{
_svc_key(): StoredToken(
access_token="svc", expires_at=NOW + timedelta(hours=1)
)
}
)
fetcher = FakeFetcher(
Error(CredError.of_upstream_unavailable("must not be called"))
)
result = await _provider(service_token_store=store, fetcher=fetcher).resolve(
SUBJECT, _spec(_cc_cfg())
)
assert isinstance(result, Ok)
assert _applied_headers(result.ok)["Authorization"] == "Bearer svc"
assert fetcher.calls == 0
async def test_client_credentials_mints_and_caches_on_miss():
store = InMemoryServiceTokenStore()
minted = StoredToken(access_token="fresh", expires_at=NOW + timedelta(hours=1))
result = await _provider(
service_token_store=store, fetcher=FakeFetcher(Ok(minted))
).resolve(SUBJECT, _spec(_cc_cfg()))
assert isinstance(result, Ok)
assert _applied_headers(result.ok)["Authorization"] == "Bearer fresh"
cached = await store.get(_svc_key())
assert isinstance(cached, Ok) and cached.ok is not None
assert cached.ok.access_token.get_secret_value() == "fresh"
async def test_client_credentials_remints_near_expiry():
store = InMemoryServiceTokenStore(
{
_svc_key(): StoredToken(
access_token="old", expires_at=NOW + timedelta(seconds=30)
)
}
)
fetcher = FakeFetcher(
Ok(StoredToken(access_token="new", expires_at=NOW + timedelta(hours=1)))
)
result = await _provider(service_token_store=store, fetcher=fetcher).resolve(
SUBJECT, _spec(_cc_cfg())
)
assert isinstance(result, Ok)
assert _applied_headers(result.ok)["Authorization"] == "Bearer new"
assert fetcher.calls == 1
async def test_client_credentials_rejected_grant_is_misconfigured():
fetcher = FakeFetcher(Error(CredError.of_misconfigured("invalid_client")))
result = await _provider(fetcher=fetcher).resolve(SUBJECT, _spec(_cc_cfg()))
assert isinstance(result, Error)
assert result.error.tag == "misconfigured"
async def test_client_credentials_endpoint_down_is_upstream_unavailable():
fetcher = FakeFetcher(
Error(CredError.of_upstream_unavailable("token endpoint timeout"))
)
result = await _provider(fetcher=fetcher).resolve(SUBJECT, _spec(_cc_cfg()))
assert isinstance(result, Error)
assert result.error.tag == "upstream_unavailable"
async def test_client_credentials_never_reads_inbound_token():
# M2M must never forward the caller bearer: the result is identical with or without one.
minted = StoredToken(access_token="svc", expires_at=NOW + timedelta(hours=1))
provider = _provider(fetcher=FakeFetcher(Ok(minted)))
without = await provider.resolve(SUBJECT, _spec(_cc_cfg()))
with_token = Subject(tenant_id="t1", subject_id="u1", inbound_token="leak-me")
present = await provider.resolve(with_token, _spec(_cc_cfg()))
assert isinstance(without, Ok) and isinstance(present, Ok)
assert (
_applied_headers(without.ok)["Authorization"]
== _applied_headers(present.ok)["Authorization"]
== "Bearer svc"
)
async def test_client_credentials_shared_across_subjects():
# Keyed by (server, resource) with no subject: every subject gets the same token.
store = InMemoryServiceTokenStore(
{
_svc_key(): StoredToken(
access_token="shared", expires_at=NOW + timedelta(hours=1)
)
}
)
provider = _provider(service_token_store=store)
u1 = await provider.resolve(SUBJECT, _spec(_cc_cfg()))
u2 = await provider.resolve(
Subject(tenant_id="t1", subject_id="u2"), _spec(_cc_cfg())
)
assert isinstance(u1, Ok) and isinstance(u2, Ok)
assert (
_applied_headers(u1.ok)["Authorization"]
== _applied_headers(u2.ok)["Authorization"]
)
async def test_client_credentials_cache_read_failure_degrades_to_mint():
# A cache read outage must not fail the request; we just mint fresh.
minted = StoredToken(access_token="fresh", expires_at=NOW + timedelta(hours=1))
provider = _provider(
service_token_store=FlakyServiceTokenStore(fail_get=True),
fetcher=FakeFetcher(Ok(minted)),
)
result = await provider.resolve(SUBJECT, _spec(_cc_cfg()))
assert isinstance(result, Ok)
assert _applied_headers(result.ok)["Authorization"] == "Bearer fresh"
async def test_client_credentials_cache_write_failure_is_best_effort():
# A cache write outage must not fail a valid mint; we return it and skip caching.
minted = StoredToken(access_token="fresh", expires_at=NOW + timedelta(hours=1))
provider = _provider(
service_token_store=FlakyServiceTokenStore(fail_put=True),
fetcher=FakeFetcher(Ok(minted)),
)
result = await provider.resolve(SUBJECT, _spec(_cc_cfg()))
assert isinstance(result, Ok)
assert _applied_headers(result.ok)["Authorization"] == "Bearer fresh"
async def test_api_key_per_user_store_error_is_upstream_unavailable():
# A store/DB outage on read is distinct from a miss: 503, not the 401/412 of "not found".
provider = _provider(credential_store=FailingCredentialStore())