feat(mcp): implement the authorization_code resolver arm

Resolve a user's authorization_code token through the injected OAuthTokenStore: present ->
Authorization: Bearer <access_token>; 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).
This commit is contained in:
Tin Chi Lo 2026-06-24 20:33:29 -07:00
parent 2e69708ef8
commit 305e70510c
2 changed files with 158 additions and 11 deletions

View file

@ -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,
},
)

View file

@ -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")),
]