diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/client_credentials_fetcher.py b/litellm/proxy/gateway/mcp/outbound_credentials/client_credentials_fetcher.py new file mode 100644 index 00000000000..c9ad5b8200d --- /dev/null +++ b/litellm/proxy/gateway/mcp/outbound_credentials/client_credentials_fetcher.py @@ -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]: ... diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py b/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py index fe8cd6fc64b..57802a7fdbe 100644 --- a/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py +++ b/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py @@ -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]: diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/service_token_store.py b/litellm/proxy/gateway/mcp/outbound_credentials/service_token_store.py new file mode 100644 index 00000000000..6be4521dd07 --- /dev/null +++ b/litellm/proxy/gateway/mcp/outbound_credentials/service_token_store.py @@ -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) diff --git a/tests/mcp_tests/gateway/test_resolver.py b/tests/mcp_tests/gateway/test_resolver.py index 219aaaa1953..e9bc8d0b68f 100644 --- a/tests/mcp_tests/gateway/test_resolver.py +++ b/tests/mcp_tests/gateway/test_resolver.py @@ -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())