mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
f73e9d61bc
commit
5cde452223
3 changed files with 50 additions and 43 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue