mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
2e69708ef8
commit
305e70510c
2 changed files with 158 additions and 11 deletions
|
|
@ -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,
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")),
|
||||
]
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue