fix(usage): address provider throughput review

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
atul naik 2026-10-03 21:55:00 +05:30
parent 45f3cf4d1b
commit 3c6dfb463f
7 changed files with 114 additions and 115 deletions

View file

@ -870,6 +870,7 @@ model LiteLLM_DailyUserSpend {
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
timed_completion_tokens BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@ -939,6 +940,7 @@ model LiteLLM_DailyOrganizationSpend {
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
timed_completion_tokens BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@ -977,6 +979,7 @@ model LiteLLM_DailyEndUserSpend {
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
timed_completion_tokens BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@unique([end_user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@ -1014,6 +1017,7 @@ model LiteLLM_DailyAgentSpend {
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
timed_completion_tokens BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@unique([agent_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@ -1051,6 +1055,7 @@ model LiteLLM_DailyTeamSpend {
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
timed_completion_tokens BigInt @default(0)
ptu_flat_cost Float @default(0.0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@ -1091,6 +1096,7 @@ model LiteLLM_DailyTagSpend {
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
timed_completion_tokens BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt

View file

@ -1,6 +1,8 @@
import asyncio
from collections.abc import Coroutine
from collections.abc import Coroutine, Iterator
from copy import deepcopy
from functools import reduce
from itertools import groupby
from typing import Final
from litellm._logging import verbose_proxy_logger
@ -13,6 +15,47 @@ from litellm.proxy.db.db_transaction_queue.base_update_queue import (
from litellm.types.services import ServiceTypes
def _daily_spend_updates(
updates: list[dict[str, BaseDailySpendTransaction]],
) -> Iterator[tuple[str, BaseDailySpendTransaction]]:
for update in updates:
yield from update.items()
def _merge_daily_spend_transactions(
existing: BaseDailySpendTransaction,
payload: BaseDailySpendTransaction,
) -> BaseDailySpendTransaction:
return {
**existing,
"spend": existing["spend"] + payload["spend"],
"prompt_tokens": existing["prompt_tokens"] + payload["prompt_tokens"],
"completion_tokens": existing["completion_tokens"] + payload["completion_tokens"],
"api_requests": existing["api_requests"] + payload["api_requests"],
"successful_requests": existing["successful_requests"] + payload["successful_requests"],
"failed_requests": existing["failed_requests"] + payload["failed_requests"],
"cache_read_input_tokens": (existing.get("cache_read_input_tokens", 0) or 0)
+ (payload.get("cache_read_input_tokens", 0) or 0),
"cache_creation_input_tokens": (existing.get("cache_creation_input_tokens", 0) or 0)
+ (payload.get("cache_creation_input_tokens", 0) or 0),
"compression_saved_tokens": (existing.get("compression_saved_tokens", 0) or 0)
+ (payload.get("compression_saved_tokens", 0) or 0),
"compression_savings_spend": (existing.get("compression_savings_spend", 0) or 0)
+ (payload.get("compression_savings_spend", 0) or 0),
"prompt_caching_savings_spend": (existing.get("prompt_caching_savings_spend", 0) or 0)
+ (payload.get("prompt_caching_savings_spend", 0) or 0),
"gateway_injected_caching_savings_spend": (existing.get("gateway_injected_caching_savings_spend", 0) or 0)
+ (payload.get("gateway_injected_caching_savings_spend", 0) or 0),
"autorouter_savings_spend": (existing.get("autorouter_savings_spend", 0) or 0)
+ (payload.get("autorouter_savings_spend", 0) or 0),
"total_response_time_ms": (existing.get("total_response_time_ms", 0) or 0)
+ (payload.get("total_response_time_ms", 0) or 0),
"timed_requests": (existing.get("timed_requests", 0) or 0) + (payload.get("timed_requests", 0) or 0),
"timed_completion_tokens": (existing.get("timed_completion_tokens", 0) or 0)
+ (payload.get("timed_completion_tokens", 0) or 0),
}
class DailySpendUpdateQueue(BaseUpdateQueue):
"""
In memory buffer for daily spend updates that should be committed to the database
@ -113,62 +156,14 @@ class DailySpendUpdateQueue(BaseUpdateQueue):
updates: list[dict[str, BaseDailySpendTransaction]],
) -> dict[str, BaseDailySpendTransaction]:
"""Aggregate updates by daily_transaction_key."""
aggregated_daily_spend_update_transactions: Final[dict[str, BaseDailySpendTransaction]] = {}
for _update in updates:
for _key, payload in _update.items():
if _key in aggregated_daily_spend_update_transactions:
daily_transaction = aggregated_daily_spend_update_transactions[_key]
daily_transaction["spend"] += payload["spend"]
daily_transaction["prompt_tokens"] += payload["prompt_tokens"]
daily_transaction["completion_tokens"] += payload["completion_tokens"]
daily_transaction["api_requests"] += payload["api_requests"]
daily_transaction["successful_requests"] += payload["successful_requests"]
daily_transaction["failed_requests"] += payload["failed_requests"]
# Add optional metrics cache_read_input_tokens and cache_creation_input_tokens
daily_transaction["cache_read_input_tokens"] = (
payload.get("cache_read_input_tokens", 0) or 0
) + daily_transaction.get("cache_read_input_tokens", 0)
daily_transaction["cache_creation_input_tokens"] = (
payload.get("cache_creation_input_tokens", 0) or 0
) + daily_transaction.get("cache_creation_input_tokens", 0)
daily_transaction["compression_saved_tokens"] = (
payload.get("compression_saved_tokens", 0) or 0
) + daily_transaction.get("compression_saved_tokens", 0)
daily_transaction["compression_savings_spend"] = (
payload.get("compression_savings_spend", 0) or 0
) + daily_transaction.get("compression_savings_spend", 0)
daily_transaction["prompt_caching_savings_spend"] = (
payload.get("prompt_caching_savings_spend", 0) or 0
) + daily_transaction.get("prompt_caching_savings_spend", 0)
daily_transaction["gateway_injected_caching_savings_spend"] = (
payload.get("gateway_injected_caching_savings_spend", 0) or 0
) + daily_transaction.get("gateway_injected_caching_savings_spend", 0)
daily_transaction["autorouter_savings_spend"] = (
payload.get("autorouter_savings_spend", 0) or 0
) + daily_transaction.get("autorouter_savings_spend", 0)
daily_transaction["total_response_time_ms"] = (
payload.get("total_response_time_ms", 0) or 0
) + daily_transaction.get("total_response_time_ms", 0)
daily_transaction["timed_requests"] = (
payload.get("timed_requests", 0) or 0
) + daily_transaction.get("timed_requests", 0)
daily_transaction["timed_completion_tokens"] = (
payload.get("timed_completion_tokens", 0) or 0
) + daily_transaction.get("timed_completion_tokens", 0)
else:
aggregated_daily_spend_update_transactions[_key] = deepcopy(payload)
return aggregated_daily_spend_update_transactions
ordered_updates: Final = sorted(_daily_spend_updates(updates), key=lambda update: update[0])
return {
key: reduce(
_merge_daily_spend_transactions,
(deepcopy(payload) for _, payload in grouped_updates),
)
for key, grouped_updates in groupby(ordered_updates, key=lambda update: update[0])
}
async def _emit_new_item_added_to_queue_event(
self,

View file

@ -247,7 +247,7 @@ def _provider_throughput(
) -> ProviderThroughputMetrics:
output_tokens_per_second: Final = (
timed_completion_tokens * 1000 / total_response_time_ms
if timed_completion_tokens > 0 and timed_requests > 0 and total_response_time_ms > 0
if timed_requests > 0 and total_response_time_ms > 0
else None
)
return ProviderThroughputMetrics(

View file

@ -107,9 +107,7 @@ async def test_add_multiple_updates(daily_spend_update_queue):
@pytest.mark.asyncio
async def test_aggregated_daily_spend_update_empty(daily_spend_update_queue):
"""Test aggregating updates from an empty queue"""
result = (
await daily_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions()
)
result = await daily_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions()
assert result == {}
@ -129,9 +127,7 @@ async def test_get_aggregated_daily_spend_update_transactions_single_key():
updates = [{test_key: test_transaction}]
# Test aggregation
result = DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions(
updates
)
result = DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions(updates)
assert len(result) == 1
assert test_key in result
@ -164,9 +160,7 @@ async def test_get_aggregated_daily_spend_update_transactions_multiple_keys():
updates = [{test_key1: test_transaction1}, {test_key2: test_transaction2}]
# Test aggregation
result = DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions(
updates
)
result = DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions(updates)
assert len(result) == 2
assert test_key1 in result
@ -221,9 +215,7 @@ async def test_get_aggregated_daily_spend_update_transactions_same_key():
updates = [{test_key: test_transaction1}, {test_key: test_transaction2}]
# Test aggregation
result = DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions(
updates
)
result = DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions(updates)
assert len(result) == 1
assert test_key in result
@ -280,9 +272,7 @@ async def test_flush_and_get_aggregated_daily_spend_update_transactions(
await daily_spend_update_queue.add_update({test_key: test_transaction2})
# Flush and get aggregated transactions
result = (
await daily_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions()
)
result = await daily_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions()
assert len(result) == 1
assert test_key in result
@ -290,9 +280,7 @@ async def test_flush_and_get_aggregated_daily_spend_update_transactions(
@pytest.mark.asyncio
async def test_queue_max_size_triggers_aggregation(
monkeypatch, daily_spend_update_queue
):
async def test_queue_max_size_triggers_aggregation(monkeypatch, daily_spend_update_queue):
"""Test that reaching MAX_SIZE_IN_MEMORY_QUEUE triggers aggregation"""
# Override MAX_SIZE_IN_MEMORY_QUEUE for testing
litellm._turn_on_debug()
@ -316,9 +304,7 @@ async def test_queue_max_size_triggers_aggregation(
assert daily_spend_update_queue.update_queue.qsize() == 1
# Verify the aggregated values
result = (
await daily_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions()
)
result = await daily_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions()
assert result[test_key]["spend"] == 6.0
assert result[test_key]["prompt_tokens"] == 600
assert result[test_key]["completion_tokens"] == 300
@ -435,9 +421,7 @@ async def test_cache_token_fields_aggregation(daily_spend_update_queue):
@pytest.mark.asyncio
async def test_queue_size_reduction_with_large_volume(
monkeypatch, daily_spend_update_queue
):
async def test_queue_size_reduction_with_large_volume(monkeypatch, daily_spend_update_queue):
"""Test that queue size is actually reduced when dealing with many items"""
# Set a smaller MAX_SIZE for testing
monkeypatch.setattr(daily_spend_update_queue, "MAX_SIZE_IN_MEMORY_QUEUE", 10)
@ -478,9 +462,7 @@ async def test_queue_size_reduction_with_large_volume(
assert daily_spend_update_queue.update_queue.qsize() <= 10
# Verify total costs are correct
result = (
await daily_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions()
)
result = await daily_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions()
print("RESULT", json.dumps(result, indent=4))
assert result[user1_key]["spend"] == 200 * 0.5 # 10.0
@ -553,6 +535,7 @@ async def test_every_optional_daily_metric_aggregates(daily_spend_update_queue):
paths, so the driver reads as zero on the dashboard however much it saved.
"""
test_key = "user1_2023-01-01_key123_claude-haiku-4-5_anthropic"
def _numeric(annotation):
# additive metrics may be declared NotRequired[float] for rows queued by a pod
# running the previous release, so unwrap before matching

View file

@ -127,9 +127,7 @@ def test_counters_increment_rather_than_overwrite(column):
def test_request_id_is_preserved_when_a_later_batch_carries_none():
sql, params = build_bulk_upsert(
TAG_TABLE, merge_by_conflict_key(TAG_TABLE, (tag_txn(request_id=None),))
)
sql, params = build_bulk_upsert(TAG_TABLE, merge_by_conflict_key(TAG_TABLE, (tag_txn(request_id=None),)))
assert '"request_id" = COALESCE(EXCLUDED."request_id", "LiteLLM_DailyTagSpend"."request_id")' in sql
assert None in params

View file

@ -103,11 +103,21 @@ async def test_update_database_attributes_router_rejected_failure_to_model_group
)
with (
patch("litellm.proxy.proxy_server.disable_spend_logs", True), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch("litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch("litellm.proxy.proxy_server.llm_router", llm_router), # test-quality-ok: get_llm_router reads this proxy_server module global at call time; no injection seam
patch(
"litellm.proxy.proxy_server.disable_spend_logs", True
), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch(
"litellm.proxy.proxy_server.prisma_client", MagicMock()
), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch(
"litellm.proxy.proxy_server.user_api_key_cache", MagicMock()
), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch(
"litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"
), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch(
"litellm.proxy.proxy_server.llm_router", llm_router
), # test-quality-ok: get_llm_router reads this proxy_server module global at call time; no injection seam
):
await db_writer.update_database(
token="test-token",
@ -320,7 +330,9 @@ async def test_a_routed_request_reaches_the_auto_router_rollup_whether_or_not_sp
}
with (
patch("litellm.proxy.proxy_server.disable_spend_logs", disable_spend_logs), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch(
"litellm.proxy.proxy_server.disable_spend_logs", disable_spend_logs
), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"),
patch(
@ -2868,17 +2880,22 @@ async def test_daily_transaction_carries_compression_saved_tokens():
@pytest.mark.asyncio
@pytest.mark.parametrize("estimate, recorded_savings, expected", [
pytest.param(None, None, -0.005, id="plain-classifier-cost"),
pytest.param({"version": 1, "status": "unknown"}, None, 0.0, id="unknown"),
pytest.param({"version": 2, "status": "unknown"}, None, 0.0, id="unknown-v2"),
pytest.param({"version": 1, "status": "unknown"}, -0.003, 0.0, id="unknown-stale-value"),
pytest.param({"version": 0, "status": "estimated"}, -0.003, 0.0, id="unsupported-version"),
pytest.param({"version": 1, "status": "estimated"}, -0.003, -0.003, id="estimated"),
pytest.param(None, -0.003, -0.003, id="legacy"),
])
@pytest.mark.parametrize(
"estimate, recorded_savings, expected",
[
pytest.param(None, None, -0.005, id="plain-classifier-cost"),
pytest.param({"version": 1, "status": "unknown"}, None, 0.0, id="unknown"),
pytest.param({"version": 2, "status": "unknown"}, None, 0.0, id="unknown-v2"),
pytest.param({"version": 1, "status": "unknown"}, -0.003, 0.0, id="unknown-stale-value"),
pytest.param({"version": 0, "status": "estimated"}, -0.003, 0.0, id="unsupported-version"),
pytest.param({"version": 1, "status": "estimated"}, -0.003, -0.003, id="estimated"),
pytest.param(None, -0.003, -0.003, id="legacy"),
],
)
async def test_daily_transaction_compression_saved_tokens_zero_when_absent(
estimate: dict[str, object] | None, recorded_savings: float | None, expected: float,
estimate: dict[str, object] | None,
recorded_savings: float | None,
expected: float,
) -> None:
"""Requests without any compression metadata produce a zero count."""
writer = DBSpendUpdateWriter()
@ -2897,12 +2914,14 @@ async def test_daily_transaction_compression_saved_tokens_zero_when_absent(
"prompt_tokens": 100,
"completion_tokens": 10,
"spend": 0.01,
"metadata": json.dumps({
"usage_object": {"prompt_tokens": 100, "completion_tokens": 10},
"routing_decision": {"savings_baseline_model": "anthropic/claude-sonnet-5", "classifier_cost": 0.005},
"autorouter_savings": recorded_savings,
"autorouter_savings_estimate": estimate,
}),
"metadata": json.dumps(
{
"usage_object": {"prompt_tokens": 100, "completion_tokens": 10},
"routing_decision": {"savings_baseline_model": "anthropic/claude-sonnet-5", "classifier_cost": 0.005},
"autorouter_savings": recorded_savings,
"autorouter_savings_estimate": estimate,
}
),
}
transaction = await writer._common_add_spend_log_transaction_to_daily_transaction(
@ -3395,9 +3414,7 @@ async def test_failed_per_entity_increment_from_redis_restores_only_what_may_sti
)
mock_redis_update_buffer.restore_transactions_to_redis.assert_awaited_once()
restored = mock_redis_update_buffer.restore_transactions_to_redis.call_args.kwargs[
"db_spend_update_transactions"
]
restored = mock_redis_update_buffer.restore_transactions_to_redis.call_args.kwargs["db_spend_update_transactions"]
assert restored["user_list_transactions"] is None
assert restored["team_list_transactions"] == {"team-1": 1.5}
assert restored["key_list_transactions"] == ({"key-1": 1.5} if safe_to_resend else None)

View file

@ -1858,7 +1858,7 @@ def test_grouping_sets_dispatcher_returns_provider_throughput_for_models_and_mod
assert day.breakdown.model_groups["public-gpt-4o"].provider_breakdown["openai"].output_tokens_per_second == 300.0
def test_grouping_sets_dispatcher_returns_no_throughput_for_legacy_rows_without_timed_tokens():
def test_grouping_sets_dispatcher_returns_zero_for_timed_requests_without_completion_tokens():
from litellm.proxy.management_endpoints.common_daily_activity import (
_GROUP_DATE_MODEL_PROVIDER,
_aggregate_grouping_sets_records_sync,
@ -1877,7 +1877,7 @@ def test_grouping_sets_dispatcher_returns_no_throughput_for_legacy_rows_without_
day = _aggregate_grouping_sets_records_sync(records=records, api_key_metadata={})["results"][0]
assert day.breakdown.models["gpt-4o"].provider_breakdown["openai"].output_tokens_per_second is None
assert day.breakdown.models["gpt-4o"].provider_breakdown["openai"].output_tokens_per_second == 0.0
def test_grouping_sets_dispatcher_keeps_ptu_flat_cost_out_of_the_provider_breakdown():