mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
test(e2e): prove the virtual key lifecycle on every gateway replica (#40023)
* test(e2e): prove the virtual key lifecycle on every replica Walks one virtual key through create, read, partial update, clear, enforce and delete against a live proxy and database, reading every write back on every gateway replica. The management suite already had single write-then-read tests for keys, but none of them proved that a partial /key/update leaves the untouched fields alone, that an explicit null clears a field, or that a write is visible on more than the one gateway that took it. Adds read_back_everywhere to the shared ProxyClient: it polls a GET path on every URL in PROXY_REPLICA_URLS until each replica's parsed body satisfies the caller's predicate, and fails naming the replica that never converged. The CLEAR sentinel in the e2e models makes an explicit JSON null expressible in a body the transport otherwise strips of None fields. Documents /key/update's merge patch semantics on the endpoint docstring. * test(e2e): prove key revocation and field preservation on every replica Applies the findings from an adversarial review of the first commit. The delete step only checked that chat was refused on the gateway that took the write, so it would have passed while a sibling gateway kept serving the deleted key. It now serves one call from every replica first, so each has the key cached and the delete has something to revoke everywhere, then polls every replica for the refusal. The file also carried its own poll loop that tested the deadline before attempting, so it gave up one attempt early and skipped the attempt landing exactly on the deadline. It now shares the harness helper, which is generic over the polled value rather than over a parsed body, so the same loop covers both the info read-back and the chat refusal. The model the enforcement step registers now carries a unique marker in its alias, matching every other deployment this suite creates, so concurrent runs never share one model group. The docstring sentence claimed an explicit null clears any field. It does not: the metadata-backed fields merge into stored metadata, where a null is a silent no-op, and only the key's own columns clear. Regenerating the dashboard types picks up the corrected text. * fix(e2e): delete a deployment that never becomes servable Registering a model posts /model/new and then waits for every replica to list it. When that wait timed out the deployment already existed in the database but its id had never been returned, so no caller could delete it and the row outlived the run. It is now deleted before the failure propagates. Found by review on the key lifecycle suite, whose module fixture registers a deployment this way, but every caller of the shared helper had the same exposure. * docs(e2e): drop the duplicated notes from the lifecycle docstrings The delete method restated what the warm-up helper already explains, and the module restated the merge patch rule that the endpoint and the request model both document.
This commit is contained in:
parent
5f2b4d27d7
commit
4b3355bdc6
7 changed files with 592 additions and 21 deletions
|
|
@ -2840,6 +2840,11 @@ async def update_key_fn(
|
|||
"""
|
||||
Update an existing API key's parameters.
|
||||
|
||||
The body is a merge patch: a field left out keeps its stored value, and on the key's own columns
|
||||
an explicit null clears it. The metadata-backed fields below are the exception, merging into the
|
||||
stored metadata instead: passing one as null leaves it unchanged, while `metadata` itself
|
||||
replaces the stored metadata wholesale.
|
||||
|
||||
Parameters:
|
||||
- key: Optional[str] - The key to update. Either key or key_alias must be provided.
|
||||
- key_alias: Optional[str] - User-friendly key alias. If key is omitted, also identifies the key to update (must match exactly one key, same as /key/delete's key_aliases)
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@
|
|||
- {id: mgmt.key.update.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "key_management_endpoints.py:2462", rationale: "Budget/model changes persist"}
|
||||
- {id: mgmt.key.update.admin_only, module: mgmt, tier: P0, surface: api, assertions: [admin_only], source: "key_management_endpoints.py:2462", rationale: "Non-admin cannot escalate perms"}
|
||||
- {id: mgmt.key.update.happy_path, module: mgmt, tier: P1, surface: ui, assertions: [happy_path], source: "key_management_endpoints.py:2462", rationale: "Key edit through the dashboard"}
|
||||
- {id: mgmt.key.update.preserves_unrelated_fields, module: mgmt, tier: P0, surface: api, assertions: [preserves_unrelated_fields], source: "key_management_endpoints.py:2829", rationale: "A partial /key/update changes only the field it names; alias, models, limits, budget window, team and metadata read back unchanged on every gateway replica"}
|
||||
- {id: mgmt.key.update.clear_persists, module: mgmt, tier: P0, surface: api, assertions: [clear_persists], source: "key_management_endpoints.py:2829", rationale: "An explicit null on /key/update clears max_budget and budget_duration, and the derived budget_reset_at with it, on every gateway replica"}
|
||||
- {id: mgmt.key.delete.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "key_management_endpoints.py:3122", rationale: "Deletion revokes future calls"}
|
||||
- {id: mgmt.key.delete.admin_only, module: mgmt, tier: P0, surface: api, assertions: [admin_only], source: "key_management_endpoints.py:3122", rationale: "Non-owner cannot delete"}
|
||||
- {id: mgmt.key.info.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "key_management_endpoints.py:3380", rationale: "Info reflects all writes"}
|
||||
|
|
|
|||
299
tests/e2e/management/test_key_lifecycle_e2e.py
Normal file
299
tests/e2e/management/test_key_lifecycle_e2e.py
Normal file
|
|
@ -0,0 +1,299 @@
|
|||
"""Live e2e: one virtual key walked through its whole lifecycle, read back on every
|
||||
gateway replica.
|
||||
|
||||
Create, read, partial update, clear, enforce, delete: one method per step, and every
|
||||
step creates its own team and key (both deleted on teardown) so a step reruns or skips
|
||||
on its own. Writes go through the control plane; read-backs poll every URL in
|
||||
PROXY_REPLICA_URLS until each replica converges, because a write that is visible on the
|
||||
gateway that took it and stale on its neighbour is exactly the failure this file exists
|
||||
to catch. Revocation is the slowest of those: a deleted key stays usable on the other
|
||||
replicas until their auth cache entry expires, so the delete step polls each of them
|
||||
rather than asserting once.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Callable, Iterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import Result, StreamingResponse, Success, UnknownApiError, unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from management_client import MODEL_ACCESS_DENIED_MARKER, ManagementClient
|
||||
from models import (
|
||||
CLEAR,
|
||||
ChatBody,
|
||||
ChatMessage,
|
||||
KeyGenerateBody,
|
||||
KeyGenerateResponse,
|
||||
KeyInfo,
|
||||
KeyInfoParams,
|
||||
KeyInfoResponse,
|
||||
KeyMetadata,
|
||||
KeyUpdateBody,
|
||||
LiteLLMParamsBody,
|
||||
TeamNewBody,
|
||||
)
|
||||
from proxy_client import Converged, NotConverged, Poller, await_converged, await_converged_everywhere
|
||||
from transport import Transport
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
BACKING_MODEL: Final = "gpt-4o-mini"
|
||||
DENIED_MODEL: Final = "gpt-5.5"
|
||||
MAX_BUDGET: Final = 25.0
|
||||
TPM_LIMIT: Final = 313131
|
||||
RPM_LIMIT: Final = 323232
|
||||
UPDATED_RPM_LIMIT: Final = 424242
|
||||
BUDGET_DURATION: Final = "30d"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CreatedKey:
|
||||
written: KeyGenerateBody
|
||||
response: KeyGenerateResponse
|
||||
|
||||
@property
|
||||
def key(self) -> str:
|
||||
return self.response.key
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def mock_deployment(client: ManagementClient) -> Iterator[str]:
|
||||
"""A deployment that answers from a canned response, so the enforcement step needs no
|
||||
provider key. The alias carries a unique marker, like every other model this suite
|
||||
registers, so concurrent runs never share one model group."""
|
||||
model_name: Final = f"e2e-key-lifecycle-{unique_marker()}"
|
||||
model_id: Final = client.proxy.create_model(model_name, LiteLLMParamsBody(model=BACKING_MODEL, mock_response="ok"))
|
||||
try:
|
||||
yield model_name
|
||||
finally:
|
||||
client.proxy.delete_model(model_id)
|
||||
|
||||
|
||||
def _await[T](client: ManagementClient, poller: Poller[T], converged: Callable[[T], bool], failure: str) -> T:
|
||||
outcome: Final = await_converged(
|
||||
poller,
|
||||
converged=converged,
|
||||
timeout=client.proxy.poll_timeout,
|
||||
interval=client.proxy.poll_interval,
|
||||
now=time.monotonic,
|
||||
sleep=time.sleep,
|
||||
)
|
||||
match outcome:
|
||||
case Converged(result=result):
|
||||
return result
|
||||
case NotConverged(last_result=last):
|
||||
pytest.fail(f"{failure}; last outcome: {last}")
|
||||
|
||||
|
||||
def _chat_poller(transport: Transport, key: str, model: str) -> Poller[StreamingResponse]:
|
||||
return lambda: transport.send(
|
||||
"/chat/completions",
|
||||
headers=transport.bearer(key),
|
||||
json=ChatBody(
|
||||
model=model,
|
||||
messages=[ChatMessage(role="user", content=f"say hi {unique_marker()}")],
|
||||
max_tokens=16,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _create_key(client: ManagementClient, resources: ResourceManager, model: str) -> CreatedKey:
|
||||
marker: Final = unique_marker()
|
||||
team_id: Final = client.create_team(TeamNewBody(team_alias=f"e2e-key-lifecycle-team-{marker}"))
|
||||
resources.defer(lambda: client.delete_team(team_id))
|
||||
written: Final = KeyGenerateBody(
|
||||
key_alias=f"e2e-key-lifecycle-{marker}",
|
||||
models=[model],
|
||||
max_budget=MAX_BUDGET,
|
||||
tpm_limit=TPM_LIMIT,
|
||||
rpm_limit=RPM_LIMIT,
|
||||
budget_duration=BUDGET_DURATION,
|
||||
metadata=KeyMetadata(tag=marker),
|
||||
team_id=team_id,
|
||||
)
|
||||
response: Final = unwrap(client.generate_key(written))
|
||||
resources.defer(lambda: client.proxy.delete_key(response.key))
|
||||
return CreatedKey(written=written, response=response)
|
||||
|
||||
|
||||
def _key_info_everywhere(
|
||||
client: ManagementClient, key: str, settled: Callable[[KeyInfo], bool]
|
||||
) -> Mapping[str, KeyInfo]:
|
||||
def converged(result: Result[KeyInfoResponse]) -> bool:
|
||||
return isinstance(result, Success) and settled(result.data.info)
|
||||
|
||||
reads: Final = client.proxy.read_back_everywhere(
|
||||
"/key/info", params=KeyInfoParams(key=key), response_type=KeyInfoResponse, converged=converged
|
||||
)
|
||||
return MappingProxyType({replica: unwrap(read).info for replica, read in reads.items()})
|
||||
|
||||
|
||||
def _is_key_not_found(result: Result[KeyInfoResponse]) -> bool:
|
||||
return isinstance(result, UnknownApiError) and result.status_code == 404
|
||||
|
||||
|
||||
def _assert_reads_back(info: KeyInfo, expected: KeyGenerateBody, replica: str) -> None:
|
||||
for field, observed, wanted in (
|
||||
("key_alias", info.key_alias, expected.key_alias),
|
||||
("models", info.models, expected.models),
|
||||
("max_budget", info.max_budget, expected.max_budget),
|
||||
("tpm_limit", info.tpm_limit, expected.tpm_limit),
|
||||
("rpm_limit", info.rpm_limit, expected.rpm_limit),
|
||||
("budget_duration", info.budget_duration, expected.budget_duration),
|
||||
("team_id", info.team_id, expected.team_id),
|
||||
("metadata", info.metadata, expected.metadata),
|
||||
):
|
||||
assert observed == wanted, f"{replica}: /key/info reports {field}={observed!r}, expected {wanted!r}"
|
||||
|
||||
|
||||
def _poll_chat_ok(client: ManagementClient, key: str, model: str) -> None:
|
||||
_ = _await(
|
||||
client,
|
||||
_chat_poller(client.proxy.transport, key, model),
|
||||
lambda outcome: outcome.ok,
|
||||
f"chat on {model} never succeeded for the key before the deadline",
|
||||
)
|
||||
|
||||
|
||||
def _warm_every_replica(client: ManagementClient, key: str, model: str) -> None:
|
||||
"""Serve one call from every replica, so each has the key in its auth cache. Without
|
||||
this the revocation check below would only prove a replica rejects a key it never
|
||||
knew, which is true of any random string."""
|
||||
for replica, transport in client.proxy.replicas.items():
|
||||
_ = _await(
|
||||
client,
|
||||
_chat_poller(transport, key, model),
|
||||
lambda outcome: outcome.ok,
|
||||
f"{replica}: chat on {model} never succeeded for the key before the deadline",
|
||||
)
|
||||
|
||||
|
||||
def _assert_chat_rejected_everywhere(client: ManagementClient, key: str, model: str) -> None:
|
||||
outcomes: Final = await_converged_everywhere(
|
||||
{replica: _chat_poller(transport, key, model) for replica, transport in client.proxy.replicas.items()},
|
||||
converged=lambda outcome: outcome.status_code == 401,
|
||||
timeout=client.proxy.poll_timeout,
|
||||
interval=client.proxy.poll_interval,
|
||||
now=time.monotonic,
|
||||
sleep=time.sleep,
|
||||
)
|
||||
for replica, outcome in outcomes.items():
|
||||
assert isinstance(outcome, Converged), (
|
||||
f"{replica}: the deleted key was still accepted on chat after "
|
||||
f"{client.proxy.poll_timeout}s, last status {outcome.last_result.status_code}"
|
||||
)
|
||||
|
||||
|
||||
class TestKeyLifecycle:
|
||||
def test_create_echoes_every_field_written(
|
||||
self, client: ManagementClient, resources: ResourceManager, mock_deployment: str
|
||||
) -> None:
|
||||
created: Final = _create_key(client, resources, mock_deployment)
|
||||
|
||||
response: Final = created.response
|
||||
for field, observed, wanted in (
|
||||
("key_alias", response.key_alias, created.written.key_alias),
|
||||
("models", response.models, created.written.models),
|
||||
("max_budget", response.max_budget, created.written.max_budget),
|
||||
("tpm_limit", response.tpm_limit, created.written.tpm_limit),
|
||||
("rpm_limit", response.rpm_limit, created.written.rpm_limit),
|
||||
("budget_duration", response.budget_duration, created.written.budget_duration),
|
||||
("team_id", response.team_id, created.written.team_id),
|
||||
("metadata", response.metadata, created.written.metadata),
|
||||
):
|
||||
assert observed == wanted, f"/key/generate echoed {field}={observed!r}, sent {wanted!r}"
|
||||
|
||||
def test_read_reflects_the_create_on_every_replica(
|
||||
self, client: ManagementClient, resources: ResourceManager, mock_deployment: str
|
||||
) -> None:
|
||||
created: Final = _create_key(client, resources, mock_deployment)
|
||||
|
||||
infos: Final = _key_info_everywhere(
|
||||
client, created.key, lambda info: info.key_alias == created.written.key_alias
|
||||
)
|
||||
for replica, info in infos.items():
|
||||
_assert_reads_back(info, created.written, replica)
|
||||
assert info.budget_reset_at is not None, (
|
||||
f"{replica}: /key/info reports no budget_reset_at for budget_duration={BUDGET_DURATION!r}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("mgmt.key.update.preserves_unrelated_fields")
|
||||
def test_partial_update_changes_only_the_named_field(
|
||||
self, client: ManagementClient, resources: ResourceManager, mock_deployment: str
|
||||
) -> None:
|
||||
created: Final = _create_key(client, resources, mock_deployment)
|
||||
before: Final = _key_info_everywhere(client, created.key, lambda info: info.rpm_limit == RPM_LIMIT)
|
||||
|
||||
_ = unwrap(client.update_key(KeyUpdateBody(key=created.key, rpm_limit=UPDATED_RPM_LIMIT)))
|
||||
|
||||
after: Final = _key_info_everywhere(client, created.key, lambda info: info.rpm_limit == UPDATED_RPM_LIMIT)
|
||||
for replica, info in after.items():
|
||||
_assert_reads_back(info, created.written.model_copy(update={"rpm_limit": UPDATED_RPM_LIMIT}), replica)
|
||||
assert info.budget_reset_at == before[replica].budget_reset_at, (
|
||||
f"{replica}: budget_reset_at moved from {before[replica].budget_reset_at!r} to "
|
||||
f"{info.budget_reset_at!r} on a /key/update that did not name budget_duration"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("mgmt.key.update.clear_persists")
|
||||
def test_explicit_null_clears_the_budget_and_its_reset_time(
|
||||
self, client: ManagementClient, resources: ResourceManager, mock_deployment: str
|
||||
) -> None:
|
||||
created: Final = _create_key(client, resources, mock_deployment)
|
||||
_ = _key_info_everywhere(client, created.key, lambda info: info.max_budget == MAX_BUDGET)
|
||||
|
||||
_ = unwrap(client.update_key(KeyUpdateBody(key=created.key, max_budget=CLEAR, budget_duration=CLEAR)))
|
||||
|
||||
cleared: Final = _key_info_everywhere(client, created.key, lambda info: info.max_budget is None)
|
||||
for replica, info in cleared.items():
|
||||
assert info.budget_duration is None, (
|
||||
f"{replica}: budget_duration={info.budget_duration!r} survived an explicit null"
|
||||
)
|
||||
assert info.budget_reset_at is None, (
|
||||
f"{replica}: clearing budget_duration left budget_reset_at={info.budget_reset_at!r}"
|
||||
)
|
||||
_assert_reads_back(
|
||||
info, created.written.model_copy(update={"max_budget": None, "budget_duration": None}), replica
|
||||
)
|
||||
|
||||
def test_key_serves_its_model_and_is_denied_others(
|
||||
self, client: ManagementClient, resources: ResourceManager, mock_deployment: str
|
||||
) -> None:
|
||||
created: Final = _create_key(client, resources, mock_deployment)
|
||||
|
||||
_poll_chat_ok(client, created.key, mock_deployment)
|
||||
|
||||
denied: Final = client.chat_status(created.key, DENIED_MODEL, f"say hi {unique_marker()}")
|
||||
assert denied.status_code == 403, (
|
||||
f"chat on {DENIED_MODEL!r} outside the key's model list must be denied 403, got "
|
||||
f"{denied.status_code}: {denied.body[:300]}"
|
||||
)
|
||||
assert MODEL_ACCESS_DENIED_MARKER in denied.body, (
|
||||
f"403 body must be a model-access denial, got: {denied.body[:300]}"
|
||||
)
|
||||
|
||||
def test_delete_revokes_info_and_chat_on_every_replica(
|
||||
self, client: ManagementClient, resources: ResourceManager, mock_deployment: str
|
||||
) -> None:
|
||||
"""The teardown's deferred delete fires again on the already-deleted key by
|
||||
design: the deferred cleanup must survive this test failing before the
|
||||
in-body delete, and a repeat /key/delete is a cheap no-op the warn-only
|
||||
teardown absorbs."""
|
||||
created: Final = _create_key(client, resources, mock_deployment)
|
||||
_warm_every_replica(client, created.key, mock_deployment)
|
||||
|
||||
client.delete_key_strict(created.key)
|
||||
|
||||
_ = client.proxy.read_back_everywhere(
|
||||
"/key/info",
|
||||
params=KeyInfoParams(key=created.key),
|
||||
response_type=KeyInfoResponse,
|
||||
converged=_is_key_not_found,
|
||||
)
|
||||
_assert_chat_rejected_everywhere(client, created.key, mock_deployment)
|
||||
|
|
@ -8,9 +8,9 @@ from __future__ import annotations
|
|||
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime
|
||||
from typing import Literal
|
||||
from typing import Final, Literal
|
||||
|
||||
from pydantic import AliasChoices, BaseModel, ConfigDict, Field, RootModel, model_validator
|
||||
from pydantic import AliasChoices, BaseModel, ConfigDict, Field, RootModel, model_serializer, model_validator
|
||||
|
||||
# ---------- keys ----------
|
||||
|
||||
|
|
@ -49,6 +49,7 @@ class KeyMetadata(BaseModel):
|
|||
logging: list[KeyLoggingCallback] | None = None
|
||||
priority: str | None = None
|
||||
batch_enqueued_token_limit: int | None = None
|
||||
tag: str | None = None
|
||||
|
||||
|
||||
class ObjectPermission(BaseModel):
|
||||
|
|
@ -81,6 +82,14 @@ class KeyGenerateBody(BaseModel):
|
|||
|
||||
class KeyGenerateResponse(BaseModel):
|
||||
key: str
|
||||
key_alias: str | None = None
|
||||
models: list[str] = []
|
||||
max_budget: float | None = None
|
||||
tpm_limit: int | None = None
|
||||
rpm_limit: int | None = None
|
||||
budget_duration: str | None = None
|
||||
team_id: str | None = None
|
||||
metadata: KeyMetadata | None = None
|
||||
|
||||
|
||||
class KeyRegenerateBody(BaseModel):
|
||||
|
|
@ -122,6 +131,7 @@ class KeyInfo(BaseModel):
|
|||
blocked: bool | None = None
|
||||
spend: float | None = None
|
||||
max_budget: float | None = None
|
||||
budget_duration: str | None = None
|
||||
budget_reset_at: str | None = None
|
||||
budget_id: str | None = None
|
||||
litellm_budget_table: LiteLLMBudgetTable | None = None
|
||||
|
|
@ -919,12 +929,34 @@ class CredentialCreateResponse(BaseModel):
|
|||
# ---------- key / team / user / organization management ----------
|
||||
|
||||
|
||||
class Cleared(BaseModel):
|
||||
"""An explicit JSON null in a merge-patch body. The transport drops `None` fields
|
||||
before sending (`exclude_none`), so `None` means "leave the stored value alone"; a
|
||||
field set to `CLEAR` reaches the wire as `null`, which tells the proxy to clear it."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
@model_serializer
|
||||
def _as_null(self) -> None:
|
||||
return None
|
||||
|
||||
|
||||
CLEAR: Final = Cleared()
|
||||
|
||||
|
||||
class KeyUpdateBody(BaseModel):
|
||||
"""POST /key/update is a merge patch: a field left `None` is dropped from the body and
|
||||
keeps its stored value, `CLEAR` sends an explicit null that clears it (`budget_duration`
|
||||
clears `budget_reset_at` with it), and `metadata` replaces the stored metadata wholesale."""
|
||||
|
||||
key: str
|
||||
models: list[str] | None = None
|
||||
key_alias: str | None = None
|
||||
tpm_limit: int | None = None
|
||||
rpm_limit: int | None = None
|
||||
max_budget: float | Cleared | None = None
|
||||
budget_duration: str | Cleared | None = None
|
||||
metadata: KeyMetadata | None = None
|
||||
|
||||
|
||||
class KeyBlockBody(BaseModel):
|
||||
|
|
|
|||
|
|
@ -16,6 +16,8 @@ from datetime import datetime
|
|||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from e2e_http import (
|
||||
AnthropicHeaders,
|
||||
AuthHeaders,
|
||||
|
|
@ -149,9 +151,7 @@ def await_servable(
|
|||
last_result: Result[ModelsListResponse] | None = None
|
||||
while True:
|
||||
t = now()
|
||||
phase_deadline = (
|
||||
started + timeout if first_seen_at is None else first_seen_at + db_sync_seconds
|
||||
)
|
||||
phase_deadline = started + timeout if first_seen_at is None else first_seen_at + db_sync_seconds
|
||||
remaining = phase_deadline - t
|
||||
if remaining <= 0:
|
||||
if (
|
||||
|
|
@ -164,9 +164,7 @@ def await_servable(
|
|||
|
||||
poll_timeout = min(request_timeout, remaining)
|
||||
last_result = list_models(poll_timeout)
|
||||
listed = isinstance(last_result, Success) and any(
|
||||
entry.id == model_name for entry in last_result.data.data
|
||||
)
|
||||
listed = isinstance(last_result, Success) and any(entry.id == model_name for entry in last_result.data.data)
|
||||
t = now()
|
||||
if not listed:
|
||||
first_seen_at = None
|
||||
|
|
@ -179,9 +177,7 @@ def await_servable(
|
|||
elif t - first_seen_at >= db_sync_seconds:
|
||||
return Servable()
|
||||
|
||||
phase_deadline = (
|
||||
started + timeout if first_seen_at is None else first_seen_at + db_sync_seconds
|
||||
)
|
||||
phase_deadline = started + timeout if first_seen_at is None else first_seen_at + db_sync_seconds
|
||||
wait = min(interval, phase_deadline - now())
|
||||
if wait > 0:
|
||||
sleep(wait)
|
||||
|
|
@ -239,6 +235,88 @@ def servable_timeout_message(
|
|||
)
|
||||
|
||||
|
||||
type Poller[T] = Callable[[], T]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Converged[T]:
|
||||
result: T
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class NotConverged[T]:
|
||||
"""The deadline passed without a read satisfying the predicate; `last_result` is
|
||||
the final read, so the caller can tell a stale body from a failed request."""
|
||||
|
||||
last_result: T
|
||||
|
||||
|
||||
type ConvergeOutcome[T] = Converged[T] | NotConverged[T]
|
||||
|
||||
|
||||
def await_converged[T](
|
||||
poll: Poller[T],
|
||||
*,
|
||||
converged: Callable[[T], bool],
|
||||
timeout: float,
|
||||
interval: float,
|
||||
now: Callable[[], float],
|
||||
sleep: Callable[[float], None],
|
||||
) -> ConvergeOutcome[T]:
|
||||
"""Poll until a read satisfies `converged` or `timeout` elapses.
|
||||
|
||||
Polls before testing the deadline, so a zero or already-spent budget still gets one
|
||||
attempt, and sleeps only min(interval, time left), so the attempt that lands exactly
|
||||
on the deadline is taken rather than skipped. Clock and sleep are injected."""
|
||||
deadline: Final = now() + timeout
|
||||
while True:
|
||||
result = poll()
|
||||
if converged(result):
|
||||
return Converged(result=result)
|
||||
remaining = deadline - now()
|
||||
if remaining <= 0:
|
||||
return NotConverged(last_result=result)
|
||||
sleep(min(interval, remaining))
|
||||
|
||||
|
||||
def await_converged_everywhere[T](
|
||||
pollers: Mapping[str, Poller[T]],
|
||||
*,
|
||||
converged: Callable[[T], bool],
|
||||
timeout: float,
|
||||
interval: float,
|
||||
now: Callable[[], float],
|
||||
sleep: Callable[[float], None],
|
||||
) -> Mapping[str, ConvergeOutcome[T]]:
|
||||
"""`await_converged` against every replica in turn, each with the full budget, so a
|
||||
replica that lags behind the one a write landed on is polled until it catches up
|
||||
rather than failing on its first stale read."""
|
||||
return MappingProxyType(
|
||||
{
|
||||
replica: await_converged(
|
||||
poll, converged=converged, timeout=timeout, interval=interval, now=now, sleep=sleep
|
||||
)
|
||||
for replica, poll in pollers.items()
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def first_lagging_replica[T](
|
||||
outcomes: Mapping[str, ConvergeOutcome[T]],
|
||||
) -> tuple[str, NotConverged[T]] | None:
|
||||
return next(
|
||||
((replica, outcome) for replica, outcome in outcomes.items() if isinstance(outcome, NotConverged)),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def converge_timeout_message(*, what: str, replica: str, timeout: float, last_result: object) -> str:
|
||||
return (
|
||||
f"{what} on {replica} never converged within {timeout}s of the write "
|
||||
f"(control/data-plane propagation issue); last read: {last_result}"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProxyClient:
|
||||
transport: Transport
|
||||
|
|
@ -290,6 +368,52 @@ class ProxyClient:
|
|||
)
|
||||
).info
|
||||
|
||||
def read_back_everywhere[R: BaseModel](
|
||||
self,
|
||||
path: str,
|
||||
*,
|
||||
params: BaseModel,
|
||||
response_type: type[R],
|
||||
converged: Callable[[Result[R]], bool],
|
||||
) -> Mapping[str, Result[R]]:
|
||||
"""GET `path` under the master key on every replica in PROXY_REPLICA_URLS (the
|
||||
data-plane URL alone when the stack exports no per-gateway addresses), polling
|
||||
each to poll_timeout until its read satisfies `converged`. Returns that read per
|
||||
replica, or fails naming the first replica that never converged and its last
|
||||
read. Behind a load balancer the single address proves one replica converged,
|
||||
not all of them; only per-gateway addresses make this a fleet-wide proof."""
|
||||
outcomes: Final = await_converged_everywhere(
|
||||
{
|
||||
url: self._body_poller(transport, path, params, response_type)
|
||||
for url, transport in self.replicas.items()
|
||||
},
|
||||
converged=converged,
|
||||
timeout=self.poll_timeout,
|
||||
interval=self.poll_interval,
|
||||
now=time.monotonic,
|
||||
sleep=time.sleep,
|
||||
)
|
||||
lagging: Final = first_lagging_replica(outcomes)
|
||||
if lagging is not None:
|
||||
replica, outcome = lagging
|
||||
raise AssertionError(
|
||||
converge_timeout_message(
|
||||
what=f"GET {path}",
|
||||
replica=replica,
|
||||
timeout=self.poll_timeout,
|
||||
last_result=outcome.last_result,
|
||||
)
|
||||
)
|
||||
return MappingProxyType(
|
||||
{replica: outcome.result for replica, outcome in outcomes.items() if isinstance(outcome, Converged)}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _body_poller[R: BaseModel](
|
||||
transport: Transport, path: str, params: BaseModel, response_type: type[R]
|
||||
) -> Poller[Result[R]]:
|
||||
return lambda: transport.get(path, headers=transport.master, params=params, response_type=response_type)
|
||||
|
||||
def model_info(self) -> list[ModelInfoEntry]:
|
||||
"""Every configured deployment with the price the proxy resolved for it
|
||||
(config override merged over cost-map defaults)."""
|
||||
|
|
@ -320,9 +444,7 @@ class ProxyClient:
|
|||
response_type=FileListResponse,
|
||||
)
|
||||
|
||||
def list_fine_tuning_jobs(
|
||||
self, key: str, params: FineTuningJobsParams
|
||||
) -> Result[FineTuningJobsResponse]:
|
||||
def list_fine_tuning_jobs(self, key: str, params: FineTuningJobsParams) -> Result[FineTuningJobsResponse]:
|
||||
return self.transport.get(
|
||||
"/v1/fine_tuning/jobs",
|
||||
headers=self.transport.bearer(key),
|
||||
|
|
@ -375,7 +497,11 @@ class ProxyClient:
|
|||
)
|
||||
).model_id
|
||||
written_at = time.monotonic()
|
||||
self._await_model_servable(body.model_name, listed_for)
|
||||
try:
|
||||
self._await_model_servable(body.model_name, listed_for)
|
||||
except BaseException:
|
||||
self.delete_model(model_id)
|
||||
raise
|
||||
settle_propagation(written_at)
|
||||
return model_id
|
||||
|
||||
|
|
|
|||
|
|
@ -1,9 +1,11 @@
|
|||
"""Harness coverage for the model barrier that gates on every replica.
|
||||
"""Harness coverage for the barriers that gate on every replica.
|
||||
|
||||
No proxy needed and no ``e2e`` marker: this pins that a model registered through
|
||||
the control plane only counts as servable once every configured replica lists it
|
||||
on /v1/models, which is what keeps a two-gateway stack from handing a test a
|
||||
model that one gateway has not reloaded yet. The fakes are plain pollers and an
|
||||
on /v1/models, and that a management write only counts as read back once every
|
||||
replica's read satisfies the caller's predicate, which is what keeps a two-gateway
|
||||
stack from handing a test a model or a key that one gateway has not caught up on
|
||||
yet. The fakes are plain pollers standing in for each replica's transport plus an
|
||||
injected clock, so nothing here monkeypatches anything.
|
||||
"""
|
||||
|
||||
|
|
@ -12,18 +14,33 @@ from __future__ import annotations
|
|||
from collections.abc import Iterable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from itertools import chain, repeat
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from e2e_config import parse_replica_urls
|
||||
from e2e_http import Success
|
||||
from models import ModelListEntry, ModelsListResponse
|
||||
from proxy_client import ModelsPoller, NotServableOn, Servable, await_servable_everywhere
|
||||
from e2e_http import Result, Success
|
||||
from models import KeyInfo, KeyInfoResponse, ModelListEntry, ModelsListResponse
|
||||
from proxy_client import (
|
||||
Poller,
|
||||
ConvergeOutcome,
|
||||
Converged,
|
||||
ModelsPoller,
|
||||
NotConverged,
|
||||
NotServableOn,
|
||||
Servable,
|
||||
await_converged_everywhere,
|
||||
await_servable_everywhere,
|
||||
first_lagging_replica,
|
||||
converge_timeout_message,
|
||||
)
|
||||
|
||||
MODEL: Final = "gpt-under-test"
|
||||
TIMEOUT: Final = 10.0
|
||||
INTERVAL: Final = 2.0
|
||||
RPM_BEFORE_UPDATE: Final = 100
|
||||
RPM_AFTER_UPDATE: Final = 200
|
||||
|
||||
|
||||
@dataclass
|
||||
|
|
@ -78,6 +95,91 @@ class TestAwaitServableEverywhere:
|
|||
assert _await(pollers) == Servable()
|
||||
|
||||
|
||||
def _key_info(rpm_limit: int) -> Success[KeyInfoResponse]:
|
||||
return Success(status_code=200, data=KeyInfoResponse(info=KeyInfo(rpm_limit=rpm_limit)))
|
||||
|
||||
|
||||
def _reads(results: Iterable[Result[KeyInfoResponse]]) -> Poller[Result[KeyInfoResponse]]:
|
||||
it: Final = iter(results)
|
||||
return lambda: next(it)
|
||||
|
||||
|
||||
def _updated(result: Result[KeyInfoResponse]) -> bool:
|
||||
return isinstance(result, Success) and result.data.info.rpm_limit == RPM_AFTER_UPDATE
|
||||
|
||||
|
||||
def _converge(
|
||||
pollers: Mapping[str, Poller[Result[KeyInfoResponse]]], clock: FakeClock
|
||||
) -> Mapping[str, ConvergeOutcome[Result[KeyInfoResponse]]]:
|
||||
return await_converged_everywhere(
|
||||
pollers,
|
||||
converged=_updated,
|
||||
timeout=TIMEOUT,
|
||||
interval=INTERVAL,
|
||||
now=clock.now,
|
||||
sleep=clock.sleep,
|
||||
)
|
||||
|
||||
|
||||
class TestAwaitConvergedEverywhere:
|
||||
def test_waits_for_the_replica_that_lags_behind_the_write(self) -> None:
|
||||
clock: Final = FakeClock()
|
||||
pollers: Final = MappingProxyType(
|
||||
{
|
||||
"gateway-1": _reads(repeat(_key_info(RPM_AFTER_UPDATE))),
|
||||
"gateway-2": _reads(
|
||||
chain(repeat(_key_info(RPM_BEFORE_UPDATE), 2), repeat(_key_info(RPM_AFTER_UPDATE)))
|
||||
),
|
||||
}
|
||||
)
|
||||
outcomes: Final = _converge(pollers, clock)
|
||||
assert outcomes == {
|
||||
"gateway-1": Converged(result=_key_info(RPM_AFTER_UPDATE)),
|
||||
"gateway-2": Converged(result=_key_info(RPM_AFTER_UPDATE)),
|
||||
}
|
||||
assert first_lagging_replica(outcomes) is None
|
||||
assert clock.elapsed == 2 * INTERVAL
|
||||
|
||||
def test_names_the_replica_that_never_converges_with_its_last_read(self) -> None:
|
||||
clock: Final = FakeClock()
|
||||
pollers: Final = MappingProxyType(
|
||||
{
|
||||
"gateway-1": _reads(repeat(_key_info(RPM_AFTER_UPDATE))),
|
||||
"gateway-2": _reads(repeat(_key_info(RPM_BEFORE_UPDATE))),
|
||||
}
|
||||
)
|
||||
outcomes: Final = _converge(pollers, clock)
|
||||
assert first_lagging_replica(outcomes) == (
|
||||
"gateway-2",
|
||||
NotConverged(last_result=_key_info(RPM_BEFORE_UPDATE)),
|
||||
)
|
||||
assert clock.elapsed == TIMEOUT
|
||||
message: Final = converge_timeout_message(
|
||||
what="GET /key/info",
|
||||
replica="gateway-2",
|
||||
timeout=TIMEOUT,
|
||||
last_result=_key_info(RPM_BEFORE_UPDATE),
|
||||
)
|
||||
assert "gateway-2" in message and "/key/info" in message and str(RPM_BEFORE_UPDATE) in message
|
||||
|
||||
def test_each_replica_gets_its_own_full_budget(self) -> None:
|
||||
"""A replica that converges late must not eat into the next replica's budget: both
|
||||
need most of the timeout here, so one shared deadline would starve the second."""
|
||||
clock: Final = FakeClock()
|
||||
slow: Final = chain(repeat(_key_info(RPM_BEFORE_UPDATE), 3), repeat(_key_info(RPM_AFTER_UPDATE)))
|
||||
pollers: Final = MappingProxyType(
|
||||
{
|
||||
"gateway-1": _reads(slow),
|
||||
"gateway-2": _reads(
|
||||
chain(repeat(_key_info(RPM_BEFORE_UPDATE), 3), repeat(_key_info(RPM_AFTER_UPDATE)))
|
||||
),
|
||||
}
|
||||
)
|
||||
outcomes: Final = _converge(pollers, clock)
|
||||
assert first_lagging_replica(outcomes) is None
|
||||
assert clock.elapsed == 2 * 3 * INTERVAL
|
||||
|
||||
|
||||
class TestParseReplicaUrls:
|
||||
def test_splits_and_trims_the_gateway_addresses(self) -> None:
|
||||
raw: Final = " http://127.0.0.1:4010/, http://127.0.0.1:4011 "
|
||||
|
|
|
|||
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -8017,6 +8017,11 @@ export interface paths {
|
|||
* Update Key Fn
|
||||
* @description Update an existing API key's parameters.
|
||||
*
|
||||
* The body is a merge patch: a field left out keeps its stored value, and on the key's own columns
|
||||
* an explicit null clears it. The metadata-backed fields below are the exception, merging into the
|
||||
* stored metadata instead: passing one as null leaves it unchanged, while `metadata` itself
|
||||
* replaces the stored metadata wholesale.
|
||||
*
|
||||
* Parameters:
|
||||
* - key: Optional[str] - The key to update. Either key or key_alias must be provided.
|
||||
* - key_alias: Optional[str] - User-friendly key alias. If key is omitted, also identifies the key to update (must match exactly one key, same as /key/delete's key_aliases)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue