litellm/tests/integration/spend/test_shutdown_flush.py
devin-ai-integration[bot] a80379baf8
fix(proxy): keep the in-flight daily spend batch when shutdown cancels the flush (#42593)
* fix(proxy): keep the in-flight daily spend batch when shutdown cancels the flush

A daily spend batch drained from the in-memory queue was dropped for good when
the scheduler tick was cancelled by shutdown, because asyncio.CancelledError
bypasses the except Exception requeue. The flush now requeues the drained rows
on cancellation and re-raises, and each daily batch upsert runs in an
interactive transaction so a statement that already reached Postgres is rolled
back with the cancel instead of committing behind the requeue

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): requeue the cancelled daily spend batch before its rollback returns

Behind a lock the rollback of the cancelled interactive transaction only
returns once the blocked statement does, which is after the shutdown flush
has already run. The commit now runs as a shielded task so the cancelled
tick requeues the batch at once and lets the rollback finish in the
background. The final flush then finds the rows and writes them exactly once

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): give the recording db a transaction seam for the bulk upsert tests

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): route the mocked daily tag spend upsert through the transaction seam

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): restore the drained Redis tag batch when shutdown cancels its commit

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yucheng <yucheng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
2026-09-22 17:58:08 -07:00

189 lines
8 KiB
Python

