mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
test(spend): reconcile concurrent requests and daily activity
This commit is contained in:
parent
b4cff58be7
commit
7e3d7178b4
3 changed files with 244 additions and 65 deletions
|
|
@ -0,0 +1,109 @@
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from math import isclose
|
||||
from typing import Final
|
||||
|
||||
from e2e_config import provider_edge_base, unique_marker
|
||||
from e2e_http import unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatBody, ChatMessage, ChatResponse, KeyGenerateBody, LiteLLMParamsBody, TeamNewBody
|
||||
from spend_e2e_client import SpendClient
|
||||
|
||||
INPUT_RATE: Final = 0.00004
|
||||
OUTPUT_RATE: Final = 0.00008
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TeamTraffic:
|
||||
team_id: str
|
||||
key: str
|
||||
responses: tuple[ChatResponse, ...]
|
||||
|
||||
@property
|
||||
def prompt_tokens(self) -> int:
|
||||
return sum(response.usage.prompt_tokens or 0 for response in self.responses if response.usage)
|
||||
|
||||
@property
|
||||
def completion_tokens(self) -> int:
|
||||
return sum(response.usage.completion_tokens or 0 for response in self.responses if response.usage)
|
||||
|
||||
@property
|
||||
def spend(self) -> float:
|
||||
return self.prompt_tokens * INPUT_RATE + self.completion_tokens * OUTPUT_RATE
|
||||
|
||||
|
||||
def create_traffic(client: SpendClient, resources: ResourceManager) -> tuple[TeamTraffic, ...]:
|
||||
base: Final = provider_edge_base("openai")
|
||||
model: Final = f"e2e-reconciliation-{unique_marker()}"
|
||||
model_id: Final = client.proxy.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(
|
||||
model="openai/gpt-5.6-luna",
|
||||
api_key="os.environ/OPENAI_API_KEY",
|
||||
api_base=None if base is None else f"{base}/v1",
|
||||
input_cost_per_token=INPUT_RATE,
|
||||
output_cost_per_token=OUTPUT_RATE,
|
||||
),
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_model(model_id))
|
||||
|
||||
def team_traffic() -> TeamTraffic:
|
||||
team: Final = client.proxy.create_team(TeamNewBody(team_alias=f"e2e-spend-{unique_marker()}"))
|
||||
resources.defer(lambda: client.proxy.delete_team(team))
|
||||
key: Final = client.proxy.generate_key(KeyGenerateBody(team_id=team, models=[model]))
|
||||
resources.defer(lambda: client.proxy.delete_key(key))
|
||||
|
||||
prompts: Final = tuple(f"Reply with one word. {index} {unique_marker()}" for index in range(7))
|
||||
|
||||
def call(index: int) -> ChatResponse:
|
||||
response: Final = unwrap(
|
||||
client.proxy.chat(
|
||||
key,
|
||||
ChatBody(
|
||||
model=model,
|
||||
messages=[ChatMessage(role="user", content=prompts[index])],
|
||||
max_completion_tokens=128,
|
||||
),
|
||||
)
|
||||
)
|
||||
assert response.id, "successful response must have an ID"
|
||||
assert response.usage is not None, "successful response must have usage"
|
||||
assert response.usage.prompt_tokens is not None and response.usage.prompt_tokens > 0
|
||||
assert response.usage.completion_tokens is not None and response.usage.completion_tokens > 0
|
||||
assert response.usage.total_tokens == response.usage.prompt_tokens + response.usage.completion_tokens
|
||||
assert not response.usage.cache_creation_input_tokens
|
||||
assert not response.usage.cache_read_input_tokens
|
||||
assert not response.usage.prompt_tokens_details or not response.usage.prompt_tokens_details.cached_tokens
|
||||
return response
|
||||
|
||||
sequential: Final = call(0)
|
||||
with ThreadPoolExecutor(max_workers=6) as pool:
|
||||
concurrent: Final = tuple(pool.map(call, range(1, 7)))
|
||||
return TeamTraffic(team, key, (sequential, *concurrent))
|
||||
|
||||
return tuple(team_traffic() for _ in range(2))
|
||||
|
||||
|
||||
def assert_logs_match(client: SpendClient, traffic: TeamTraffic) -> None:
|
||||
expected_ids: Final = frozenset(response.id for response in traffic.responses)
|
||||
assert len(expected_ids) == len(traffic.responses), "responses must have distinct IDs"
|
||||
rows: Final = client.poll_logs_for_key(
|
||||
traffic.key,
|
||||
min_rows=len(traffic.responses),
|
||||
predicate=lambda values: frozenset(row.request_id for row in values) == expected_ids,
|
||||
)
|
||||
assert frozenset(row.request_id for row in rows) == expected_ids, "stored IDs must equal returned response IDs"
|
||||
assert len(rows) == len(traffic.responses), "expected exactly one scoped spend row per response"
|
||||
by_id: Final = {row.request_id: row for row in rows}
|
||||
for response in traffic.responses:
|
||||
row = by_id[response.id]
|
||||
usage = response.usage
|
||||
assert usage is not None and usage.prompt_tokens is not None and usage.completion_tokens is not None
|
||||
assert row.team_id == traffic.team_id
|
||||
assert row.status == "success"
|
||||
assert row.cache_hit != "True"
|
||||
assert row.prompt_tokens == usage.prompt_tokens
|
||||
assert row.completion_tokens == usage.completion_tokens
|
||||
assert row.total_tokens == usage.total_tokens
|
||||
expected_cost = usage.prompt_tokens * INPUT_RATE + usage.completion_tokens * OUTPUT_RATE
|
||||
assert row.spend is not None and isclose(row.spend, expected_cost, rel_tol=1e-6, abs_tol=1e-9)
|
||||
|
|
@ -17,13 +17,12 @@ fails the test; a pricing or token-count drift does not.
|
|||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from math import isclose
|
||||
|
||||
import pytest
|
||||
|
||||
from e2e_http import Result, Success
|
||||
from e2e_http import Success
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatResponse, LiteLLMParamsBody, SpendLogs, SpendLogsParams
|
||||
from models import LiteLLMParamsBody, SpendLogs, SpendLogsParams
|
||||
from spend_e2e_client import SpendClient, SpendLogRow, is_ok, unique_marker, unwrap
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
|
@ -280,51 +279,18 @@ def test_key_spend_equals_sum_of_logs(client: SpendClient, scoped_key: str) -> N
|
|||
), f"key aggregate {key_spend} != sum of logs {logs_total}; rows: {_summarize(rows)}"
|
||||
|
||||
|
||||
@pytest.mark.replayable
|
||||
@pytest.mark.covers("quota_management.spend_tracking.concurrent_burst.loses_no_spend")
|
||||
def test_burst_of_concurrent_calls_loses_no_spend(
|
||||
client: SpendClient, scoped_key: str
|
||||
client: SpendClient, resources: ResourceManager
|
||||
) -> None:
|
||||
"""Six concurrent calls on one key: every call lands its own spend row under a
|
||||
distinct request_id and the key aggregate equals the sum of the rows.
|
||||
Sequential accuracy is covered by test_key_spend_equals_sum_of_logs; this pins
|
||||
the concurrent increment path (parallel writers racing on one key's counter),
|
||||
where a lost update can never be reproduced by sequential calls."""
|
||||
burst = 6
|
||||
from spend_reconciliation import assert_logs_match, create_traffic
|
||||
|
||||
def call(idx: int) -> Result[ChatResponse]:
|
||||
return client.chat(
|
||||
scoped_key,
|
||||
"gemini-2.5-flash",
|
||||
f"burst call {idx} {unique_marker()}",
|
||||
max_tokens=16,
|
||||
)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=burst) as pool:
|
||||
results = tuple(pool.map(call, range(burst)))
|
||||
failed = [r for r in results if not is_ok(r)]
|
||||
assert not failed, f"{len(failed)}/{burst} burst calls failed; first: {failed[0]}"
|
||||
|
||||
rows = client.poll_logs_for_key(
|
||||
scoped_key,
|
||||
min_rows=burst,
|
||||
predicate=lambda rs: len([r for r in rs if (r.spend or 0) > 0]) >= burst,
|
||||
)
|
||||
costed = [r for r in rows if (r.spend or 0) > 0]
|
||||
assert len(costed) >= burst, (
|
||||
f"only {len(costed)}/{burst} burst calls produced a costed row - "
|
||||
f"rows lost under concurrency: {_summarize(rows)}"
|
||||
)
|
||||
request_ids = [r.request_id for r in costed]
|
||||
assert len(set(request_ids)) == len(request_ids), (
|
||||
f"concurrent rows collapsed onto shared request_ids: {_summarize(rows)}"
|
||||
)
|
||||
|
||||
logs_total = sum((r.spend or 0) for r in rows)
|
||||
key_spend = client.poll_key_spend(scoped_key, minimum=logs_total * 0.999)
|
||||
assert _approx_equal(key_spend, logs_total), (
|
||||
f"key aggregate {key_spend} != sum of {len(rows)} rows {logs_total} - "
|
||||
f"spend increments lost under concurrency: {_summarize(rows)}"
|
||||
)
|
||||
traffic = create_traffic(client, resources)
|
||||
for team in traffic:
|
||||
assert_logs_match(client, team)
|
||||
key_spend = client.poll_key_spend(team.key, minimum=team.spend * 0.999999)
|
||||
assert isclose(key_spend, team.spend, rel_tol=1e-6, abs_tol=1e-9)
|
||||
|
||||
|
||||
@pytest.mark.covers("quota_management.spend_tracking.pagination.keeps_total")
|
||||
|
|
|
|||
|
|
@ -7,13 +7,17 @@ missing start/end dates are rejected.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from math import isclose
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from e2e_http import ProbeResult
|
||||
from models import DateRangeParams
|
||||
from lifecycle import ResourceManager
|
||||
from pydantic import BaseModel
|
||||
from spend_e2e_client import SpendClient
|
||||
from spend_reconciliation import assert_logs_match, create_traffic
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
|
@ -24,22 +28,45 @@ class TeamDailyActivityParams(BaseModel):
|
|||
start_date: str | None = None
|
||||
end_date: str | None = None
|
||||
page: int = 1
|
||||
page_size: int = 1
|
||||
team_ids: str | None = None
|
||||
|
||||
|
||||
class TeamDailyActivityRow(BaseModel):
|
||||
date: str
|
||||
metrics: TeamDailyActivityMetrics
|
||||
breakdown: TeamDailyActivityBreakdown
|
||||
|
||||
|
||||
class TeamDailyActivityMetrics(BaseModel):
|
||||
spend: float
|
||||
total_tokens: int
|
||||
prompt_tokens: int
|
||||
completion_tokens: int
|
||||
api_requests: int
|
||||
successful_requests: int
|
||||
failed_requests: int
|
||||
|
||||
|
||||
class TeamDailyActivityEntity(BaseModel):
|
||||
metrics: TeamDailyActivityMetrics
|
||||
|
||||
|
||||
class TeamDailyActivityBreakdown(BaseModel):
|
||||
entities: dict[str, TeamDailyActivityEntity]
|
||||
|
||||
|
||||
class TeamDailyActivityMetadata(BaseModel):
|
||||
page: int
|
||||
total_pages: int
|
||||
has_more: bool
|
||||
total_spend: float
|
||||
total_prompt_tokens: int
|
||||
total_completion_tokens: int
|
||||
total_tokens: int
|
||||
total_api_requests: int
|
||||
total_successful_requests: int
|
||||
total_failed_requests: int
|
||||
|
||||
|
||||
class TeamDailyActivityResponse(BaseModel):
|
||||
|
|
@ -47,32 +74,109 @@ class TeamDailyActivityResponse(BaseModel):
|
|||
metadata: TeamDailyActivityMetadata
|
||||
|
||||
|
||||
def _range_days(days: int) -> DateRangeParams:
|
||||
end = datetime.now(timezone.utc).date()
|
||||
start = end - timedelta(days=days)
|
||||
return DateRangeParams(start_date=start.isoformat(), end_date=end.isoformat())
|
||||
|
||||
|
||||
def _probe(client: SpendClient, params: BaseModel) -> ProbeResult:
|
||||
return client.proxy.transport.probe(ROUTE, params=params)
|
||||
|
||||
|
||||
class TestTeamDailyActivity:
|
||||
@pytest.mark.replayable
|
||||
@pytest.mark.covers("mgmt.team.daily_activity.happy_path")
|
||||
@pytest.mark.parametrize("days", [1, 7, 30])
|
||||
def test_valid_date_range_returns_results_and_metadata(self, client: SpendClient, days: int) -> None:
|
||||
result = _probe(client, _range_days(days))
|
||||
assert result.status_code == 200, (
|
||||
f"{ROUTE} range={days}d must be 200, got {result.status_code}: {result.body[:600]}"
|
||||
def test_valid_date_range_returns_results_and_metadata(
|
||||
self, client: SpendClient, resources: ResourceManager
|
||||
) -> None:
|
||||
started: Final = datetime.now(timezone.utc).date()
|
||||
traffic: Final = create_traffic(client, resources)
|
||||
for team in traffic:
|
||||
assert_logs_match(client, team)
|
||||
ended: Final = datetime.now(timezone.utc).date()
|
||||
team_ids: Final = ",".join(team.team_id for team in traffic)
|
||||
|
||||
def fetch(
|
||||
page: int, start: str = started.isoformat(), end: str = ended.isoformat()
|
||||
) -> TeamDailyActivityResponse:
|
||||
result: Final = _probe(
|
||||
client,
|
||||
TeamDailyActivityParams(
|
||||
start_date=start,
|
||||
end_date=end,
|
||||
page=page,
|
||||
page_size=1,
|
||||
team_ids=team_ids,
|
||||
),
|
||||
)
|
||||
assert result.status_code == 200, f"daily activity failed: {result.status_code} {result.body[:300]}"
|
||||
return TeamDailyActivityResponse.model_validate_json(result.body)
|
||||
|
||||
def pages() -> tuple[TeamDailyActivityResponse, ...]:
|
||||
first: Final = fetch(1)
|
||||
assert first.metadata.total_pages <= len(traffic) * 2, "unexpected extra scoped daily groups"
|
||||
return (first, *(fetch(page) for page in range(2, first.metadata.total_pages + 1)))
|
||||
|
||||
deadline: Final = time.monotonic() + client.proxy.poll_timeout
|
||||
while True:
|
||||
observed = pages()
|
||||
if sum(page.metadata.total_api_requests for page in observed) >= sum(len(t.responses) for t in traffic):
|
||||
break
|
||||
if time.monotonic() >= deadline:
|
||||
break
|
||||
time.sleep(client.proxy.poll_interval)
|
||||
|
||||
assert len(observed) >= 2, "two teams must exercise a page boundary"
|
||||
for index, page in enumerate(observed, 1):
|
||||
assert page.metadata.page == index
|
||||
assert page.metadata.total_pages == len(observed)
|
||||
assert page.metadata.has_more == (index < len(observed))
|
||||
assert len(page.results) == 1, "each fetched daily group must appear in results"
|
||||
row = page.results[0]
|
||||
assert started <= datetime.fromisoformat(row.date).date() <= ended
|
||||
assert len(row.breakdown.entities) == 1
|
||||
assert row.metrics.total_tokens == page.metadata.total_tokens
|
||||
assert row.metrics.prompt_tokens == page.metadata.total_prompt_tokens
|
||||
assert row.metrics.completion_tokens == page.metadata.total_completion_tokens
|
||||
assert row.metrics.api_requests == page.metadata.total_api_requests
|
||||
assert row.metrics.successful_requests == page.metadata.total_successful_requests
|
||||
assert row.metrics.failed_requests == page.metadata.total_failed_requests
|
||||
assert isclose(row.metrics.spend, page.metadata.total_spend, rel_tol=1e-6, abs_tol=1e-9)
|
||||
|
||||
entities: Final = tuple(
|
||||
(team_id, entity.metrics)
|
||||
for page in observed
|
||||
for row in page.results
|
||||
for team_id, entity in row.breakdown.entities.items()
|
||||
)
|
||||
parsed = TeamDailyActivityResponse.model_validate_json(result.body)
|
||||
assert parsed.metadata.page == 1
|
||||
assert parsed.metadata.total_pages >= 1
|
||||
if parsed.results:
|
||||
first = parsed.results[0]
|
||||
assert first.date
|
||||
assert first.metrics.spend >= 0
|
||||
assert first.metrics.total_tokens >= 0
|
||||
assert frozenset(team_id for team_id, _ in entities) == frozenset(team.team_id for team in traffic)
|
||||
for team in traffic:
|
||||
metrics = tuple(metrics for team_id, metrics in entities if team_id == team.team_id)
|
||||
assert sum(m.api_requests for m in metrics) == len(team.responses)
|
||||
assert sum(m.successful_requests for m in metrics) == len(team.responses)
|
||||
assert sum(m.failed_requests for m in metrics) == 0
|
||||
assert sum(m.prompt_tokens for m in metrics) == team.prompt_tokens
|
||||
assert sum(m.completion_tokens for m in metrics) == team.completion_tokens
|
||||
assert sum(m.total_tokens for m in metrics) == team.prompt_tokens + team.completion_tokens
|
||||
assert isclose(sum(m.spend for m in metrics), team.spend, rel_tol=1e-6, abs_tol=1e-9)
|
||||
assert isclose(
|
||||
sum(page.metadata.total_spend for page in observed),
|
||||
sum(team.spend for team in traffic),
|
||||
rel_tol=1e-6,
|
||||
abs_tol=1e-9,
|
||||
)
|
||||
assert sum(page.metadata.total_tokens for page in observed) == sum(
|
||||
team.prompt_tokens + team.completion_tokens for team in traffic
|
||||
)
|
||||
|
||||
empty_date: Final = (started - timedelta(days=7)).isoformat()
|
||||
empty: Final = fetch(1, empty_date, empty_date)
|
||||
assert empty.results == []
|
||||
assert empty.metadata.total_pages == 0
|
||||
assert empty.metadata.page == 1
|
||||
assert not empty.metadata.has_more
|
||||
assert empty.metadata.total_spend == 0
|
||||
assert empty.metadata.total_tokens == 0
|
||||
assert empty.metadata.total_api_requests == 0
|
||||
assert empty.metadata.total_prompt_tokens == 0
|
||||
assert empty.metadata.total_completion_tokens == 0
|
||||
assert empty.metadata.total_successful_requests == 0
|
||||
assert empty.metadata.total_failed_requests == 0
|
||||
|
||||
@pytest.mark.covers("mgmt.team.daily_activity.missing_start_date_rejected")
|
||||
def test_missing_start_date_is_rejected(self, client: SpendClient) -> None:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue