mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(proxy): roll up only closed days into LiteLLM_DailyGlobalSpend and split the key-free read at the marker
The write path no longer dual-writes the global table. The cron rolls up closed UTC days only, so a pod still flushing the current day can never leave the global table short. The key-free arm reads days through the marker from the global table and later days from LiteLLM_DailyUserSpend in one UNION ALL, and the marker comes from the config cache rather than a per-request database lookup. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
c8a2d8c349
commit
ad8de0e192
8 changed files with 204 additions and 456 deletions
|
|
@ -25,41 +25,29 @@ SpendRow = Mapping[str, object]
|
|||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DailySpendTable:
|
||||
"""A daily rollup table and the unique constraint its upserts arbitrate on."""
|
||||
"""The physical table behind one entity's daily rollup."""
|
||||
|
||||
name: str
|
||||
key_columns: tuple[str, ...]
|
||||
entity_id_column: str
|
||||
carries_request_id: bool = False
|
||||
|
||||
|
||||
DAILY_SPEND_TABLES: Final[Mapping[DailySpendEntity, DailySpendTable]] = MappingProxyType(
|
||||
{
|
||||
"user": DailySpendTable(name="LiteLLM_DailyUserSpend", entity_id_column="user_id"),
|
||||
"team": DailySpendTable(name="LiteLLM_DailyTeamSpend", entity_id_column="team_id"),
|
||||
"org": DailySpendTable(name="LiteLLM_DailyOrganizationSpend", entity_id_column="organization_id"),
|
||||
"end_user": DailySpendTable(name="LiteLLM_DailyEndUserSpend", entity_id_column="end_user_id"),
|
||||
"agent": DailySpendTable(name="LiteLLM_DailyAgentSpend", entity_id_column="agent_id"),
|
||||
"tag": DailySpendTable(name="LiteLLM_DailyTagSpend", entity_id_column="tag", carries_request_id=True),
|
||||
}
|
||||
)
|
||||
|
||||
# The unique constraint's columns after the entity id, in constraint order. A NULL can
|
||||
# never match itself in a unique index, so every one of these is normalized to '': the
|
||||
# conflict target has to be NULL-free or the row is re-inserted on every single flush.
|
||||
_KEY_COLUMNS: Final = ("date", "api_key", "model", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint")
|
||||
|
||||
|
||||
def _entity_table(name: str, entity_id_column: str, carries_request_id: bool = False) -> DailySpendTable:
|
||||
return DailySpendTable(
|
||||
name=name, key_columns=(entity_id_column, *_KEY_COLUMNS), carries_request_id=carries_request_id
|
||||
)
|
||||
|
||||
|
||||
DAILY_SPEND_TABLES: Final[Mapping[DailySpendEntity, DailySpendTable]] = MappingProxyType(
|
||||
{
|
||||
"user": _entity_table("LiteLLM_DailyUserSpend", "user_id"),
|
||||
"team": _entity_table("LiteLLM_DailyTeamSpend", "team_id"),
|
||||
"org": _entity_table("LiteLLM_DailyOrganizationSpend", "organization_id"),
|
||||
"end_user": _entity_table("LiteLLM_DailyEndUserSpend", "end_user_id"),
|
||||
"agent": _entity_table("LiteLLM_DailyAgentSpend", "agent_id"),
|
||||
"tag": _entity_table("LiteLLM_DailyTagSpend", "tag", carries_request_id=True),
|
||||
}
|
||||
)
|
||||
|
||||
GLOBAL_SPEND_TABLE: Final = DailySpendTable(
|
||||
name="LiteLLM_DailyGlobalSpend",
|
||||
key_columns=("date", "model", "model_group", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint"),
|
||||
)
|
||||
|
||||
_COUNTER_COLUMNS: Final = (
|
||||
"prompt_tokens",
|
||||
"completion_tokens",
|
||||
|
|
@ -104,7 +92,7 @@ def _as_float(value: object) -> float:
|
|||
|
||||
def conflict_key(table: DailySpendTable, transaction: SpendRow) -> tuple[str, ...]:
|
||||
"""The tuple the database arbitrates the upsert on, normalized free of NULLs."""
|
||||
return tuple(_as_text(transaction.get(column)) for column in table.key_columns)
|
||||
return tuple(_as_text(transaction.get(column)) for column in (table.entity_id_column, *_KEY_COLUMNS))
|
||||
|
||||
|
||||
def _merge(group: Sequence[SpendRow]) -> SpendRow:
|
||||
|
|
@ -142,11 +130,7 @@ def _row_params(
|
|||
return (
|
||||
str(uuid.uuid4()),
|
||||
*key,
|
||||
*(
|
||||
()
|
||||
if "model_group" in table.key_columns
|
||||
else (None if transaction.get("model_group") is None else _as_text(transaction.get("model_group")),)
|
||||
),
|
||||
None if transaction.get("model_group") is None else _as_text(transaction.get("model_group")),
|
||||
*(_as_int(transaction.get(column)) for column in _COUNTER_COLUMNS),
|
||||
*(_as_float(transaction.get(column)) for column in _SPEND_COLUMNS),
|
||||
*((None if request_id is None else _as_text(request_id),) if table.carries_request_id else ()),
|
||||
|
|
@ -156,25 +140,26 @@ def _row_params(
|
|||
def _insert_columns(table: DailySpendTable) -> tuple[str, ...]:
|
||||
return (
|
||||
"id",
|
||||
*table.key_columns,
|
||||
*(() if "model_group" in table.key_columns else ("model_group",)),
|
||||
table.entity_id_column,
|
||||
*_KEY_COLUMNS,
|
||||
"model_group",
|
||||
*_COUNTER_COLUMNS,
|
||||
*_SPEND_COLUMNS,
|
||||
*(("request_id",) if table.carries_request_id else ()),
|
||||
)
|
||||
|
||||
|
||||
def _upsert_statement(
|
||||
def build_bulk_upsert(
|
||||
table: DailySpendTable,
|
||||
batch: Sequence[tuple[tuple[str, ...], SpendRow]],
|
||||
first_param: int,
|
||||
) -> str:
|
||||
) -> tuple[str, tuple[SqlValue, ...]]:
|
||||
"""The single statement writing one merged batch, plus its positional arguments."""
|
||||
columns: Final = _insert_columns(table)
|
||||
quoted_table: Final = f'"{table.name}"'
|
||||
rows: Final = ", ".join(
|
||||
"("
|
||||
+ ", ".join(
|
||||
f"${first_param + row_index * len(columns) + offset}::{_CASTS.get(column, 'text')}"
|
||||
f"${row_index * len(columns) + offset + 1}::{_CASTS.get(column, 'text')}"
|
||||
for offset, column in enumerate(columns)
|
||||
)
|
||||
+ ", (NOW() AT TIME ZONE 'UTC'))"
|
||||
|
|
@ -191,44 +176,11 @@ def _upsert_statement(
|
|||
if table.carries_request_id
|
||||
else ""
|
||||
)
|
||||
return (
|
||||
sql: Final = (
|
||||
f'INSERT INTO {quoted_table} ({_quoted(columns)}, "updated_at")\n'
|
||||
f"VALUES {rows}\n"
|
||||
f"ON CONFLICT ({_quoted(table.key_columns)}) DO UPDATE SET\n"
|
||||
f"ON CONFLICT ({_quoted((table.entity_id_column, *_KEY_COLUMNS))}) DO UPDATE SET\n"
|
||||
f" {increments}{request_id_update},\n"
|
||||
f" \"updated_at\" = (NOW() AT TIME ZONE 'UTC')"
|
||||
)
|
||||
|
||||
|
||||
def _params(table: DailySpendTable, batch: Sequence[tuple[tuple[str, ...], SpendRow]]) -> tuple[SqlValue, ...]:
|
||||
return tuple(value for key, transaction in batch for value in _row_params(table, key, transaction))
|
||||
|
||||
|
||||
def build_bulk_upsert(
|
||||
table: DailySpendTable,
|
||||
batch: Sequence[tuple[tuple[str, ...], SpendRow]],
|
||||
) -> tuple[str, tuple[SqlValue, ...]]:
|
||||
"""The single statement writing one merged batch, plus its positional arguments."""
|
||||
return _upsert_statement(table, batch, first_param=1), _params(table, batch)
|
||||
|
||||
|
||||
def build_bulk_upsert_with_global_rollup(
|
||||
table: DailySpendTable,
|
||||
batch: Sequence[tuple[tuple[str, ...], SpendRow]],
|
||||
) -> tuple[str, tuple[SqlValue, ...]]:
|
||||
"""One statement writing a batch to its table and, atomically, its key-free rollup
|
||||
to ``LiteLLM_DailyGlobalSpend``.
|
||||
|
||||
A data-modifying CTE runs both inserts in the same snapshot and transaction, so a
|
||||
batch that lands in one table lands in both and a retried deadlock replays both.
|
||||
Postgres does not order the CTE against the main statement, so two writers can still
|
||||
deadlock across the tables; the caller's deadlock retry covers that, and each insert
|
||||
takes its own rows in key order so same-table lock order stays deterministic.
|
||||
"""
|
||||
global_batch: Final = merge_by_conflict_key(GLOBAL_SPEND_TABLE, tuple(row for _, row in batch))
|
||||
entity_params: Final = _params(table, batch)
|
||||
sql: Final = (
|
||||
f"WITH entity_rows AS (\n{_upsert_statement(table, batch, first_param=1)}\nRETURNING 1)\n"
|
||||
f"{_upsert_statement(GLOBAL_SPEND_TABLE, global_batch, first_param=len(entity_params) + 1)}"
|
||||
)
|
||||
return sql, (*entity_params, *_params(GLOBAL_SPEND_TABLE, global_batch))
|
||||
return sql, tuple(value for key, transaction in batch for value in _row_params(table, key, transaction))
|
||||
|
|
|
|||
|
|
@ -46,7 +46,6 @@ from litellm.proxy._types import (
|
|||
from litellm.proxy.db.daily_spend_bulk_upsert import (
|
||||
DAILY_SPEND_TABLES,
|
||||
build_bulk_upsert,
|
||||
build_bulk_upsert_with_global_rollup,
|
||||
merge_by_conflict_key,
|
||||
)
|
||||
from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import (
|
||||
|
|
@ -1940,11 +1939,7 @@ class DBSpendUpdateWriter:
|
|||
merged_batch = merge_by_conflict_key(
|
||||
table=table, transactions=tuple(transactions_to_process.values())
|
||||
)
|
||||
sql, params = (
|
||||
build_bulk_upsert_with_global_rollup(table=table, batch=merged_batch)
|
||||
if entity_type == "user"
|
||||
else build_bulk_upsert(table=table, batch=merged_batch)
|
||||
)
|
||||
sql, params = build_bulk_upsert(table=table, batch=merged_batch)
|
||||
await prisma_client.db.execute_raw(sql, *params)
|
||||
except Exception as batch_error:
|
||||
# Log detailed error information for debugging batch upsert failures
|
||||
|
|
|
|||
|
|
@ -11,8 +11,7 @@ from typing_extensions import ReadOnly, TypedDict
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT
|
||||
from litellm.proxy._types import CommonProxyErrors
|
||||
from litellm.proxy.db.daily_spend_bulk_upsert import GLOBAL_SPEND_TABLE
|
||||
from litellm.proxy.spend_tracking.daily_global_spend_rollup import reconciled_through
|
||||
from litellm.proxy.spend_tracking.daily_global_spend_rollup import GLOBAL_SPEND_TABLE_NAME, reconciled_through
|
||||
from litellm.proxy.spend_tracking.key_metadata_recovery import (
|
||||
attach_user_emails,
|
||||
recover_double_hashed_key_metadata,
|
||||
|
|
@ -736,28 +735,62 @@ def _rollup_metric_select(table_name: str) -> str:
|
|||
_MODEL_GROUP_EXPR: Final = "COALESCE(NULLIF(model_group, ''), model)"
|
||||
|
||||
|
||||
async def key_free_source_table(prisma_client: PrismaClient, query: _AggregatedQueryKwargs) -> str | None:
|
||||
"""The table the key-free arm reads from, when the global rollup can answer instead of the per-key table.
|
||||
_KEY_FREE_SOURCE_COLUMNS: Final = (
|
||||
"date",
|
||||
"model",
|
||||
"model_group",
|
||||
"custom_llm_provider",
|
||||
"mcp_namespaced_tool_name",
|
||||
"endpoint",
|
||||
"spend",
|
||||
"prompt_tokens",
|
||||
"completion_tokens",
|
||||
"cache_read_input_tokens",
|
||||
"cache_creation_input_tokens",
|
||||
"compression_saved_tokens",
|
||||
"compression_savings_spend",
|
||||
"prompt_caching_savings_spend",
|
||||
"gateway_injected_caching_savings_spend",
|
||||
"autorouter_savings_spend",
|
||||
"api_requests",
|
||||
"successful_requests",
|
||||
"failed_requests",
|
||||
)
|
||||
|
||||
Only an unfiltered read of the user table has the same rows as ``LiteLLM_DailyGlobalSpend``,
|
||||
and only through the day the reconcile marker has reached: the writer keeps that day
|
||||
current, later days are covered once the next run advances the marker.
|
||||
|
||||
async def global_rollup_reconciled_through(prisma_client: PrismaClient, query: _AggregatedQueryKwargs) -> str | None:
|
||||
"""The last day ``LiteLLM_DailyGlobalSpend`` can answer the key-free arm for, or None to
|
||||
read it all from the per-key table.
|
||||
|
||||
Only an unfiltered read of the user table sums to the same rows as the global table. The
|
||||
marker read is served from the config cache, so this is not a database round trip per request.
|
||||
"""
|
||||
if query["table_name"] != "litellm_dailyuserspend":
|
||||
return None
|
||||
if query["entity_id"] is not None or query["api_key"] is not None or query["exclude_entity_ids"]:
|
||||
return None
|
||||
_, adjusted_end = _adjust_dates_for_timezone(
|
||||
query["start_date"], query["end_date"], query["timezone_offset_minutes"], query["include_current_utc_day"]
|
||||
)
|
||||
try:
|
||||
marker: Final = await reconciled_through(prisma_client)
|
||||
return await reconciled_through(prisma_client)
|
||||
except Exception as exc: # noqa: BLE001 # the per-key table is always a correct answer, so never fail the read
|
||||
verbose_proxy_logger.warning("Could not read the daily global spend marker, using the per-key table: %s", exc)
|
||||
return None
|
||||
if marker is None or adjusted_end > marker:
|
||||
return None
|
||||
return GLOBAL_SPEND_TABLE.name
|
||||
|
||||
|
||||
def _key_free_source(pg_table: str, where_clause: str, marker_param: str | None) -> str:
|
||||
"""The relation the key-free arm aggregates: the per-key table alone, or the global rollup
|
||||
for days through the marker plus the per-key table for the days still open after it."""
|
||||
if marker_param is None:
|
||||
return f'"{pg_table}"\n WHERE {where_clause}'
|
||||
columns: Final = ", ".join(_KEY_FREE_SOURCE_COLUMNS)
|
||||
return f"""(
|
||||
SELECT {columns}
|
||||
FROM "{GLOBAL_SPEND_TABLE_NAME}"
|
||||
WHERE {where_clause} AND date <= {marker_param}
|
||||
UNION ALL
|
||||
SELECT {columns}
|
||||
FROM "{pg_table}"
|
||||
WHERE {where_clause} AND date > {marker_param}
|
||||
) AS key_free_source"""
|
||||
|
||||
|
||||
def _build_aggregated_sql_query(
|
||||
|
|
@ -772,15 +805,16 @@ def _build_aggregated_sql_query(
|
|||
exclude_entity_ids: list[str] | None = None, # mutable-ok: filter union shared with the paginated path
|
||||
timezone_offset_minutes: int | None = None,
|
||||
include_current_utc_day: bool = False,
|
||||
key_free_table: str | None = None,
|
||||
global_rollup_through: str | None = None,
|
||||
) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params
|
||||
"""Build the GROUPING SETS query for aggregated daily activity.
|
||||
|
||||
One statement, two UNION ALL arms over the same WHERE clause. The first arm is
|
||||
key-free: grand total, per-date totals and the (date, model / model_group /
|
||||
provider / mcp / endpoint) rollups, so its row count never grows with the number
|
||||
of keys; it reads ``key_free_table`` when given (the global rollup, whose row count
|
||||
never grew with the number of keys to begin with) and the entity table otherwise.
|
||||
of keys. With ``global_rollup_through`` it reads days through that marker from
|
||||
``LiteLLM_DailyGlobalSpend`` (whose row count never grew with the number of keys to
|
||||
begin with) and only the days after it from the per-key table.
|
||||
The second arm emits the (date, <dimension>, api_key) rollups for the
|
||||
USAGE_TOP_API_KEYS_LIMIT highest-spend keys only. Both arms share the 7-bit
|
||||
group_level bitmask (date, api_key, model, model_group, provider, mcp, endpoint).
|
||||
|
|
@ -806,8 +840,8 @@ def _build_aggregated_sql_query(
|
|||
exclude_entity_ids=exclude_entity_ids,
|
||||
)
|
||||
sentinel_param: Final = f"${len(where_params) + 1}"
|
||||
marker_param: Final = None if global_rollup_through is None else f"${len(where_params) + 2}"
|
||||
metric_select: Final = _rollup_metric_select(table_name)
|
||||
key_free_source: Final = key_free_table or pg_table
|
||||
|
||||
# TODO: drop the successful_requests/failed_requests aggregates (and the
|
||||
# total_successful_requests metadata they feed) once the admin UI reads SGR
|
||||
|
|
@ -826,8 +860,7 @@ def _build_aggregated_sql_query(
|
|||
| GROUPING(model, {_MODEL_GROUP_EXPR},
|
||||
custom_llm_provider, mcp_namespaced_tool_name,
|
||||
endpoint) AS group_level,{metric_select}
|
||||
FROM "{key_free_source}"
|
||||
WHERE {where_clause}
|
||||
FROM {_key_free_source(pg_table, where_clause, marker_param)}
|
||||
GROUP BY GROUPING SETS (
|
||||
(date),
|
||||
(date, model),
|
||||
|
|
@ -869,7 +902,8 @@ def _build_aggregated_sql_query(
|
|||
))
|
||||
"""
|
||||
|
||||
return sql_query, [*where_params, PTU_SENTINEL_API_KEY]
|
||||
marker_params: Final = () if global_rollup_through is None else (global_rollup_through,)
|
||||
return sql_query, [*where_params, PTU_SENTINEL_API_KEY, *marker_params]
|
||||
|
||||
|
||||
def _build_entity_rollup_sql_query(
|
||||
|
|
@ -1418,7 +1452,8 @@ async def get_daily_activity_aggregated(
|
|||
include_current_utc_day=include_current_utc_day,
|
||||
)
|
||||
sql_query, sql_params = _build_aggregated_sql_query(
|
||||
**query_kwargs, key_free_table=await key_free_source_table(prisma_client, query_kwargs)
|
||||
**query_kwargs,
|
||||
global_rollup_through=await global_rollup_reconciled_through(prisma_client, query_kwargs),
|
||||
)
|
||||
entity_query: Final = _build_entity_rollup_sql_query(**query_kwargs) if include_entity_breakdown else None
|
||||
|
||||
|
|
|
|||
|
|
@ -1,9 +1,10 @@
|
|||
"""Reconcile ``LiteLLM_DailyGlobalSpend`` from ``LiteLLM_DailyUserSpend``, one day per transaction.
|
||||
"""Roll closed UTC days of ``LiteLLM_DailyUserSpend`` up into ``LiteLLM_DailyGlobalSpend``.
|
||||
|
||||
The spend writer keeps both tables in step from the moment it is deployed; this job rolls up
|
||||
the days before that and records how far it has reached in ``LiteLLM_Config`` so usage reads
|
||||
know when the global table can answer for a date range. It runs as a background cron, never
|
||||
in a Prisma migration, since on a large deployment the aggregate is minutes of work.
|
||||
Only days that are over get rolled up, so a pod still flushing per-key spend for the current
|
||||
day can never leave the global table short; usage reads serve days through the recorded
|
||||
marker from the global table and later days live from the per-key table. The marker lives in
|
||||
``LiteLLM_Config``. This runs as a background cron, never in a Prisma migration, since on a
|
||||
large deployment the first backfill is minutes of work.
|
||||
"""
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
|
@ -19,7 +20,6 @@ from litellm.constants import (
|
|||
DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS,
|
||||
DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM,
|
||||
)
|
||||
from litellm.proxy.db.daily_spend_bulk_upsert import GLOBAL_SPEND_TABLE
|
||||
from litellm.repositories.config_repository import ConfigRepository
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -27,8 +27,11 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
_DAY_TRANSACTION_TIMEOUT: Final = timedelta(minutes=10)
|
||||
_REPLAY_DAYS: Final = 1
|
||||
GLOBAL_SPEND_TABLE_NAME: Final = "LiteLLM_DailyGlobalSpend"
|
||||
# The unique constraint, in constraint order. NULL never matches itself in a unique index, so
|
||||
# every column is normalized to '' or the same group would be inserted again on every run.
|
||||
_KEY_COLUMNS: Final = ("date", "model", "model_group", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint")
|
||||
_METRIC_COLUMNS: Final = (
|
||||
"prompt_tokens",
|
||||
"completion_tokens",
|
||||
|
|
@ -51,23 +54,21 @@ def _quoted(columns: tuple[str, ...]) -> str:
|
|||
|
||||
|
||||
def _reconcile_day_sql() -> str:
|
||||
key_columns: Final = GLOBAL_SPEND_TABLE.key_columns
|
||||
normalized_keys: Final = ", ".join(f"COALESCE(\"{column}\", '')" for column in key_columns)
|
||||
normalized_keys: Final = ", ".join(f"COALESCE(\"{column}\", '')" for column in _KEY_COLUMNS)
|
||||
sums: Final = ", ".join(f'SUM("{column}")' for column in _METRIC_COLUMNS)
|
||||
overwrite: Final = ", ".join(f'"{column}" = EXCLUDED."{column}"' for column in _METRIC_COLUMNS)
|
||||
return (
|
||||
f'INSERT INTO "{GLOBAL_SPEND_TABLE.name}" ("id", {_quoted(key_columns)}, {_quoted(_METRIC_COLUMNS)}, '
|
||||
f'INSERT INTO "{GLOBAL_SPEND_TABLE_NAME}" ("id", {_quoted(_KEY_COLUMNS)}, {_quoted(_METRIC_COLUMNS)}, '
|
||||
'"updated_at")\n'
|
||||
f"SELECT gen_random_uuid()::text, {normalized_keys}, {sums}, (NOW() AT TIME ZONE 'UTC')\n"
|
||||
'FROM "LiteLLM_DailyUserSpend" WHERE "date" = $1\n'
|
||||
f"GROUP BY {normalized_keys}\n"
|
||||
f"ON CONFLICT ({_quoted(key_columns)}) DO UPDATE SET {overwrite}, "
|
||||
f"ON CONFLICT ({_quoted(_KEY_COLUMNS)}) DO UPDATE SET {overwrite}, "
|
||||
"\"updated_at\" = (NOW() AT TIME ZONE 'UTC')"
|
||||
)
|
||||
|
||||
|
||||
RECONCILE_DAY_SQL: Final = _reconcile_day_sql()
|
||||
_LOCK_GLOBAL_TABLE_SQL: Final = f'LOCK TABLE "{GLOBAL_SPEND_TABLE.name}" IN EXCLUSIVE MODE'
|
||||
_PENDING_DAYS_SQL: Final = (
|
||||
'SELECT DISTINCT "date" FROM "LiteLLM_DailyUserSpend" WHERE "date" >= $1 AND "date" <= $2 ORDER BY "date"'
|
||||
)
|
||||
|
|
@ -134,19 +135,19 @@ def _first_pending_day(marker: str | None) -> str:
|
|||
|
||||
|
||||
async def pending_days(prisma_client: "PrismaClient", today: date) -> tuple[str, ...]:
|
||||
"""Every UTC day through today still to roll up, oldest first; the marker day and the one
|
||||
before it are replayed so rows flushed by a pre-writer pod during a rolling deploy are folded in."""
|
||||
"""Every closed UTC day (strictly before today) still to roll up, oldest first. The marker
|
||||
day and the one before it are replayed so per-key rows that landed after their day was
|
||||
rolled up (a flush straddling midnight, a late retry) are folded in."""
|
||||
marker: Final = await reconciled_through(prisma_client)
|
||||
rows: Final = await prisma_client.db.query_raw(_PENDING_DAYS_SQL, _first_pending_day(marker), today.isoformat())
|
||||
return tuple(sorted({*(_DateRow.model_validate(row).date for row in rows), today.isoformat()}))
|
||||
last_closed_day: Final = (today - timedelta(days=1)).isoformat()
|
||||
rows: Final = await prisma_client.db.query_raw(_PENDING_DAYS_SQL, _first_pending_day(marker), last_closed_day)
|
||||
return tuple(_DateRow.model_validate(row).date for row in rows)
|
||||
|
||||
|
||||
async def reconcile_day(prisma_client: "PrismaClient", day: str) -> None:
|
||||
"""Rewrite one day of the global table from the per-key sums; the table lock keeps the
|
||||
writer's increments out between the aggregate and the overwrite so none are lost."""
|
||||
async with prisma_client.db.tx(timeout=_DAY_TRANSACTION_TIMEOUT) as transaction:
|
||||
await transaction.execute_raw(_LOCK_GLOBAL_TABLE_SQL)
|
||||
await transaction.execute_raw(RECONCILE_DAY_SQL, day)
|
||||
"""Rewrite one day of the global table from the per-key sums. Idempotent: a rerun
|
||||
overwrites every group with the same totals."""
|
||||
await prisma_client.db.execute_raw(RECONCILE_DAY_SQL, day)
|
||||
|
||||
|
||||
async def run_daily_global_spend_reconcile(
|
||||
|
|
|
|||
|
|
@ -1,19 +1,12 @@
|
|||
"""Tests for the single-statement daily spend upsert (LIT-5291)."""
|
||||
|
||||
import pathlib
|
||||
import re
|
||||
from typing import Final
|
||||
|
||||
import psycopg
|
||||
import pytest
|
||||
from psycopg.rows import dict_row
|
||||
from pytest_postgresql import factories
|
||||
|
||||
from litellm.proxy.db.daily_spend_bulk_upsert import (
|
||||
DAILY_SPEND_TABLES,
|
||||
GLOBAL_SPEND_TABLE,
|
||||
build_bulk_upsert,
|
||||
build_bulk_upsert_with_global_rollup,
|
||||
conflict_key,
|
||||
merge_by_conflict_key,
|
||||
)
|
||||
|
|
@ -192,146 +185,3 @@ async def test_writer_survives_a_transaction_whose_key_columns_are_null():
|
|||
_, params = prisma_client.db.statements[0]
|
||||
assert None not in params[:9]
|
||||
assert transactions == {}
|
||||
|
||||
|
||||
def user_txn(**overrides):
|
||||
txn = {**tag_txn(), "user_id": "u-1", **overrides}
|
||||
del txn["tag"]
|
||||
del txn["request_id"]
|
||||
return txn
|
||||
|
||||
|
||||
def _bound_rows(insert_sql: str, params: tuple[object, ...]) -> list[dict[str, object]]:
|
||||
"""Each VALUES row of one INSERT as a column -> bound value mapping, consuming params in order."""
|
||||
header = re.search(r"INSERT INTO \"[A-Za-z_]+\" \(([^)]*)\)", insert_sql)
|
||||
assert header is not None, insert_sql
|
||||
columns = [c.strip('"') for c in header.group(1).split(", ") if c != '"updated_at"']
|
||||
row_count = insert_sql.split("ON CONFLICT", 1)[0].count("(NOW() AT TIME ZONE 'UTC'))")
|
||||
return [dict(zip(columns, params[i * len(columns) : (i + 1) * len(columns)])) for i in range(row_count)]
|
||||
|
||||
|
||||
def test_global_rollup_folds_every_key_and_user_into_one_row_per_dimension_tuple():
|
||||
"""The global table has no api_key or user_id, so a batch spread over many keys and
|
||||
users must collapse to one row per (date, model, group, provider, mcp, endpoint)."""
|
||||
batch = merge_by_conflict_key(
|
||||
USER_TABLE,
|
||||
tuple(user_txn(user_id=f"u-{i}", api_key=f"sk-{i}", spend=1.0, api_requests=1) for i in range(5))
|
||||
+ (user_txn(user_id="u-0", api_key="sk-0", model="claude", spend=10.0, api_requests=3),),
|
||||
)
|
||||
|
||||
sql, params = build_bulk_upsert_with_global_rollup(USER_TABLE, batch)
|
||||
|
||||
entity_insert, global_insert = sql.split("RETURNING 1)")
|
||||
entity_rows = _bound_rows(entity_insert, params)
|
||||
global_rows = _bound_rows(global_insert, params[len(entity_rows) * len(entity_rows[0]) :])
|
||||
assert len(entity_rows) == 6
|
||||
assert 'INSERT INTO "LiteLLM_DailyGlobalSpend"' in global_insert
|
||||
assert [(r["model"], r["spend"], r["api_requests"]) for r in global_rows] == [
|
||||
("claude", 10.0, 3),
|
||||
("gpt-4o-mini", 5.0, 5),
|
||||
]
|
||||
assert all("api_key" not in r and "user_id" not in r for r in global_rows)
|
||||
conflict = re.search(r"ON CONFLICT \(([^)]*)\)", global_insert)
|
||||
assert conflict is not None
|
||||
assert conflict.group(1) == ", ".join(f'"{c}"' for c in GLOBAL_SPEND_TABLE.key_columns)
|
||||
|
||||
|
||||
def test_global_rollup_params_follow_the_entity_params_in_one_placeholder_sequence():
|
||||
"""Both inserts bind from one flat tuple, so the global arm's placeholders must start
|
||||
exactly where the entity arm's stop or every value lands one column off."""
|
||||
batch = merge_by_conflict_key(USER_TABLE, (user_txn(),))
|
||||
|
||||
sql, params = build_bulk_upsert_with_global_rollup(USER_TABLE, batch)
|
||||
|
||||
placeholders = [int(n) for n in re.findall(r"\$(\d+)::", sql)]
|
||||
assert placeholders == list(range(1, len(params) + 1))
|
||||
|
||||
|
||||
_bulk_upsert_postgresql_proc: Final = factories.postgresql_proc()
|
||||
_bulk_upsert_postgresql: Final = factories.postgresql("_bulk_upsert_postgresql_proc")
|
||||
|
||||
_MIGRATIONS_DIR: Final = (
|
||||
pathlib.Path(__file__).resolve().parents[4] / "litellm-proxy-extras" / "litellm_proxy_extras" / "migrations"
|
||||
)
|
||||
_GLOBAL_SPEND_MIGRATION: Final = _MIGRATIONS_DIR / "20260915000000_add_daily_global_spend" / "migration.sql"
|
||||
|
||||
_DAILY_USER_SPEND_DDL: Final = """
|
||||
CREATE TABLE "LiteLLM_DailyUserSpend" (
|
||||
id TEXT PRIMARY KEY,
|
||||
user_id TEXT,
|
||||
date TEXT NOT NULL,
|
||||
api_key TEXT NOT NULL,
|
||||
model TEXT,
|
||||
model_group TEXT,
|
||||
custom_llm_provider TEXT,
|
||||
mcp_namespaced_tool_name TEXT,
|
||||
endpoint TEXT,
|
||||
prompt_tokens BIGINT DEFAULT 0,
|
||||
completion_tokens BIGINT DEFAULT 0,
|
||||
cache_read_input_tokens BIGINT DEFAULT 0,
|
||||
cache_creation_input_tokens BIGINT DEFAULT 0,
|
||||
compression_saved_tokens BIGINT DEFAULT 0,
|
||||
compression_savings_spend DOUBLE PRECISION DEFAULT 0,
|
||||
prompt_caching_savings_spend DOUBLE PRECISION DEFAULT 0,
|
||||
gateway_injected_caching_savings_spend DOUBLE PRECISION DEFAULT 0,
|
||||
autorouter_savings_spend DOUBLE PRECISION DEFAULT 0,
|
||||
spend DOUBLE PRECISION DEFAULT 0,
|
||||
api_requests BIGINT DEFAULT 0,
|
||||
successful_requests BIGINT DEFAULT 0,
|
||||
failed_requests BIGINT DEFAULT 0,
|
||||
created_at TIMESTAMP DEFAULT now(),
|
||||
updated_at TIMESTAMP,
|
||||
UNIQUE (user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint)
|
||||
)
|
||||
"""
|
||||
|
||||
|
||||
def _execute_dollar_sql(conn: psycopg.Connection, sql: str, params: tuple[object, ...]) -> None:
|
||||
converted: Final = re.sub(r"\$(\d+)", r"%(p\1)s", sql)
|
||||
conn.execute(
|
||||
converted, # pyright: ignore[reportArgumentType] # psycopg stubs want a literal-typed query
|
||||
{f"p{i}": v for i, v in enumerate(params, start=1)},
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
|
||||
def test_global_rollup_equals_the_per_key_sums_after_repeated_flushes(_bulk_upsert_postgresql: psycopg.Connection):
|
||||
"""Against real Postgres and the shipped migration: two flushes of a mixed batch leave
|
||||
the global table exactly equal to the per-key table summed over user and key, with the
|
||||
NULL and '' spellings of a dimension folded into one row."""
|
||||
conn: Final = _bulk_upsert_postgresql
|
||||
conn.execute(_DAILY_USER_SPEND_DDL) # pyright: ignore[reportArgumentType] # DDL literal
|
||||
conn.execute(_GLOBAL_SPEND_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # DDL literal
|
||||
conn.commit()
|
||||
|
||||
batch = merge_by_conflict_key(
|
||||
USER_TABLE,
|
||||
(
|
||||
user_txn(user_id="u-1", api_key="sk-1", spend=1.0, prompt_tokens=10),
|
||||
user_txn(user_id="u-2", api_key="sk-2", spend=2.0, prompt_tokens=20),
|
||||
user_txn(user_id="u-1", api_key="sk-3", model=None, custom_llm_provider=None, spend=4.0),
|
||||
user_txn(user_id="u-3", api_key="sk-4", model="", custom_llm_provider="", spend=8.0),
|
||||
),
|
||||
)
|
||||
sql, params = build_bulk_upsert_with_global_rollup(USER_TABLE, batch)
|
||||
_execute_dollar_sql(conn, sql, params)
|
||||
_execute_dollar_sql(conn, sql, params)
|
||||
|
||||
with conn.cursor(row_factory=dict_row) as cur:
|
||||
global_rows = cur.execute(
|
||||
'SELECT model, spend, prompt_tokens, api_requests FROM "LiteLLM_DailyGlobalSpend" ORDER BY model'
|
||||
).fetchall()
|
||||
per_key = cur.execute(
|
||||
"""
|
||||
SELECT COALESCE(model, '') AS model, SUM(spend) AS spend, SUM(prompt_tokens) AS prompt_tokens,
|
||||
SUM(api_requests) AS api_requests
|
||||
FROM "LiteLLM_DailyUserSpend" GROUP BY COALESCE(model, '') ORDER BY 1
|
||||
"""
|
||||
).fetchall()
|
||||
|
||||
assert [row["model"] for row in global_rows] == ["", "gpt-4o-mini"]
|
||||
assert [(r["model"], r["spend"], int(r["prompt_tokens"]), int(r["api_requests"])) for r in global_rows] == [
|
||||
(r["model"], float(r["spend"]), int(r["prompt_tokens"]), int(r["api_requests"])) for r in per_key
|
||||
]
|
||||
assert global_rows[0]["spend"] == pytest.approx(24.0)
|
||||
assert global_rows[1]["spend"] == pytest.approx(6.0)
|
||||
|
|
|
|||
|
|
@ -254,19 +254,14 @@ class _RecordingPrisma:
|
|||
|
||||
|
||||
def _row_values(statement: Statement, column: str) -> list[object]:
|
||||
"""Every row's value for one column of the first INSERT, read out of the flat parameter tuple.
|
||||
|
||||
The user-table statement chains a global rollup INSERT after its own, so the row count
|
||||
comes from the first INSERT's VALUES rather than from the parameter count.
|
||||
"""
|
||||
"""Every row's value for one column, read out of the flat parameter tuple."""
|
||||
sql, params = statement
|
||||
header = re.search(r"INSERT INTO \"[A-Za-z_]+\" \(([^)]*)\)", sql)
|
||||
assert header is not None, sql
|
||||
columns = header.group(1).split(", ")
|
||||
stride = len(columns) - 1 # updated_at is inlined, not bound
|
||||
offset = columns.index(f'"{column}"')
|
||||
rows = sql.split("ON CONFLICT", 1)[0].count("(NOW() AT TIME ZONE 'UTC'))")
|
||||
return [params[row * stride + offset] for row in range(rows)]
|
||||
return [params[row * stride + offset] for row in range(len(params) // stride)]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1468,98 +1463,6 @@ async def test_update_daily_spend_keeps_failed_transactions_for_retry():
|
|||
assert daily_spend_transactions == expected
|
||||
|
||||
|
||||
def _entity_txn(entity_field: str, entity_id: str, api_key: str) -> dict[str, object]:
|
||||
txn = _daily_txn()
|
||||
del txn["user_id"]
|
||||
return {**txn, entity_field: entity_id, "api_key": api_key}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_flush_writes_the_global_rollup_in_the_same_statement():
|
||||
"""The user flush is the one place per-key spend becomes key-free spend, so a batch spread
|
||||
over many keys must land in LiteLLM_DailyGlobalSpend as one row in the same statement.
|
||||
A separate statement would let a crash between the two leave the tables out of sync."""
|
||||
prisma_client = _RecordingPrisma()
|
||||
txns = {f"k{i}": _entity_txn("user_id", f"user-{i}", f"sk-{i}") for i in range(4)}
|
||||
|
||||
await DBSpendUpdateWriter._update_daily_spend(
|
||||
n_retry_times=0,
|
||||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
daily_spend_transactions=txns,
|
||||
entity_type="user",
|
||||
entity_id_field="user_id",
|
||||
)
|
||||
|
||||
assert len(prisma_client.db.statements) == 1
|
||||
sql, params = prisma_client.db.statements[0]
|
||||
assert sql.count('INSERT INTO "LiteLLM_DailyUserSpend"') == 1
|
||||
assert sql.count('INSERT INTO "LiteLLM_DailyGlobalSpend"') == 1
|
||||
assert sql.index('"LiteLLM_DailyUserSpend"') < sql.index('"LiteLLM_DailyGlobalSpend"')
|
||||
global_insert = sql.split('INSERT INTO "LiteLLM_DailyGlobalSpend"', 1)[1]
|
||||
assert global_insert.split("ON CONFLICT", 1)[0].count("(NOW() AT TIME ZONE 'UTC'))") == 1
|
||||
assert "api_key" not in global_insert
|
||||
assert params.count(0.4) == 1
|
||||
assert txns == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("entity_type", "entity_field"),
|
||||
[
|
||||
("team", "team_id"),
|
||||
("org", "organization_id"),
|
||||
("tag", "tag"),
|
||||
("end_user", "end_user_id"),
|
||||
("agent", "agent_id"),
|
||||
],
|
||||
)
|
||||
async def test_other_entity_flushes_leave_the_global_table_alone(entity_type, entity_field):
|
||||
"""Every entity table sees the same request, so writing the rollup from more than one of
|
||||
them would count each request once per entity type."""
|
||||
prisma_client = _RecordingPrisma()
|
||||
txn = _entity_txn(entity_field, "e-1", "sk-1")
|
||||
if entity_type == "tag":
|
||||
txn["request_id"] = "req-1"
|
||||
|
||||
await DBSpendUpdateWriter._update_daily_spend(
|
||||
n_retry_times=0,
|
||||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
daily_spend_transactions={"k": txn},
|
||||
entity_type=entity_type,
|
||||
entity_id_field=entity_field,
|
||||
)
|
||||
|
||||
(sql, _params) = prisma_client.db.statements[0]
|
||||
assert "LiteLLM_DailyGlobalSpend" not in sql
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_failed_chained_user_flush_keeps_every_transaction_for_retry():
|
||||
def raise_outage():
|
||||
raise ValueError("simulated database outage")
|
||||
|
||||
prisma_client = _RecordingPrisma(execute_raw=raise_outage)
|
||||
txns = {f"k{i}": _entity_txn("user_id", f"user-{i}", f"sk-{i}") for i in range(3)}
|
||||
expected = dict(txns)
|
||||
mock_proxy_logging = MagicMock()
|
||||
mock_proxy_logging.failure_handler = AsyncMock()
|
||||
|
||||
with pytest.raises(ValueError, match="simulated database outage"):
|
||||
await DBSpendUpdateWriter._update_daily_spend(
|
||||
n_retry_times=0,
|
||||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
daily_spend_transactions=txns,
|
||||
entity_type="user",
|
||||
entity_id_field="user_id",
|
||||
)
|
||||
|
||||
assert txns == expected
|
||||
assert 'INSERT INTO "LiteLLM_DailyGlobalSpend"' in prisma_client.db.statements[0][0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_commit_key_spend_updates_includes_last_active():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import pathlib
|
||||
import re
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
|
@ -10,11 +11,6 @@ import pytest
|
|||
from psycopg.rows import dict_row
|
||||
from pytest_postgresql import factories
|
||||
|
||||
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
|
||||
|
||||
|
||||
import pathlib
|
||||
|
||||
from litellm.constants import (
|
||||
DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM,
|
||||
PTU_SENTINEL_API_KEY,
|
||||
|
|
@ -29,10 +25,11 @@ from litellm.proxy.management_endpoints.common_daily_activity import (
|
|||
get_api_key_metadata,
|
||||
get_daily_activity,
|
||||
get_daily_activity_aggregated,
|
||||
key_free_source_table,
|
||||
global_rollup_reconciled_through,
|
||||
update_metrics,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.daily_global_spend_rollup import RECONCILE_DAY_SQL
|
||||
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
|
||||
from litellm.proxy.utils import evict_config_param
|
||||
from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
||||
DailySpendMetadata,
|
||||
|
|
@ -1632,7 +1629,9 @@ def _prisma_with_marker(marker: str | None) -> MagicMock:
|
|||
prisma.db = MagicMock()
|
||||
prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
|
||||
prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
|
||||
row = None if marker is None else SimpleNamespace(param_name="m", param_value=f'{{"reconciled_through": "{marker}"}}')
|
||||
row = (
|
||||
None if marker is None else SimpleNamespace(param_name="m", param_value=f'{{"reconciled_through": "{marker}"}}')
|
||||
)
|
||||
prisma.get_generic_data = AsyncMock(return_value=row)
|
||||
return prisma
|
||||
|
||||
|
|
@ -1657,9 +1656,9 @@ def _unfiltered_user_query(**overrides):
|
|||
@pytest.mark.parametrize(
|
||||
("marker", "overrides", "expected"),
|
||||
[
|
||||
("2026-06-02", {}, "LiteLLM_DailyGlobalSpend"),
|
||||
("2026-06-02", {"model": "gpt-5"}, "LiteLLM_DailyGlobalSpend"),
|
||||
("2026-06-01", {}, None),
|
||||
("2026-06-02", {}, "2026-06-02"),
|
||||
("2026-06-02", {"model": "gpt-5"}, "2026-06-02"),
|
||||
("2026-05-01", {}, "2026-05-01"),
|
||||
(None, {}, None),
|
||||
("2026-06-02", {"api_key": "sk-1"}, None),
|
||||
("2026-06-02", {"api_key": []}, None),
|
||||
|
|
@ -1668,33 +1667,50 @@ def _unfiltered_user_query(**overrides):
|
|||
("2026-06-02", {"table_name": "litellm_dailyteamspend", "entity_id_field": "team_id"}, None),
|
||||
],
|
||||
)
|
||||
async def test_key_free_source_table_routes_only_unfiltered_user_reads_within_the_marker(marker, overrides, expected):
|
||||
"""Anything that filters by key or entity has no counterpart in the global table, and a
|
||||
range the reconcile has not reached must stay on the per-key table."""
|
||||
async def test_global_rollup_marker_is_used_only_for_unfiltered_user_reads(marker, overrides, expected):
|
||||
"""Anything that filters by key or entity has no counterpart in the global table; the
|
||||
SQL splits the range at the marker itself, so the marker passes through unchanged."""
|
||||
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
|
||||
prisma = _prisma_with_marker(marker)
|
||||
|
||||
assert await key_free_source_table(prisma, _unfiltered_user_query(**overrides)) == expected
|
||||
assert await global_rollup_reconciled_through(prisma, _unfiltered_user_query(**overrides)) == expected
|
||||
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_free_source_table_judges_the_timezone_extended_end_not_the_requested_one():
|
||||
"""A caller west of UTC asking through their local today gets today's UTC bucket added to
|
||||
the range; the marker must cover that extended day, not just the requested end."""
|
||||
async def test_global_rollup_marker_read_failure_falls_back_to_the_per_key_table():
|
||||
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
|
||||
today_utc: Final = datetime.now(timezone.utc).date()
|
||||
yesterday: Final = (today_utc - timedelta(days=1)).isoformat()
|
||||
query: Final = _unfiltered_user_query(
|
||||
start_date=yesterday, end_date=yesterday, timezone_offset_minutes=24 * 60, include_current_utc_day=True
|
||||
)
|
||||
prisma = _prisma_with_marker(None)
|
||||
prisma.get_generic_data = AsyncMock(side_effect=RuntimeError("db down"))
|
||||
|
||||
assert await key_free_source_table(_prisma_with_marker(yesterday), query) is None
|
||||
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
|
||||
assert await key_free_source_table(_prisma_with_marker(today_utc.isoformat()), query) == "LiteLLM_DailyGlobalSpend"
|
||||
assert await global_rollup_reconciled_through(prisma, _unfiltered_user_query()) is None
|
||||
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
|
||||
|
||||
|
||||
def test_aggregated_sql_splits_the_key_free_arm_at_the_marker_and_keeps_the_key_arm_per_key():
|
||||
sql, params = _build_aggregated_sql_query(**_unfiltered_user_query(), global_rollup_through="2026-06-01")
|
||||
marker_param: Final = f"${len(params)}"
|
||||
|
||||
assert params[-1] == "2026-06-01"
|
||||
assert (
|
||||
f'FROM "LiteLLM_DailyGlobalSpend"\n WHERE date >= $1 AND date <= $2 AND date <= {marker_param}'
|
||||
in sql
|
||||
)
|
||||
assert (
|
||||
f'FROM "LiteLLM_DailyUserSpend"\n WHERE date >= $1 AND date <= $2 AND date > {marker_param}' in sql
|
||||
)
|
||||
key_arm: Final = sql.split("UNION ALL\n (WITH top_api_keys")[1]
|
||||
assert "LiteLLM_DailyGlobalSpend" not in key_arm
|
||||
assert marker_param not in key_arm
|
||||
|
||||
|
||||
def test_aggregated_sql_without_a_marker_reads_the_per_key_table_only():
|
||||
sql, params = _build_aggregated_sql_query(**_unfiltered_user_query())
|
||||
|
||||
assert "LiteLLM_DailyGlobalSpend" not in sql
|
||||
assert params[-1] == PTU_SENTINEL_API_KEY
|
||||
|
||||
|
||||
_GLOBAL_SPEND_MIGRATION: Final = (
|
||||
pathlib.Path(__file__).resolve().parents[4]
|
||||
/ "litellm-proxy-extras"
|
||||
|
|
@ -1706,12 +1722,12 @@ _GLOBAL_SPEND_MIGRATION: Final = (
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_daily_activity_aggregated_reads_the_global_table_for_the_key_free_arm(
|
||||
async def test_get_daily_activity_aggregated_serves_closed_days_from_the_global_table_and_open_days_live(
|
||||
_aggregated_postgresql: psycopg.Connection,
|
||||
):
|
||||
"""With the range reconciled, the key-free arm reads LiteLLM_DailyGlobalSpend while the
|
||||
per-key arm stays on the user table, and the response is identical to the all-per-key
|
||||
read: same totals, same rollups, same top keys."""
|
||||
"""Day 1 is rolled up and day 2 is still open (never rolled up), so a marker of day 1 must
|
||||
give the same response as reading everything per-key: day 1 from the global table, day 2
|
||||
live, one grand total across both. The per-key arm stays on the user table throughout."""
|
||||
n_keys: Final = USAGE_TOP_API_KEYS_LIMIT + 3
|
||||
rows: Final = [
|
||||
(
|
||||
|
|
@ -1734,11 +1750,10 @@ async def test_get_daily_activity_aggregated_reads_the_global_table_for_the_key_
|
|||
_seed_daily_user_spend(_aggregated_postgresql, rows)
|
||||
with _aggregated_postgresql.cursor() as cur:
|
||||
cur.execute(_GLOBAL_SPEND_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # DDL literal
|
||||
for day in ("2026-06-01", "2026-06-02"):
|
||||
cur.execute(
|
||||
re.sub(r"\$(\d+)", r"%(p\1)s", RECONCILE_DAY_SQL), # pyright: ignore[reportArgumentType] # $N -> psycopg
|
||||
{"p1": day},
|
||||
)
|
||||
cur.execute(
|
||||
re.sub(r"\$(\d+)", r"%(p\1)s", RECONCILE_DAY_SQL), # pyright: ignore[reportArgumentType] # $N -> psycopg
|
||||
{"p1": "2026-06-01"},
|
||||
)
|
||||
_aggregated_postgresql.commit()
|
||||
|
||||
async def read(marker: str | None, sql_seen: list[str]):
|
||||
|
|
@ -1760,16 +1775,19 @@ async def test_get_daily_activity_aggregated_reads_the_global_table_for_the_key_
|
|||
per_key_sql: Final[list[str]] = [] # mutable-ok: out-param for the query_raw shim
|
||||
global_sql: Final[list[str]] = [] # mutable-ok: out-param for the query_raw shim
|
||||
from_per_key = await read(None, per_key_sql)
|
||||
from_global = await read("2026-06-02", global_sql)
|
||||
from_global = await read("2026-06-01", global_sql)
|
||||
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
|
||||
|
||||
assert per_key_sql[0].count('FROM "LiteLLM_DailyGlobalSpend"') == 0
|
||||
assert global_sql[0].count('FROM "LiteLLM_DailyGlobalSpend"') == 1
|
||||
assert global_sql[0].count('FROM "LiteLLM_DailyUserSpend"') == 2
|
||||
assert global_sql[0].count('FROM "LiteLLM_DailyUserSpend"') == 3
|
||||
assert from_global.model_dump() == from_per_key.model_dump()
|
||||
assert from_global.metadata.total_spend == pytest.approx(2 * sum(float(i + 1) for i in range(n_keys)))
|
||||
assert {day.date.isoformat() for day in from_global.results} == {"2026-06-01", "2026-06-02"}
|
||||
assert len(from_global.results[0].breakdown.api_keys) == USAGE_TOP_API_KEYS_LIMIT
|
||||
assert set(from_global.results[0].breakdown.model_groups) == {"gpt-5", "claude"}
|
||||
|
||||
|
||||
def _no_spend_record():
|
||||
"""A rollup row for a key with no spend, where SUM() returns NULL (None)."""
|
||||
return SimpleNamespace(
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@
|
|||
|
||||
import pathlib
|
||||
import re
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import date
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
|
@ -13,11 +12,7 @@ from psycopg.rows import dict_row
|
|||
from pytest_postgresql import factories
|
||||
|
||||
from litellm.constants import DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM
|
||||
from litellm.proxy.db.daily_spend_bulk_upsert import (
|
||||
DAILY_SPEND_TABLES,
|
||||
build_bulk_upsert_with_global_rollup,
|
||||
merge_by_conflict_key,
|
||||
)
|
||||
from litellm.proxy.db.daily_spend_bulk_upsert import DAILY_SPEND_TABLES, build_bulk_upsert, merge_by_conflict_key
|
||||
from litellm.proxy.spend_tracking.daily_global_spend_rollup import (
|
||||
RECONCILE_DAY_SQL,
|
||||
reconciled_through,
|
||||
|
|
@ -45,21 +40,6 @@ class _FakeConfigTable:
|
|||
return _FakeConfigRow(where["param_name"], data["update"]["param_value"])
|
||||
|
||||
|
||||
class _FakeTransaction:
|
||||
def __init__(self, prisma: "_FakePrisma") -> None:
|
||||
self._prisma = prisma
|
||||
|
||||
async def execute_raw(self, sql: str, *params: str) -> int:
|
||||
if "LOCK TABLE" in sql:
|
||||
self._prisma.locks_taken += 1
|
||||
return 0
|
||||
(day,) = params
|
||||
if day in self._prisma.failing_days:
|
||||
raise RuntimeError(f"day {day} exploded")
|
||||
self._prisma.reconciled.append(day)
|
||||
return 1
|
||||
|
||||
|
||||
class _FakeDb:
|
||||
def __init__(self, prisma: "_FakePrisma") -> None:
|
||||
self._prisma = prisma
|
||||
|
|
@ -69,19 +49,21 @@ class _FakeDb:
|
|||
first, last = params
|
||||
return [{"date": d} for d in sorted(self._prisma.user_days) if first <= d <= last]
|
||||
|
||||
@asynccontextmanager
|
||||
async def tx(self, timeout: object):
|
||||
yield _FakeTransaction(self._prisma)
|
||||
async def execute_raw(self, sql: str, *params: str) -> int:
|
||||
(day,) = params
|
||||
if day in self._prisma.failing_days:
|
||||
raise RuntimeError(f"day {day} exploded")
|
||||
self._prisma.reconciled.append(day)
|
||||
return 1
|
||||
|
||||
|
||||
class _FakePrisma:
|
||||
"""Enough of PrismaClient for the reconcile: per-key dates, a config table, and a transaction."""
|
||||
"""Enough of PrismaClient for the reconcile: per-key dates, a config table, and execute_raw."""
|
||||
|
||||
def __init__(self, user_days: tuple[str, ...], failing_days: frozenset[str] = frozenset()) -> None:
|
||||
self.user_days = user_days
|
||||
self.failing_days = failing_days
|
||||
self.reconciled: list[str] = []
|
||||
self.locks_taken = 0
|
||||
self.db = _FakeDb(self)
|
||||
|
||||
async def get_generic_data(self, key: str, value: str, table_name: str) -> _FakeConfigRow | None:
|
||||
|
|
@ -97,33 +79,45 @@ async def _fresh_marker_cache():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_first_run_rolls_up_every_historical_day_and_today_then_marks_today():
|
||||
"""Before any marker exists, every day with per-key rows is rolled up, plus today even
|
||||
with no rows yet, so reads for ranges ending today can switch to the global table."""
|
||||
prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-03", "2026-09-14"))
|
||||
async def test_first_run_rolls_up_every_closed_day_and_never_today():
|
||||
"""Before any marker exists every closed day with per-key rows is rolled up. Today is left
|
||||
out: pods are still flushing it, so it is served live from the per-key table until it closes."""
|
||||
prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-03", "2026-09-14", "2026-09-15"))
|
||||
|
||||
result = await run_daily_global_spend_reconcile(prisma, today=TODAY)
|
||||
|
||||
assert result.days_reconciled == ("2026-09-01", "2026-09-03", "2026-09-14", "2026-09-15")
|
||||
assert result.days_reconciled == ("2026-09-01", "2026-09-03", "2026-09-14")
|
||||
assert result.failed_day is None
|
||||
assert result.reconciled_through == "2026-09-15"
|
||||
assert await reconciled_through(prisma) == "2026-09-15"
|
||||
assert prisma.locks_taken == 4
|
||||
assert result.reconciled_through == "2026-09-14"
|
||||
assert await reconciled_through(prisma) == "2026-09-14"
|
||||
assert "2026-09-15" not in prisma.reconciled
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_later_run_replays_the_marker_day_and_the_day_before_only():
|
||||
"""Days older than marker-1 are settled; the marker day and its predecessor are replayed so
|
||||
rows a pre-writer pod flushed around midnight during a rolling deploy get folded in."""
|
||||
per-key rows that landed after their day was rolled up get folded in."""
|
||||
prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-12", "2026-09-13", "2026-09-14"))
|
||||
await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 13))
|
||||
await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14))
|
||||
prisma.reconciled.clear()
|
||||
|
||||
result = await run_daily_global_spend_reconcile(prisma, today=TODAY)
|
||||
|
||||
assert result.days_reconciled == ("2026-09-12", "2026-09-13", "2026-09-14", "2026-09-15")
|
||||
assert result.days_reconciled == ("2026-09-12", "2026-09-13", "2026-09-14")
|
||||
assert "2026-09-01" not in prisma.reconciled
|
||||
assert await reconciled_through(prisma) == "2026-09-15"
|
||||
assert await reconciled_through(prisma) == "2026-09-14"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_run_with_no_new_closed_days_keeps_the_marker():
|
||||
prisma = _FakePrisma(user_days=("2026-09-13",))
|
||||
await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14))
|
||||
prisma.reconciled.clear()
|
||||
|
||||
result = await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14))
|
||||
|
||||
assert result.days_reconciled == ("2026-09-13",)
|
||||
assert result.reconciled_through == "2026-09-13"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -149,16 +143,16 @@ async def test_the_next_run_resumes_from_the_failed_day():
|
|||
|
||||
result = await run_daily_global_spend_reconcile(prisma, today=TODAY)
|
||||
|
||||
assert result.days_reconciled == ("2026-09-01", "2026-09-02", "2026-09-03", "2026-09-15")
|
||||
assert await reconciled_through(prisma) == "2026-09-15"
|
||||
assert result.days_reconciled == ("2026-09-01", "2026-09-02", "2026-09-03")
|
||||
assert await reconciled_through(prisma) == "2026-09-03"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_failure_with_nothing_done_reports_the_previous_marker_and_alerts():
|
||||
"""A pre-writer pod flushing rows for the day before the marker is exactly the replay case;
|
||||
when that replay fails the marker must stay put and the operator must hear about it."""
|
||||
"""A late flush for the day before the marker is exactly the replay case; when that replay
|
||||
fails the marker must stay put and the operator must hear about it."""
|
||||
prisma = _FakePrisma(user_days=("2026-09-13",))
|
||||
await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 13))
|
||||
await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14))
|
||||
prisma.user_days = ("2026-09-12", "2026-09-13")
|
||||
prisma.failing_days = frozenset({"2026-09-12"})
|
||||
alert = AsyncMock()
|
||||
|
|
@ -212,7 +206,7 @@ async def test_scheduled_run_runs_and_releases_the_lock_when_it_wins():
|
|||
|
||||
result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock, today=TODAY)
|
||||
|
||||
assert result is not None and result.days_reconciled == ("2026-09-13", "2026-09-15")
|
||||
assert result is not None and result.days_reconciled == ("2026-09-13",)
|
||||
lock.release_lock.assert_awaited_once()
|
||||
|
||||
|
||||
|
|
@ -226,7 +220,7 @@ async def test_scheduled_run_proceeds_when_the_lock_cannot_be_acquired_or_read()
|
|||
|
||||
result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock, today=TODAY)
|
||||
|
||||
assert result is not None and result.days_reconciled == ("2026-09-13", "2026-09-15")
|
||||
assert result is not None and result.days_reconciled == ("2026-09-13",)
|
||||
lock.release_lock.assert_not_awaited()
|
||||
|
||||
|
||||
|
|
@ -341,9 +335,9 @@ def _normalized(rows: list[dict[str, object]]) -> list[tuple[object, ...]]:
|
|||
|
||||
|
||||
def test_reconcile_day_sql_makes_the_global_day_equal_the_per_key_sums(_rollup_postgresql: psycopg.Connection):
|
||||
"""Against real Postgres and the shipped migration: rows the writer never saw (a
|
||||
pre-writer pod's flush, NULL and '' dimension spellings) end up folded into the global
|
||||
day, running the day twice changes nothing, and other days are left alone."""
|
||||
"""Against real Postgres and the shipped migration: writer-shaped rows and legacy rows
|
||||
(NULL and '' dimension spellings) fold into one global day, running the day twice changes
|
||||
nothing, and other days are left alone."""
|
||||
conn: Final = _rollup_postgresql
|
||||
conn.execute(_DAILY_USER_SPEND_DDL) # pyright: ignore[reportArgumentType] # DDL literal
|
||||
conn.execute(_GLOBAL_SPEND_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # DDL literal
|
||||
|
|
@ -353,7 +347,7 @@ def test_reconcile_day_sql_makes_the_global_day_equal_the_per_key_sums(_rollup_p
|
|||
USER_TABLE,
|
||||
(_user_txn(api_key="sk-1", spend=1.0), _user_txn(api_key="sk-2", user_id="u-2", spend=2.0, prompt_tokens=20)),
|
||||
)
|
||||
_execute_dollar_sql(conn, *build_bulk_upsert_with_global_rollup(USER_TABLE, written_batch))
|
||||
_execute_dollar_sql(conn, *build_bulk_upsert(USER_TABLE, written_batch))
|
||||
|
||||
conn.execute(
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue