refactor(mcp): stores return Result so a DB outage is distinct from a miss

CredentialStore.get, TokenStore.get, and TokenStore.put returned bare values (T | None /
None), so a store/DB outage could only surface as a raised exception across the seam, and a
read could not distinguish 'not found' from 'backend down'. Change them to return
Result[..., CredError]: Ok(None) is a genuine miss (-> 401/412 to start the dance), while
Error(upstream_unavailable) is a store/DB outage (-> 503), per the plan's fail-closed-loud
invariant. The resolver's api_key and authorization_code arms unwrap with isinstance and
propagate the store error; a put failure on the refresh path also surfaces 503. Adds tests
for the per-user and authorization_code store-error paths (31 tests, gates green).
This commit is contained in:
Tin Chi Lo 2026-06-17 18:20:49 -07:00
parent 109ba26397
commit a871c85bc7
4 changed files with 92 additions and 28 deletions

View file

@ -12,6 +12,9 @@ from typing import Protocol
from pydantic import BaseModel, ConfigDict
from ..result import Ok, Result
from .types import CredError
class CredentialKey(BaseModel):
"""Identifies a per-user secret. Per-tenant / per-user isolation is the key shape."""
@ -25,12 +28,14 @@ class CredentialKey(BaseModel):
class CredentialStore(Protocol):
"""Fetches the per-subject secret for an `api_key` per-user / BYOK server.
Returns a `Result` so a store/DB outage (`Error(upstream_unavailable)`) is distinct from a
genuine miss (`Ok(None)`); the resolver maps the former to 503 and the latter to a 401/412.
Async because the durable body queries Prisma / Redis on LiteLLM's async stack; a
synchronous read would block the event loop. Defined async now, before any caller, so the
signature does not break when that body lands.
synchronous read would block the event loop. Both are settled now, before any caller, so
the signature does not break when that body lands.
"""
async def get(self, key: CredentialKey) -> str | None: ...
async def get(self, key: CredentialKey) -> Result[str | None, CredError]: ...
class InMemoryCredentialStore:
@ -39,5 +44,5 @@ class InMemoryCredentialStore:
def __init__(self, seeded: dict[CredentialKey, str] | None = None) -> None:
self._values: dict[CredentialKey, str] = dict(seeded or {})
async def get(self, key: CredentialKey) -> str | None:
return self._values.get(key)
async def get(self, key: CredentialKey) -> Result[str | None, CredError]:
return Ok(self._values.get(key))

View file

