mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
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:
parent
109ba26397
commit
a871c85bc7
4 changed files with 92 additions and 28 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue