test(spend): reconcile concurrent requests and daily activity

This commit is contained in:
Yuneng Jiang 2026-09-14 21:52:01 -07:00
parent b4cff58be7
commit 7e3d7178b4
No known key found for this signature in database
3 changed files with 244 additions and 65 deletions

View file

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

View file

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

View file

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