@ -102,7 +102,10 @@ class UpstreamCredentialProvider:
StaticHeaderAuth(config.header_for(source.value.get_secret_value()))
)
case PerUserEnvVar():
value = await self._per_user_value(subject, server)
fetched = await self._per_user_value(subject, server)
if isinstance(fetched, Error):
return Error(fetched.error) # store/DB down -> 503
value = fetched.ok
if value is None:
return Error(
CredError.of_precondition_required(
@ -111,7 +114,10 @@ class UpstreamCredentialProvider:
)
return Ok(StaticHeaderAuth(config.header_for(value)))
case Byok():
value = await self._per_user_value(subject, server)
fetched = await self._per_user_value(subject, server)
if isinstance(fetched, Error):
return Error(fetched.error) # store/DB down -> 503
value = fetched.ok
if value is None:
return Error(
CredError.of_unauthorized(
@ -121,7 +127,9 @@ class UpstreamCredentialProvider:
return Ok(StaticHeaderAuth(config.header_for(value)))
assert_never(config.key_source)
async def _per_user_value(self, subject: Subject, server: ServerSpec) -> str | None:
async def _per_user_value(
self, subject: Subject, server: ServerSpec
) -> Result[str | None, CredError]:
return await self._credential_store.get(
CredentialKey(
tenant_id=subject.tenant_id,
@ -155,7 +163,10 @@ class UpstreamCredentialProvider:
server_id=server.server_id,
resource=server.resource,
)
token = await self._token_store.get(key)
fetched = await self._token_store.get(key)
if isinstance(fetched, Error):
return Error(fetched.error) # store/DB down -> 503
token = fetched.ok
if token is None:
return Error(
CredError.of_unauthorized(
@ -171,13 +182,15 @@ class UpstreamCredentialProvider:
)
)
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)
if isinstance(refreshed, Error):
return Error(
refreshed.error
) # refresh rejected (401) or endpoint down (503)
new_token = refreshed.ok
persisted = await self._token_store.put(key, new_token)
if isinstance(persisted, Error):
return Error(persisted.error) # store/DB down on write -> 503
return Ok(_bearer(new_token))
def _is_near_expiry(self, token: StoredToken) -> bool:
return self._clock.now() >= token.expires_at - _REFRESH_BUFFER

View file

@ -14,6 +14,9 @@ from typing import Protocol
from pydantic import BaseModel, ConfigDict, SecretStr
from ..result import Ok, Result
from .types import CredError
class TokenKey(BaseModel):
"""Identifies a per-user upstream token. Per-tenant / per-user / per-audience isolation."""
@ -35,11 +38,18 @@ class StoredToken(BaseModel):
class TokenStore(Protocol):
"""Persists and retrieves per-`(subject, server, resource)` OAuth tokens."""
"""Persists and retrieves per-`(subject, server, resource)` OAuth tokens.
async def get(self, key: TokenKey) -> StoredToken | None: ...
Both methods return a `Result` so a store/DB outage (`Error(upstream_unavailable)`, 503)
is distinct from a genuine miss (`Ok(None)`, which drives the OAuth dance). Async because
the durable body queries Prisma / Redis on LiteLLM's async stack.
"""
async def put(self, key: TokenKey, token: StoredToken) -> None: ...
async def get(self, key: TokenKey) -> Result[StoredToken | None, CredError]: ...
async def put(
self, key: TokenKey, token: StoredToken
) -> Result[None, CredError]: ...
class InMemoryTokenStore:
@ -48,10 +58,11 @@ class InMemoryTokenStore:
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 get(self, key: TokenKey) -> Result[StoredToken | None, CredError]:
return Ok(self._tokens.get(key))
async def put(self, key: TokenKey, token: StoredToken) -> None:
async def put(self, key: TokenKey, token: StoredToken) -> Result[None, CredError]:
self._tokens[key] = (
token # mutable-ok: an in-memory store's backing must be mutable
)
return Ok(None)

View file

@ -11,8 +11,10 @@ import pytest
from pydantic import SecretStr, ValidationError
from litellm.proxy.gateway.mcp._spike_exhaustiveness import http_status
from litellm.proxy.gateway.mcp.outbound_credentials.clock import Clock
from litellm.proxy.gateway.mcp.outbound_credentials.credential_store import (
CredentialKey,
CredentialStore,
InMemoryCredentialStore,
)
from litellm.proxy.gateway.mcp.outbound_credentials.httpx_auth import (
@ -22,10 +24,14 @@ from litellm.proxy.gateway.mcp.outbound_credentials.httpx_auth import (
from litellm.proxy.gateway.mcp.outbound_credentials.resolver import (
UpstreamCredentialProvider,
)
from litellm.proxy.gateway.mcp.outbound_credentials.token_refresher import (
TokenRefresher,
)
from litellm.proxy.gateway.mcp.outbound_credentials.token_store import (
InMemoryTokenStore,
StoredToken,
TokenKey,
TokenStore,
)
from litellm.proxy.gateway.mcp.outbound_credentials.types import (
ApiKeyConfig,
@ -62,12 +68,25 @@ class FakeRefresher:
return self._result
class FailingCredentialStore:
async def get(self, key: CredentialKey) -> Result[str | None, CredError]:
return Error(CredError.of_upstream_unavailable("credential store down"))
class FailingTokenStore:
async def get(self, key: TokenKey) -> Result[StoredToken | None, CredError]:
return Error(CredError.of_upstream_unavailable("token store down"))
async def put(self, key: TokenKey, token: StoredToken) -> Result[None, CredError]:
return Error(CredError.of_upstream_unavailable("token store down"))
def _provider(
*,
credential_store: InMemoryCredentialStore | None = None,
token_store: InMemoryTokenStore | None = None,
refresher: FakeRefresher | None = None,
clock: FixedClock | None = None,
credential_store: CredentialStore | None = None,
token_store: TokenStore | None = None,
refresher: TokenRefresher | None = None,
clock: Clock | None = None,
) -> UpstreamCredentialProvider:
return UpstreamCredentialProvider(
credential_store=credential_store or InMemoryCredentialStore(),
@ -300,8 +319,9 @@ async def test_authorization_code_refreshes_proactively_near_expiry():
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"
assert isinstance(persisted, Ok)
assert persisted.ok is not None
assert persisted.ok.access_token.get_secret_value() == "new"
async def test_authorization_code_refresh_rejected_fails_closed():
@ -352,3 +372,18 @@ def test_stored_token_secrets_are_masked():
assert "ACCESS-SECRET" not in dumped
assert "REFRESH-SECRET" not in dumped
assert "**********" in dumped
async def test_api_key_per_user_store_error_is_upstream_unavailable():
# A store/DB outage on read is distinct from a miss: 503, not the 401/412 of "not found".
provider = _provider(credential_store=FailingCredentialStore())
result = await provider.resolve(SUBJECT, _spec(ApiKeyConfig(key_source=Byok())))
assert isinstance(result, Error)
assert result.error.tag == "upstream_unavailable"
async def test_authorization_code_store_error_is_upstream_unavailable():
provider = _provider(token_store=FailingTokenStore())
result = await provider.resolve(SUBJECT, _spec(_authz_cfg()))
assert isinstance(result, Error)
assert result.error.tag == "upstream_unavailable"