From 5cde452223b933853325b2a33b1a54533ba76404 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Wed, 17 Jun 2026 17:16:52 -0700 Subject: [PATCH] refactor(mcp): make resolve() and CredentialStore async (Greptile P2) The real CredentialStore body reads Prisma/Redis on LiteLLM's async stack, and the OAuth-flow and SigV4 arms will do async I/O (token endpoints, RFC 8693, STS). A synchronous resolver would block the event loop or force a breaking signature change once a caller exists. Define it async now, before any runtime caller: CredentialStore.get, resolve(), and every arm are async, and the tests await resolve() (asyncio_mode is auto). await sequences resolution before the upstream call and the Result type still forces handling the missing case, so async introduces no credential-less-call race. --- .../outbound_credentials/credential_store.py | 11 +++-- .../mcp/outbound_credentials/resolver.py | 38 ++++++++-------- tests/mcp_tests/gateway/test_resolver.py | 44 ++++++++++--------- 3 files changed, 50 insertions(+), 43 deletions(-) diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/credential_store.py b/litellm/proxy/gateway/mcp/outbound_credentials/credential_store.py index 02129833b75..9ab28eff8ad 100644 --- a/litellm/proxy/gateway/mcp/outbound_credentials/credential_store.py +++ b/litellm/proxy/gateway/mcp/outbound_credentials/credential_store.py @@ -23,9 +23,14 @@ class CredentialKey(BaseModel): class CredentialStore(Protocol): - """Fetches the per-subject secret for an `api_key` per-user / BYOK server.""" + """Fetches the per-subject secret for an `api_key` per-user / BYOK server. - def get(self, key: CredentialKey) -> str | None: ... + 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. + """ + + async def get(self, key: CredentialKey) -> str | None: ... class InMemoryCredentialStore: @@ -34,5 +39,5 @@ class InMemoryCredentialStore: def __init__(self, seeded: dict[CredentialKey, str] | None = None) -> None: self._values: dict[CredentialKey, str] = dict(seeded or {}) - def get(self, key: CredentialKey) -> str | None: + async def get(self, key: CredentialKey) -> str | None: return 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 29d1ec0c7cc..972d694025a 100644 --- a/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py +++ b/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py @@ -47,33 +47,33 @@ class UpstreamCredentialProvider: def __init__(self, credential_store: CredentialStore) -> None: self._credential_store = credential_store - def resolve( + async def resolve( self, subject: Subject, server: ServerSpec ) -> Result[httpx.Auth, CredError]: match server.config: case AuthorizationCodeConfig() as config: - return self._authorization_code(subject, server, config) + return await self._authorization_code(subject, server, config) case ClientCredentialsConfig() as config: - return self._client_credentials(subject, server, config) + return await self._client_credentials(subject, server, config) case TokenExchangeConfig() as config: - return self._token_exchange(subject, server, config) + return await self._token_exchange(subject, server, config) case ApiKeyConfig() as config: - return self._api_key(subject, server, config) + return await self._api_key(subject, server, config) case PassthroughConfig() as config: - return self._passthrough(subject, server, config) + return await self._passthrough(subject, server, config) case NoneConfig() as config: - return self._none(subject, server, config) + return await self._none(subject, server, config) case AwsSigV4Config() as config: - return self._aws_sigv4(subject, server, config) + return await self._aws_sigv4(subject, server, config) assert_never(server.config) # --- implemented arms ----------------------------------------------------------------- - def _none( + async def _none( self, subject: Subject, server: ServerSpec, config: NoneConfig ) -> Result[httpx.Auth, CredError]: return Ok(NoOpAuth()) - def _api_key( + async def _api_key( self, subject: Subject, server: ServerSpec, config: ApiKeyConfig ) -> Result[httpx.Auth, CredError]: match config.key_source: @@ -84,7 +84,7 @@ class UpstreamCredentialProvider: StaticHeaderAuth(config.header_for(source.value.get_secret_value())) ) case PerUserEnvVar(): - value = self._per_user_value(subject, server) + value = await self._per_user_value(subject, server) if value is None: return Error( CredError.of_precondition_required( @@ -93,7 +93,7 @@ class UpstreamCredentialProvider: ) return Ok(StaticHeaderAuth(config.header_for(value))) case Byok(): - value = self._per_user_value(subject, server) + value = await self._per_user_value(subject, server) if value is None: return Error( CredError.of_unauthorized( @@ -103,8 +103,8 @@ class UpstreamCredentialProvider: return Ok(StaticHeaderAuth(config.header_for(value))) assert_never(config.key_source) - def _per_user_value(self, subject: Subject, server: ServerSpec) -> str | None: - return self._credential_store.get( + async def _per_user_value(self, subject: Subject, server: ServerSpec) -> str | None: + return await self._credential_store.get( CredentialKey( tenant_id=subject.tenant_id, subject_id=subject.subject_id, @@ -112,7 +112,7 @@ class UpstreamCredentialProvider: ) ) - def _passthrough( + async def _passthrough( self, subject: Subject, server: ServerSpec, config: PassthroughConfig ) -> Result[httpx.Auth, CredError]: # The one arm that forwards a caller-supplied token, and only one the client obtained @@ -126,22 +126,22 @@ class UpstreamCredentialProvider: ) # --- arms awaiting their collaborators (typed stubs, fail closed) ---------------------- - def _authorization_code( + async def _authorization_code( self, subject: Subject, server: ServerSpec, config: AuthorizationCodeConfig ) -> Result[httpx.Auth, CredError]: return _todo(AuthSpecKind.authorization_code) - def _client_credentials( + async def _client_credentials( self, subject: Subject, server: ServerSpec, config: ClientCredentialsConfig ) -> Result[httpx.Auth, CredError]: return _todo(AuthSpecKind.client_credentials) - def _token_exchange( + async def _token_exchange( self, subject: Subject, server: ServerSpec, config: TokenExchangeConfig ) -> Result[httpx.Auth, CredError]: return _todo(AuthSpecKind.token_exchange) - def _aws_sigv4( + async def _aws_sigv4( self, subject: Subject, server: ServerSpec, config: AwsSigV4Config ) -> Result[httpx.Auth, CredError]: return _todo(AuthSpecKind.aws_sigv4) diff --git a/tests/mcp_tests/gateway/test_resolver.py b/tests/mcp_tests/gateway/test_resolver.py index 50890c8dd12..a3a61077b79 100644 --- a/tests/mcp_tests/gateway/test_resolver.py +++ b/tests/mcp_tests/gateway/test_resolver.py @@ -63,8 +63,8 @@ def test_discriminated_union_picks_the_variant_by_kind(): assert isinstance(spec.config, NoneConfig) -def test_none_attaches_no_credential(): - result = PROVIDER.resolve(SUBJECT, _spec(NoneConfig())) +async def test_none_attaches_no_credential(): + result = await PROVIDER.resolve(SUBJECT, _spec(NoneConfig())) assert isinstance(result, Ok) assert isinstance(result.ok, NoOpAuth) assert "Authorization" not in _applied_headers(result.ok) @@ -80,9 +80,9 @@ def test_none_attaches_no_credential(): ("raw", "k"), ], ) -def test_api_key_emits_the_right_scheme(scheme: str, expected: str): +async def test_api_key_emits_the_right_scheme(scheme: str, expected: str): config = ApiKeyConfig(scheme=scheme, key_source=SharedKey(value="k")) # type: ignore[arg-type] - result = PROVIDER.resolve(SUBJECT, _spec(config)) + result = await PROVIDER.resolve(SUBJECT, _spec(config)) assert isinstance(result, Ok) assert isinstance(result.ok, StaticHeaderAuth) assert _applied_headers(result.ok)["Authorization"] == expected @@ -98,38 +98,40 @@ def test_secret_fields_are_masked_in_serialization(): @pytest.mark.parametrize("source", [Byok(), PerUserEnvVar()]) -def test_api_key_per_user_pulls_the_subject_credential(source: object): +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) - result = provider.resolve(SUBJECT, _spec(ApiKeyConfig(key_source=source))) # type: ignore[arg-type] + 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" -def test_api_key_byok_missing_returns_unauthorized(): +async def test_api_key_byok_missing_returns_unauthorized(): # Missing BYOK credential -> 401 + WWW-Authenticate (the user must provide it). - result = PROVIDER.resolve(SUBJECT, _spec(ApiKeyConfig(key_source=Byok()))) + result = await PROVIDER.resolve(SUBJECT, _spec(ApiKeyConfig(key_source=Byok()))) assert isinstance(result, Error) assert result.error.tag == "unauthorized" -def test_api_key_env_var_missing_returns_precondition_required(): +async def test_api_key_env_var_missing_returns_precondition_required(): # Missing per-user env var -> 412 (a setup precondition), distinct from BYOK's 401. - result = PROVIDER.resolve(SUBJECT, _spec(ApiKeyConfig(key_source=PerUserEnvVar()))) + result = await PROVIDER.resolve( + SUBJECT, _spec(ApiKeyConfig(key_source=PerUserEnvVar())) + ) assert isinstance(result, Error) assert result.error.tag == "precondition_required" -def test_api_key_per_user_isolated_by_subject(): +async def test_api_key_per_user_isolated_by_subject(): # The stored key belongs to (t1,u1,s1); a different subject must not receive it. store = InMemoryCredentialStore( {CredentialKey(tenant_id="t1", subject_id="u1", server_id="s1"): "u1-secret"} ) provider = UpstreamCredentialProvider(store) other = Subject(tenant_id="t1", subject_id="u2") - result = provider.resolve(other, _spec(ApiKeyConfig(key_source=Byok()))) + result = await provider.resolve(other, _spec(ApiKeyConfig(key_source=Byok()))) assert isinstance(result, Error) assert result.error.tag == "unauthorized" @@ -140,26 +142,26 @@ def test_missing_status_maps_byok_401_distinct_from_env_var_412(): assert http_status(CredError.of_precondition_required("env var missing")) == 412 -def test_passthrough_forwards_the_inbound_token(): +async def test_passthrough_forwards_the_inbound_token(): subject = Subject(tenant_id="t1", subject_id="u1", inbound_token="upstream-tok") - result = PROVIDER.resolve(subject, _spec(PassthroughConfig())) + result = await PROVIDER.resolve(subject, _spec(PassthroughConfig())) assert isinstance(result, Ok) assert _applied_headers(result.ok)["Authorization"] == "Bearer upstream-tok" -def test_passthrough_without_a_token_fails_closed(): - result = PROVIDER.resolve(SUBJECT, _spec(PassthroughConfig())) +async def test_passthrough_without_a_token_fails_closed(): + result = await PROVIDER.resolve(SUBJECT, _spec(PassthroughConfig())) assert isinstance(result, Error) assert result.error.tag == "unauthorized" -def test_self_contained_arms_never_read_the_inbound_token(): +async def test_self_contained_arms_never_read_the_inbound_token(): # The #30559 guard: none/api_key must produce the identical credential whether or not a # caller bearer is present, proving they never forward the gateway-bound token. with_token = Subject(tenant_id="t1", subject_id="u1", inbound_token="leak-me") for config in (NoneConfig(), ApiKeyConfig(key_source=SharedKey(value="k"))): - without = PROVIDER.resolve(SUBJECT, _spec(config)) - present = PROVIDER.resolve(with_token, _spec(config)) + without = await PROVIDER.resolve(SUBJECT, _spec(config)) + present = await PROVIDER.resolve(with_token, _spec(config)) assert isinstance(without, Ok) and isinstance(present, Ok) assert _applied_headers(without.ok).get("Authorization") == _applied_headers( present.ok @@ -190,7 +192,7 @@ def test_self_contained_arms_never_read_the_inbound_token(): {"kind": "aws_sigv4", "region": "us-east-1"}, ], ) -def test_unimplemented_arms_fail_closed(config: dict): - result = PROVIDER.resolve(SUBJECT, _spec(config)) +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"