import json
import os
import signal
import threading
import uuid
from collections.abc import Callable, Iterator
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import httpx
import psycopg
import pytest
import yaml
from integration._support.client import Gateway, delete_key_if_present, eventually, string_value
from integration._support.database import read_rows
from integration._support.process import OwnedProxy, owned_proxy_process
from integration._support.wire import Reply, Request, wire_server
REQUESTS_WHILE_BLOCKED: Final = 6
CANCEL_LOG_LINE: Final = "in-flight scheduled job(s) for shutdown"
BATCH_DRAINED_LOG_LINE: Final = f"flushed {REQUESTS_WHILE_BLOCKED} daily spend update items from in-memory queue"
MODEL_INSERT_ARRIVED_LOG_LINE: Final = "path=/model/new"
def _api_requests(table: str, column: str, identity: str) -> int:
rows: Final = read_rows(
f'SELECT coalesce(sum(api_requests), 0)::int AS total FROM "{table}" WHERE {column}=%s', (identity,)
)
total: Final = rows[0]["total"]
assert isinstance(total, int)
return total
def _waiting_on(table: str) -> int:
rows: Final = read_rows(
"SELECT count(*)::int AS waiting FROM pg_stat_activity WHERE wait_event_type='Lock' AND query LIKE %s",
(f'%"{table}"%',),
)
waiting: Final = rows[0]["waiting"]
assert isinstance(waiting, int)
return waiting
def _provider(request: Request) -> Reply:
if request.method != "POST":
return Reply(status=404, body=b'{"error":"not scripted"}')
assert request.target == "/v1/chat/completions"
return Reply(
body=json.dumps(
{
"id": "chatcmpl-" + uuid.uuid4().hex,
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4},
}
).encode()
)
@dataclass(frozen=True, slots=True)
class _Shutdown:
owner: str
team: str
owned: OwnedProxy
key: str
model: str
def chat(self) -> None:
body: Final = {"model": self.model, "messages": [{"role": "user", "content": f"spend {uuid.uuid4().hex}"}]}
assert self.owned.gateway.request("POST", "/v1/chat/completions", body, key=self.key).status_code == 200
def daily_user_requests(self) -> int:
return _api_requests("LiteLLM_DailyUserSpend", "user_id", self.owner)
def logged(self, line: str, times: int = 1) -> bool:
return self.owned.log.read_text(errors="replace").count(line) >= times
def chat_while_spend_update_is_blocked(self, blocker: psycopg.Connection, table: str) -> None:
blocker.execute(f'LOCK TABLE "{table}" IN EXCLUSIVE MODE')
for _ in range(REQUESTS_WHILE_BLOCKED):
self.chat()
eventually(lambda: _waiting_on(table), lambda waiting: waiting == 1, seconds=30)
def start_blocked_model_insert(self) -> threading.Thread:
body: Final = {
"model_name": f"integration-blocked-{uuid.uuid4().hex}",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "integration-provider-key"},
"model_info": {},
}
def insert() -> None:
try:
self.owned.gateway.request("POST", "/model/new", body)
except httpx.TransportError:
pass
thread: Final = threading.Thread(target=insert, daemon=True)
thread.start()
return thread
def terminate_once(self, blocked: Callable[[], bool], release: Callable[[], None]) -> None:
eventually(blocked, lambda state: state, seconds=60)
self.owned.process.send_signal(signal.SIGTERM)
eventually(lambda: self.logged(CANCEL_LOG_LINE), lambda seen: seen, seconds=60)
release()
self.owned.process.wait(timeout=120)
def _config_with_pool_limit(tmp_path: Path, pool_limit: int) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["general_settings"]["database_connection_pool_limit"] = pool_limit
config["general_settings"]["database_connection_pool_timeout"] = 60
path: Final = tmp_path / f"pool-{pool_limit}.yaml"
path.write_text(yaml.safe_dump(config))
return path
@contextmanager
def _proxy_with_one_seeded_row(gateway: Gateway, tmp_path: Path, pool_limit: int) -> Iterator[_Shutdown]:
owner: Final = f"integration-owner-{uuid.uuid4().hex}"
with gateway.scenario() as scenario, wire_server(_provider) as wire:
model: Final = scenario.model(api_base=wire.url + "/v1", num_retries=0)
team: Final = scenario.team(models=[model])
with owned_proxy_process(
gateway,
tmp_path,
{
"LITELLM_LOG": "DEBUG",
"GRACEFUL_SHUTDOWN_TIMEOUT": "1",
"SCHEDULED_JOB_SHUTDOWN_FINISH_TIMEOUT_SECONDS": "1",
"SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS": "5",
},
config=_config_with_pool_limit(tmp_path, pool_limit),
) as owned:
key: Final = string_value(
owned.gateway.post("/key/generate", {"user_id": owner, "team_id": team, "models": [model]})["key"]
)
scenario.cleanups.callback(delete_key_if_present, gateway, key)
shutdown: Final = _Shutdown(owner, team, owned, key, model)
shutdown.chat()
eventually(shutdown.daily_user_requests, lambda total: total == 1, seconds=60)
yield shutdown
assert _api_requests("LiteLLM_DailyUserSpend", "user_id", owner) == 1 + REQUESTS_WHILE_BLOCKED
assert _api_requests("LiteLLM_DailyTeamSpend", "team_id", team) == 1 + REQUESTS_WHILE_BLOCKED
@pytest.mark.covers("quota_management.spend_tracking.shutdown_cancel_keeps_in_flight_daily_batch")
def test_daily_spend_batch_cancelled_while_waiting_for_a_pool_connection_is_written_by_the_final_flush(
gateway: Gateway, tmp_path: Path
) -> None:
with (
_proxy_with_one_seeded_row(gateway, tmp_path, pool_limit=2) as shutdown,
psycopg.connect(os.environ["DATABASE_URL"]) as models,
psycopg.connect(os.environ["DATABASE_URL"]) as memberships,
):
models.execute('LOCK TABLE "LiteLLM_ProxyModelTable" IN EXCLUSIVE MODE')
first: Final = shutdown.start_blocked_model_insert()
eventually(lambda: _waiting_on("LiteLLM_ProxyModelTable"), lambda waiting: waiting == 1, seconds=30)
shutdown.chat_while_spend_update_is_blocked(memberships, "LiteLLM_TeamMembership")
second: Final = shutdown.start_blocked_model_insert()
eventually(lambda: shutdown.logged(MODEL_INSERT_ARRIVED_LOG_LINE, times=2), lambda seen: seen, seconds=30)
memberships.rollback()
eventually(lambda: _waiting_on("LiteLLM_ProxyModelTable"), lambda waiting: waiting == 2, seconds=30)
shutdown.terminate_once(lambda: shutdown.logged(BATCH_DRAINED_LOG_LINE), models.rollback)
first.join(timeout=30)
second.join(timeout=30)
@pytest.mark.covers("quota_management.spend_tracking.shutdown_cancel_keeps_in_flight_daily_batch")
def test_daily_spend_batch_cancelled_while_waiting_for_a_row_lock_is_written_exactly_once(
gateway: Gateway, tmp_path: Path
) -> None:
with (
_proxy_with_one_seeded_row(gateway, tmp_path, pool_limit=10) as shutdown,
psycopg.connect(os.environ["DATABASE_URL"]) as holder,
psycopg.connect(os.environ["DATABASE_URL"]) as memberships,
):
holder.execute('SELECT 1 FROM "LiteLLM_DailyUserSpend" WHERE user_id=%s FOR UPDATE', (shutdown.owner,))
shutdown.chat_while_spend_update_is_blocked(memberships, "LiteLLM_TeamMembership")
memberships.rollback()
shutdown.terminate_once(
lambda: shutdown.logged(BATCH_DRAINED_LOG_LINE) and _waiting_on("LiteLLM_DailyUserSpend") == 1,
holder.rollback,
)