mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
fix(e2e): assert the batch cost join on an in-run failed batch instead of a cross-run baton
The batch list is served from LiteLLM_ManagedObjectTable whenever the managed
files hook is loaded, and the Buildkite e2e stacks bundle a fresh Postgres per
build, so a prior run's marker batch is never listed and the baton could only
ever pass vacuously. Each run now creates a batch OpenAI fails at validation
within seconds, retrieves it by its raw provider id with the same key until it
is failed, and asserts the {provider_batch_id}_batch_cost row that retrieve
writes joins the key's token hash and alias
This commit is contained in:
parent
07a0ca1074
commit
654afb5477
2 changed files with 45 additions and 81 deletions
|
|
@ -61,4 +61,4 @@
|
|||
- {id: quota_management.spend_tracking.key_attribution.joins_key, module: quota_management, tier: P1, behavior: spend_tracking, variant: key_attribution, assertions: [joins_key], exercised_on: [chat_completions, messages, responses, embeddings, batches, files, google_native, rust_control_plane], source: "proxy/spend_tracking/spend_tracking_utils.py", rationale: "Every spend row a virtual key writes across chat, queued chat, messages, responses, embeddings, the Gemini passthrough, file upload, batch create, and a replayed callback log carries api_key equal to the key's token hash and the key alias, the join the usage APIs depend on; a re-hashed token shows up as an unattributed key-hash-* row (#39568, #39572)"}
|
||||
- {id: quota_management.spend_tracking.key_attribution.reports_alias_and_email, module: quota_management, tier: P1, behavior: spend_tracking, variant: key_attribution, assertions: [reports_alias_and_email], exercised_on: [chat_completions, messages, responses, embeddings, batches, files, google_native, rust_control_plane], source: "proxy/management_endpoints/internal_user_endpoints.py", rationale: "/spend/logs?api_key= returns every one of the key's rows with its alias and /user/daily/activity aggregates them under the key's token with key_alias and user_email; /spend/logs carries no email field, so the email is asserted on daily activity only"}
|
||||
- {id: quota_management.spend_tracking.key_attribution.health_rows_keep_service_account, module: quota_management, tier: P1, behavior: spend_tracking, variant: key_attribution, assertions: [health_rows_keep_service_account], exercised_on: [chat_completions], source: "proxy/health_check.py", rationale: "A /health probe's spend row stays keyed by the literal litellm-internal-health-check service account rather than a hash of it, so health spend never appears as an unattributed key"}
|
||||
- {id: quota_management.spend_tracking.key_attribution.batch_cost_joins_key, module: quota_management, tier: P1, behavior: spend_tracking, variant: key_attribution, assertions: [batch_cost_joins_key], exercised_on: [batches], source: "proxy/batches_endpoints/endpoints.py", rationale: "Retrieving a completed batch by its raw provider id prices it inline against the retrieving key, so the {provider_batch_id}_batch_cost row must carry that key's token hash and alias; the target is the newest completed cross-run marker batch this proxy has not billed yet, since completion lags by up to 24h, and a cold start with no completed marker is a documented vacuous pass. The CheckBatchCost poller's own unified-id row needs the batch create row in the same database, which a stack booted fresh per run never holds for a completed marker"}
|
||||
- {id: quota_management.spend_tracking.key_attribution.batch_cost_joins_key, module: quota_management, tier: P1, behavior: spend_tracking, variant: key_attribution, assertions: [batch_cost_joins_key], exercised_on: [batches], source: "proxy/batches_endpoints/endpoints.py", rationale: "The retrieve that first sees a batch in a terminal state writes its {provider_batch_id}_batch_cost row, so the batch each run creates is one OpenAI fails at validation within seconds and the test retrieves it by its raw provider id with the same key until it is failed; that raw-id retrieve is never owned by the CheckBatchCost poller, so it prices the batch inline against the retrieving key and the row must carry that key's token hash and alias. A completed batch with a positive cost is out of one run's reach: OpenAI's completion window is 24h and a stack booted fresh per run lists no earlier run's batches"}
|
||||
|
|
|
|||
|
|
@ -12,19 +12,17 @@ re-hashed token (v1.99.0's regression, #39568 and #39572) shows up as a
|
|||
key-hash-* row with no alias and no email in the customer's usage exports.
|
||||
|
||||
The health-check service account writes rows too; those must stay keyed by the
|
||||
literal service-account name, never by a hash of it. A batch's cost row lands
|
||||
only once the batch completes, up to 24h later, so the batch cost path rides a
|
||||
cross-run baton like the batches suite: each run submits a one-line marker batch
|
||||
and never cancels it, and the newest completed marker from any run that this
|
||||
proxy has not billed yet is retrieved by its raw provider id with this run's
|
||||
key. That retrieve is the writer under test: a raw id is never poller-owned, so
|
||||
the proxy prices it inline against the retrieving key, and the
|
||||
{provider_batch_id}_batch_cost row must join this run's token with its alias.
|
||||
The CheckBatchCost poller's own row (the unified id, billed against the
|
||||
submitting key) needs the batch's create row in the same database, which a stack
|
||||
booted fresh per run never holds for a completed marker, so that writer is out
|
||||
of this test's reach. A cold start with no completed marker is a documented
|
||||
vacuous pass, never a skip.
|
||||
literal service-account name, never by a hash of it. A batch's cost row is
|
||||
written by the retrieve that first sees the batch in a terminal state, so the
|
||||
batch the run creates is one OpenAI fails at validation within seconds (its one
|
||||
line targets /v1/embeddings under a /v1/chat/completions batch), and the test
|
||||
retrieves it by its raw provider id with the same key until it is failed. A raw
|
||||
id is never owned by the CheckBatchCost poller, so that retrieve prices the batch
|
||||
inline against the retrieving key and its {provider_batch_id}_batch_cost row
|
||||
must join the key's token with its alias. A completed batch with a positive
|
||||
cost is out of a single run's reach (OpenAI's completion window is 24h, and a
|
||||
stack booted fresh per run lists no earlier run's batches), so the poller's own
|
||||
row is not asserted here.
|
||||
|
||||
/spend/logs carries no email field, so the email assertion lives on
|
||||
/user/daily/activity alone; /spend/logs is held to the alias in metadata.
|
||||
|
|
@ -38,7 +36,7 @@ from datetime import datetime, timedelta, timezone
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from models import ChatMessage, KeyGenerateBody, SpendLogsParams
|
||||
from models import KeyGenerateBody
|
||||
from proxy_client import Converged, await_converged
|
||||
from pydantic import BaseModel
|
||||
from spend_e2e_client import (
|
||||
|
|
@ -64,11 +62,9 @@ BATCH_MODEL: Final = "openai-gpt-4o-mini"
|
|||
BATCH_BACKEND_MODEL: Final = "gpt-4o-mini"
|
||||
BATCH_PROVIDER: Final = "openai"
|
||||
HEALTH_SERVICE_ACCOUNT: Final = "litellm-internal-health-check"
|
||||
BATON_MARKER_KEY: Final = "litellm_e2e_suite"
|
||||
BATON_MARKER_VALUE: Final = "key-attribution-baton"
|
||||
BATON_POLL_SECONDS: Final = 30.0
|
||||
BATON_POLL_INTERVAL_SECONDS: Final = 10.0
|
||||
BATON_LIST_LIMIT: Final = 100
|
||||
BATCH_TERMINAL_STATUSES: Final = frozenset({"completed", "failed", "cancelled", "expired"})
|
||||
FAILED_BATCH_POLL_SECONDS: Final = 120.0
|
||||
FAILED_BATCH_POLL_INTERVAL_SECONDS: Final = 5.0
|
||||
MAX_TOKENS: Final = 8
|
||||
REPLAY_RESPONSE_COST: Final = 0.0001
|
||||
REPLAY_PROMPT_TOKENS: Final = 5
|
||||
|
|
@ -86,17 +82,16 @@ WRITE_PATHS: Final = (
|
|||
)
|
||||
|
||||
|
||||
class BatchLineBody(BaseModel):
|
||||
class EmbeddingLineBody(BaseModel):
|
||||
model: str
|
||||
messages: list[ChatMessage]
|
||||
max_tokens: int
|
||||
input: str
|
||||
|
||||
|
||||
class BatchLine(BaseModel):
|
||||
class EmbeddingLine(BaseModel):
|
||||
custom_id: str
|
||||
method: str = "POST"
|
||||
url: str = "/v1/chat/completions"
|
||||
body: BatchLineBody
|
||||
url: str = "/v1/embeddings"
|
||||
body: EmbeddingLineBody
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -134,26 +129,19 @@ def _call_id(name: str, sent: StreamingResponse) -> WritePath:
|
|||
return WritePath(name=name, request_id=sent.call_id)
|
||||
|
||||
|
||||
def _batch_jsonl(marker: str) -> bytes:
|
||||
line: Final = BatchLine(
|
||||
custom_id=marker,
|
||||
body=BatchLineBody(
|
||||
model=BATCH_BACKEND_MODEL,
|
||||
messages=[ChatMessage(role="user", content=f"Reply with the word ok. {marker}")],
|
||||
max_tokens=MAX_TOKENS,
|
||||
),
|
||||
)
|
||||
def _endpoint_mismatched_jsonl(marker: str) -> bytes:
|
||||
line: Final = EmbeddingLine(custom_id=marker, body=EmbeddingLineBody(model=BATCH_BACKEND_MODEL, input=marker))
|
||||
return f"{line.model_dump_json()}\n".encode()
|
||||
|
||||
|
||||
def _drive_batch(client: SpendClient, identity: AttributedKey, marker: str) -> tuple[WritePath, WritePath]:
|
||||
uploaded: Final = client.upload_batch_file(identity.key, BATCH_MODEL, _batch_jsonl(marker))
|
||||
uploaded: Final = client.upload_batch_file(identity.key, BATCH_MODEL, _endpoint_mismatched_jsonl(marker))
|
||||
created: Final = client.create_batch(
|
||||
identity.key,
|
||||
BatchCreateBody(
|
||||
input_file_id=uploaded.id,
|
||||
model=BATCH_MODEL,
|
||||
metadata={BATON_MARKER_KEY: BATON_MARKER_VALUE, "run": marker},
|
||||
metadata={"run": marker},
|
||||
),
|
||||
)
|
||||
return (
|
||||
|
|
@ -204,48 +192,26 @@ def _drive_every_write_path(client: SpendClient, identity: AttributedKey) -> tup
|
|||
)
|
||||
|
||||
|
||||
def _completed_batons(listed: list[BatchObject]) -> tuple[BatchObject, ...]:
|
||||
return tuple(
|
||||
sorted(
|
||||
(
|
||||
batch
|
||||
for batch in listed
|
||||
if batch.status == "completed" and (batch.metadata or {}).get(BATON_MARKER_KEY) == BATON_MARKER_VALUE
|
||||
),
|
||||
key=lambda batch: batch.created_at or 0,
|
||||
reverse=True,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _provider_batch_id(unified_batch_id: str) -> str:
|
||||
encoded: Final = unified_batch_id.removeprefix("batch_")
|
||||
decoded: Final = base64.urlsafe_b64decode(encoded + "=" * (-len(encoded) % 4)).decode()
|
||||
return decoded.removeprefix("litellm:").split(";", 1)[0]
|
||||
|
||||
|
||||
def _batch_cost_request_id(unified_batch_id: str) -> str:
|
||||
return f"{_provider_batch_id(unified_batch_id)}_batch_cost"
|
||||
def _driven_batch_id(driven: DrivenKey) -> str:
|
||||
return next(path.request_id for path in driven.paths if path.name == "batch_create")
|
||||
|
||||
|
||||
def _newest_unbilled_completed_baton(client: SpendClient, key: str) -> BatchObject | None:
|
||||
def _await_terminal_batch(client: SpendClient, key: str, provider_batch_id: str) -> BatchObject:
|
||||
outcome: Final = await_converged(
|
||||
lambda: _completed_batons(client.list_batches(key, BATCH_MODEL, limit=BATON_LIST_LIMIT)),
|
||||
converged=lambda completed: bool(completed),
|
||||
timeout=BATON_POLL_SECONDS,
|
||||
interval=BATON_POLL_INTERVAL_SECONDS,
|
||||
lambda: client.retrieve_batch(key, provider_batch_id, provider=BATCH_PROVIDER),
|
||||
converged=lambda batch: batch.status in BATCH_TERMINAL_STATUSES,
|
||||
timeout=FAILED_BATCH_POLL_SECONDS,
|
||||
interval=FAILED_BATCH_POLL_INTERVAL_SECONDS,
|
||||
now=time.monotonic,
|
||||
sleep=time.sleep,
|
||||
)
|
||||
completed: Final = outcome.result if isinstance(outcome, Converged) else ()
|
||||
return next(
|
||||
(
|
||||
batch
|
||||
for batch in completed
|
||||
if not client.proxy.spend_logs(SpendLogsParams(request_id=_batch_cost_request_id(batch.id)))
|
||||
),
|
||||
None,
|
||||
)
|
||||
return outcome.result if isinstance(outcome, Converged) else outcome.last_result
|
||||
|
||||
|
||||
def _health_rows_between(client: SpendClient, started_at: datetime) -> list[SpendLogRow]:
|
||||
|
|
@ -416,23 +382,21 @@ class TestKeyAttribution:
|
|||
"quota_management.spend_tracking.key_attribution.batch_cost_joins_key",
|
||||
exercised_on=["batches"],
|
||||
)
|
||||
def test_completed_batch_cost_row_joins_the_key(self, client: SpendClient, driven: DrivenKey) -> None:
|
||||
completed: Final = _newest_unbilled_completed_baton(client, driven.identity.key)
|
||||
if completed is None:
|
||||
return
|
||||
provider_batch_id: Final = _provider_batch_id(completed.id)
|
||||
fetched: Final = client.retrieve_batch(driven.identity.key, provider_batch_id, provider=BATCH_PROVIDER)
|
||||
assert fetched.status == "completed", f"listed-completed marker retrieved as {fetched.status!r}"
|
||||
cost_request_id: Final = _batch_cost_request_id(completed.id)
|
||||
rows: Final = client.proxy.poll_logs_for_request_id(
|
||||
cost_request_id,
|
||||
predicate=lambda found: any((row.spend or 0) > 0 for row in found),
|
||||
def test_terminal_batch_cost_row_joins_the_retrieving_key(self, client: SpendClient, driven: DrivenKey) -> None:
|
||||
provider_batch_id: Final = _provider_batch_id(_driven_batch_id(driven))
|
||||
fetched: Final = _await_terminal_batch(client, driven.identity.key, provider_batch_id)
|
||||
assert fetched.status == "failed", (
|
||||
f"endpoint-mismatched batch {provider_batch_id} is {fetched.status!r} after "
|
||||
f"{FAILED_BATCH_POLL_SECONDS:.0f}s, so its terminal cost row cannot be asserted"
|
||||
)
|
||||
priced: Final = [row for row in rows if (row.spend or 0) > 0]
|
||||
assert priced, f"retrieving completed batch {provider_batch_id} wrote no positive-cost row under {cost_request_id}"
|
||||
cost_request_id: Final = f"{provider_batch_id}_batch_cost"
|
||||
rows: Final = client.proxy.poll_logs_for_request_id(cost_request_id)
|
||||
assert rows, f"retrieving failed batch {provider_batch_id} wrote no cost row under {cost_request_id}"
|
||||
call_types: Final = tuple(sorted({row.call_type or "" for row in rows}))
|
||||
assert call_types == ("aretrieve_batch",), f"cost rows under {cost_request_id} carry call types {call_types}"
|
||||
unjoined: Final = [
|
||||
(row.call_type, row.api_key, row.metadata.user_api_key_alias if row.metadata else None)
|
||||
for row in priced
|
||||
for row in rows
|
||||
if row.api_key != driven.identity.token
|
||||
or row.metadata is None
|
||||
or row.metadata.user_api_key_alias != driven.identity.alias
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue