mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
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:
parent
a871c85bc7
commit
108e75c3cd
4 changed files with 299 additions and 12 deletions
|
|
@ -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]: ...
|
||||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue