diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/clock.py b/litellm/proxy/gateway/mcp/outbound_credentials/clock.py new file mode 100644 index 00000000000..46f7a1b58d8 --- /dev/null +++ b/litellm/proxy/gateway/mcp/outbound_credentials/clock.py @@ -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) diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py b/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py index 972d694025a..53434665a37 100644 --- a/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py +++ b/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py @@ -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") diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/token_refresher.py b/litellm/proxy/gateway/mcp/outbound_credentials/token_refresher.py new file mode 100644 index 00000000000..1368b311028 --- /dev/null +++ b/litellm/proxy/gateway/mcp/outbound_credentials/token_refresher.py @@ -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]: ... diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/token_store.py b/litellm/proxy/gateway/mcp/outbound_credentials/token_store.py new file mode 100644 index 00000000000..b5742fd8f91 --- /dev/null +++ b/litellm/proxy/gateway/mcp/outbound_credentials/token_store.py @@ -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 + ) diff --git a/tests/mcp_tests/gateway/test_resolver.py b/tests/mcp_tests/gateway/test_resolver.py index a3a61077b79..c9decbb0ac9 100644 --- a/tests/mcp_tests/gateway/test_resolver.py +++ b/tests/mcp_tests/gateway/test_resolver.py @@ -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