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.
This commit is contained in:
Tin Chi Lo 2026-06-17 17:16:52 -07:00
parent f73e9d61bc
commit 5cde452223
3 changed files with 50 additions and 43 deletions

View file

@ -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)

View file

@ -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)

View file

@ -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"