mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
feat(mcp): implement the authorization_code arm (per-user OAuth token)
Per-user 3LO: resolve() reads the stored upstream token for (tenant, subject, server, resource) from an injected async TokenStore, returns it as a Bearer if valid, refreshes it proactively within a 60s window via an injected TokenRefresher (the real body is SDK-backed; a fake is used in tests), and fails closed with unauthorized when there is no token or it is expired without a refresh token. Refresh failures surface the refresher's own CredError (unauthorized when the grant is rejected, upstream_unavailable when the endpoint is unreachable). The OAuth dance that populates the store is the separate AS surface (S7); this arm only reads and refreshes, and never replays the inbound caller bearer. Design A (explicit lookup + proactive refresh returning a Bearer snapshot, refresher backed by the SDK) over returning the SDK OAuthClientProvider directly: testable, explicit lifecycle, clean fail-closed Result, with reactive-401 left to the transport. Adds StoredToken/TokenKey + TokenStore (InMemoryTokenStore), Clock (SystemClock), and the TokenRefresher port; the provider now takes these by DI. Access and refresh tokens are SecretStr. 29 tests, gates green.
This commit is contained in:
parent
5cde452223
commit
0d6b9ab86f
5 changed files with 339 additions and 24 deletions
19
litellm/proxy/gateway/mcp/outbound_credentials/clock.py
Normal file
19
litellm/proxy/gateway/mcp/outbound_credentials/clock.py
Normal file
|
|
@ -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)
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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]: ...
|
||||
|
|
@ -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
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue