From 305e70510cdc07faee2b7602549344ae5f5b2022 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Wed, 24 Jun 2026 20:33:29 -0700 Subject: [PATCH] feat(mcp): implement the authorization_code resolver arm Resolve a user's authorization_code token through the injected OAuthTokenStore: present -> Authorization: Bearer ; absent -> the RFC 9728 WWW-Authenticate OAuth challenge; store unavailable -> the same challenge (not a 500), since a transient outage is not a definite absence. UpstreamCredentialProvider gains the oauth_token_store collaborator (fail-closed null default); per-subject isolation comes from keying the fetch on subject_id. Not live until to_server_spec maps authorization_code and a v1-backed token source is wired (next steps). --- .../outbound_credentials/resolver.py | 78 ++++++++++++++-- .../outbound_credentials/test_resolver.py | 91 ++++++++++++++++++- 2 files changed, 158 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py index 969bbf01ec8..134e7438d4e 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py @@ -7,13 +7,15 @@ no precedence cascade. It is wildcard-free with an `assert_never` tail, so addin an arm fails the type gate (basedpyright `reportMatchNotExhaustive`); a bypassed gate fails loudly at runtime instead of returning `None`. -`none` and `api_key` (shared-key source) are live; the remaining arms are `not_implemented` -stubs that each land in a follow-up PR with their injected seam. The self-contained arms read -straight from the config and need no collaborator. Pure v2: no imports from v1. +`none` and `api_key` (shared-key source) are live, as is `authorization_code`, which reads the +user's token from the injected `OAuthTokenStore`. The remaining arms are `not_implemented` stubs +that each land in a follow-up PR with their seam. Pure v2: no imports from v1. """ from __future__ import annotations +from typing import Optional + import httpx from typing_extensions import assert_never @@ -21,6 +23,11 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth impo NoOpAuth, StaticHeaderAuth, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( + OAuthToken, + OAuthTokenStore, + TokenStoreUnavailable, +) from litellm.proxy._experimental.mcp_server.outbound_credentials.result import ( Error, Ok, @@ -43,13 +50,26 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( ) +class _NullOAuthTokenStore: + """Fail-closed default: with no token store wired, every user reads as not authorized.""" + + async def fetch(self, user_id: str, server_id: str) -> Optional[OAuthToken]: + return None + + class UpstreamCredentialProvider: """Produces the one `httpx.Auth` for a `(subject, upstream)` pair, per declared mode. - Collaborators (the per-mode credential stores and token fetchers) are injected as each arm - is built; the live `none` and `api_key`-shared arms read from the config and need none. + Collaborators (the per-mode credential stores and token fetchers) are injected as each arm is + built; the live `none` and `api_key`-shared arms read from the config and need none, while + `authorization_code` reads the user's token from the injected `OAuthTokenStore`. """ + def __init__(self, oauth_token_store: Optional[OAuthTokenStore] = None) -> None: + self._oauth_token_store: OAuthTokenStore = ( + oauth_token_store or _NullOAuthTokenStore() + ) + async def resolve_credentials( self, subject: Subject, server: ServerSpec ) -> Result[httpx.Auth, CredError]: @@ -65,7 +85,14 @@ class UpstreamCredentialProvider: case TokenExchangeConfig(): return _not_implemented(AuthSpecKind.token_exchange) case AuthorizationCodeConfig(): - return _not_implemented(AuthSpecKind.authorization_code) + token = await self._authz_token(subject, server) + if token is None: + return Error(_oauth_challenge(server.server_id)) + return Ok( + StaticHeaderAuth( + f"Bearer {token.access_token}", header_name="Authorization" + ) + ) case AwsSigV4Config(): return _not_implemented(AuthSpecKind.aws_sigv4) assert_never(server.config) @@ -86,8 +113,47 @@ class UpstreamCredentialProvider: ) assert_never(config.key_source) + async def _authz_token( + self, subject: Subject, server: ServerSpec + ) -> Optional[OAuthToken]: + """The user's authorization_code token, or None when absent or the store is unreachable. + + A store outage is mapped to None (the OAuth challenge), not raised, so a transient outage + does not 500; it is the store, not this resolver, that declines to cache the failure. + """ + try: + return await self._oauth_token_store.fetch( + subject.subject_id, server.server_id + ) + except TokenStoreUnavailable: + return None + def _not_implemented(kind: AuthSpecKind) -> Result[httpx.Auth, CredError]: return Error( CredError.of_not_implemented(f"{kind.value}: resolver arm not implemented yet") ) + + +_OAUTH_WWW_AUTHENTICATE = ( + 'Bearer resource_metadata="/.well-known/oauth-protected-resource"' +) + + +def _oauth_challenge(server_id: str) -> CredError: + """The 401 an authorization_code server returns when the user has no usable token. + + Carries the RFC 9728 ``WWW-Authenticate`` challenge that drives the OAuth flow, plus an + ``authorization_required`` body. The exact body is reconciled with v1 when the v1-backed token + source lands. + """ + message = "Authorization required: complete the OAuth flow for this server." + return CredError.of_unauthorized( + message, + www_authenticate=_OAUTH_WWW_AUTHENTICATE, + body={ + "error": "authorization_required", + "server_id": server_id, + "message": message, + }, + ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py index 75be6dfc157..b1d14c639c6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py @@ -1,9 +1,9 @@ """Tests for the resolver dispatch: live arms produce auth, stubbed arms fail closed. -`none` and `api_key` (shared-key source) are implemented; every other arm, plus the `api_key` -BYOK source, returns a typed `not_implemented` error until its mode lands. Parametrizing the -stubs over one config each also guards reachability: a dropped `case` would hit `assert_never` -and raise instead of returning the stub. +`none`, `api_key` (shared-key source), and `authorization_code` are implemented; every other arm, +plus the `api_key` BYOK source, returns a typed `not_implemented` error until its mode lands. +Parametrizing the stubs over one config each also guards reachability: a dropped `case` would hit +`assert_never` and raise instead of returning the stub. """ import httpx @@ -28,6 +28,10 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials import ( TokenExchangeConfig, UpstreamCredentialProvider, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( + OAuthToken, + TokenStoreUnavailable, +) _SUBJECT = Subject(tenant_id="", subject_id="") @@ -84,12 +88,89 @@ async def test_api_key_shared_honors_authorization_scheme(): assert _emitted(result.ok)["Authorization"] == "Bearer tok" +class _FakeTokenStore: + """An OAuthTokenStore returning a canned per-user token (None == not authorized).""" + + def __init__(self, by_user: dict) -> None: + self._by_user = by_user + + async def fetch(self, user_id: str, server_id: str): + return self._by_user.get((user_id, server_id)) + + +@pytest.mark.asyncio +async def test_authorization_code_emits_bearer_for_a_stored_token(): + store = _FakeTokenStore({("alice", "s"): OAuthToken(access_token="at-alice")}) + result = await UpstreamCredentialProvider( + oauth_token_store=store + ).resolve_credentials( + Subject(tenant_id="", subject_id="alice"), _spec(AuthorizationCodeConfig()) + ) + assert isinstance(result, Ok) + assert _emitted(result.ok)["Authorization"] == "Bearer at-alice" + + +@pytest.mark.asyncio +async def test_authorization_code_without_token_is_unauthorized_with_challenge(): + result = await UpstreamCredentialProvider( + oauth_token_store=_FakeTokenStore({}) + ).resolve_credentials( + Subject(tenant_id="", subject_id="alice"), _spec(AuthorizationCodeConfig()) + ) + assert isinstance(result, Error) + assert result.error.tag == "unauthorized" + challenge = result.error.unauthorized + assert challenge.www_authenticate is not None + assert challenge.body is not None + assert challenge.body["error"] == "authorization_required" + + +@pytest.mark.asyncio +async def test_authorization_code_store_unavailable_surfaces_the_challenge(): + class _Unavailable: + async def fetch(self, user_id: str, server_id: str): + raise TokenStoreUnavailable("down") + + result = await UpstreamCredentialProvider( + oauth_token_store=_Unavailable() + ).resolve_credentials( + Subject(tenant_id="", subject_id="alice"), _spec(AuthorizationCodeConfig()) + ) + assert isinstance(result, Error) + assert result.error.tag == "unauthorized" + + +@pytest.mark.asyncio +async def test_authorization_code_with_no_store_wired_is_unauthorized(): + result = await UpstreamCredentialProvider().resolve_credentials( + Subject(tenant_id="", subject_id="alice"), _spec(AuthorizationCodeConfig()) + ) + assert isinstance(result, Error) + assert result.error.tag == "unauthorized" + + +@pytest.mark.asyncio +async def test_authorization_code_isolates_by_subject(): + store = _FakeTokenStore({("alice", "s"): OAuthToken(access_token="at-alice")}) + provider = UpstreamCredentialProvider(oauth_token_store=store) + alice = await provider.resolve_credentials( + Subject(tenant_id="", subject_id="alice"), _spec(AuthorizationCodeConfig()) + ) + bob = await provider.resolve_credentials( + Subject(tenant_id="", subject_id="bob"), _spec(AuthorizationCodeConfig()) + ) + assert ( + isinstance(alice, Ok) + and _emitted(alice.ok)["Authorization"] == "Bearer at-alice" + ) + assert isinstance(bob, Error) and bob.error.tag == "unauthorized" + + _STUBBED = [ ("api_key_byok", ApiKeyConfig(key_source=Byok())), ("passthrough", PassthroughConfig()), ("client_credentials", ClientCredentialsConfig()), ("token_exchange", TokenExchangeConfig()), - ("authorization_code", AuthorizationCodeConfig()), ("aws_sigv4", AwsSigV4Config(region="us-east-1")), ]