diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/credential_store.py b/litellm/proxy/gateway/mcp/outbound_credentials/credential_store.py index 9ab28eff8ad..2cd122493bb 100644 --- a/litellm/proxy/gateway/mcp/outbound_credentials/credential_store.py +++ b/litellm/proxy/gateway/mcp/outbound_credentials/credential_store.py @@ -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)) diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py b/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py index 8b4624b1a6b..fe8cd6fc64b 100644 --- a/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py +++ b/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py @@ -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 diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/token_store.py b/litellm/proxy/gateway/mcp/outbound_credentials/token_store.py index b5742fd8f91..6d36c3bfa8f 100644 --- a/litellm/proxy/gateway/mcp/outbound_credentials/token_store.py +++ b/litellm/proxy/gateway/mcp/outbound_credentials/token_store.py @@ -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) diff --git a/tests/mcp_tests/gateway/test_resolver.py b/tests/mcp_tests/gateway/test_resolver.py index f2600dc6d04..219aaaa1953 100644 --- a/tests/mcp_tests/gateway/test_resolver.py +++ b/tests/mcp_tests/gateway/test_resolver.py @@ -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"