feat(proxy): enhance response cost tracking with retry logic and batch cleanup

- Added retry mechanism for fetching responses in case of transient errors, improving reliability.
- Updated `_expire_stale_rows` method to use `update_many` for marking stale rows, enhancing database interaction.
- Introduced constants for retry attempts and delay, allowing for configurable response polling behavior.
- Enhanced logging to provide detailed information on stale object cleanup runs.
- Updated unit tests to cover new retry logic and ensure correct status updates for various response states.
This commit is contained in:
harish-berri 2026-05-11 22:45:52 +00:00
parent f3a041dec3
commit b4e598fe36
4 changed files with 195 additions and 65 deletions

View file

@ -3,14 +3,18 @@ Polls LiteLLM_ManagedObjectTable to check if the response is complete.
Cost tracking is handled automatically by litellm.aget_responses().
"""
import asyncio
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any, Dict, List, Optional
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS,
MAX_OBJECTS_PER_POLL_CYCLE,
RESPONSES_COST_POLL_RETRY_ATTEMPTS,
RESPONSES_COST_POLL_RETRY_DELAY_SECONDS,
STALE_OBJECT_CLEANUP_MAX_BATCHES_PER_POLL_CYCLE,
STALE_OBJECT_CLEANUP_BATCH_SIZE,
)
@ -36,32 +40,35 @@ class CheckResponsesCost:
async def _expire_stale_rows(
self, cutoff: datetime, batch_size: int
) -> int:
"""Execute the bounded UPDATE that marks stale rows as 'stale_expired'.
Isolated so it can be swapped / mocked in tests without touching the
orchestration logic in ``_cleanup_stale_managed_objects``.
Uses PostgreSQL syntax (``$1::timestamptz``, ``LIMIT``, double-quoted
identifiers) which is the only dialect the proxy supports — every
``schema.prisma`` in the repo sets ``provider = "postgresql"``.
Same pattern as ``spend_log_cleanup.py``.
"""
return await self.prisma_client.db.execute_raw(
"""
UPDATE "LiteLLM_ManagedObjectTable"
SET "status" = 'stale_expired'
WHERE "id" IN (
SELECT "id" FROM "LiteLLM_ManagedObjectTable"
WHERE "file_purpose" = 'response'
AND "status" NOT IN ('completed', 'complete', 'failed', 'expired', 'cancelled', 'stale_expired')
AND "created_at" < $1::timestamptz
ORDER BY "created_at" ASC
LIMIT $2
)
""",
cutoff,
batch_size,
"""Mark up to `batch_size` stale response rows as `stale_expired`."""
stale_rows: List[Any] = await self.prisma_client.db.litellm_managedobjecttable.find_many(
where={
"file_purpose": "response",
"status": {
"not_in": [
"completed",
"complete",
"failed",
"expired",
"cancelled",
"stale_expired",
]
},
"created_at": {"lt": cutoff},
},
take=batch_size,
order={"created_at": "asc"},
select={"id": True},
)
stale_ids = [row.id for row in stale_rows]
if not stale_ids:
return 0
await self.prisma_client.db.litellm_managedobjecttable.update_many(
where={"id": {"in": stale_ids}},
data={"status": "stale_expired"},
)
return len(stale_ids)
async def _cleanup_stale_managed_objects(self) -> None:
"""
@ -74,14 +81,56 @@ class CheckResponsesCost:
rows per invocation to avoid overwhelming the DB when there is a large
backlog.
"""
cutoff = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS)
result = await self._expire_stale_rows(cutoff, STALE_OBJECT_CLEANUP_BATCH_SIZE)
if result > 0:
cutoff = datetime.now(timezone.utc) - timedelta(
days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS
)
total_marked = 0
runs = 0
for _ in range(STALE_OBJECT_CLEANUP_MAX_BATCHES_PER_POLL_CYCLE):
result = await self._expire_stale_rows(cutoff, STALE_OBJECT_CLEANUP_BATCH_SIZE)
runs += 1
total_marked += result
if result < STALE_OBJECT_CLEANUP_BATCH_SIZE:
break
if total_marked > 0:
verbose_proxy_logger.warning(
f"CheckResponsesCost: marked {result} stale managed objects "
f"(older than {MANAGED_OBJECT_STALENESS_CUTOFF_DAYS} days) as stale_expired"
f"CheckResponsesCost: marked {total_marked} stale managed objects "
f"(older than {MANAGED_OBJECT_STALENESS_CUTOFF_DAYS} days) as stale_expired "
f"across {runs} cleanup run(s)"
)
async def _fetch_response_with_retries(
self, unified_object_id: str, metadata: Dict[str, str]
) -> Optional[Any]:
from litellm.proxy.hooks.responses_id_security import ResponsesIDSecurity
responses_id_security, _, _ = ResponsesIDSecurity()._decrypt_response_id(
unified_object_id
)
last_error: Optional[Exception] = None
for attempt in range(1, RESPONSES_COST_POLL_RETRY_ATTEMPTS + 1):
try:
return await litellm.aget_responses(
response_id=responses_id_security,
litellm_metadata=metadata,
)
except Exception as e:
last_error = e
if attempt == RESPONSES_COST_POLL_RETRY_ATTEMPTS:
break
await asyncio.sleep(RESPONSES_COST_POLL_RETRY_DELAY_SECONDS)
verbose_proxy_logger.info(
"Skipping job %s after %d failed poll attempt(s): %s",
unified_object_id,
RESPONSES_COST_POLL_RETRY_ATTEMPTS,
last_error,
)
return None
async def check_responses_cost(self):
"""
Check if background responses are complete and track their cost.
@ -107,23 +156,21 @@ class CheckResponsesCost:
)
verbose_proxy_logger.debug(f"Found {len(jobs)} response jobs to check")
completed_jobs = []
status_to_job_ids: Dict[str, List[str]] = {
"completed": [],
"failed": [],
"cancelled": [],
"expired": [],
}
for job in jobs:
unified_object_id = job.unified_object_id
try:
from litellm.proxy.hooks.responses_id_security import (
ResponsesIDSecurity,
)
# Get the stored response object to extract model information
stored_response = job.file_object
stored_response = job.file_object or {}
model_name = stored_response.get("model", None)
# Decrypt the response ID
responses_id_security, _, _ = ResponsesIDSecurity()._decrypt_response_id(unified_object_id)
# Prepare metadata with model information for cost tracking
litellm_metadata = {
"user_api_key_user_id": job.created_by or "default-user-id",
@ -133,16 +180,17 @@ class CheckResponsesCost:
if model_name:
litellm_metadata["model"] = model_name
litellm_metadata["model_group"] = model_name # Use same value for model_group
response = await litellm.aget_responses(
response_id=responses_id_security,
litellm_metadata=litellm_metadata,
response = await self._fetch_response_with_retries(
unified_object_id=unified_object_id,
metadata=litellm_metadata,
)
if response is None:
continue
verbose_proxy_logger.debug(
f"Response {unified_object_id} status: {response.status}, model: {model_name}"
)
except Exception as e:
verbose_proxy_logger.info(
f"Skipping job {unified_object_id} due to error: {e}"
@ -154,21 +202,20 @@ class CheckResponsesCost:
verbose_proxy_logger.info(
f"Response {unified_object_id} is complete. Cost automatically tracked by aget_responses."
)
completed_jobs.append(job)
elif response.status in ["failed", "cancelled"]:
status_to_job_ids["completed"].append(job.id)
elif response.status in ["failed", "cancelled", "expired"]:
verbose_proxy_logger.info(
f"Response {unified_object_id} has status {response.status}, marking as complete"
f"Response {unified_object_id} has status {response.status}, marking as {response.status}"
)
completed_jobs.append(job)
status_to_job_ids[response.status].append(job.id)
# Mark completed jobs in the database
if len(completed_jobs) > 0:
for status, job_ids in status_to_job_ids.items():
if not job_ids:
continue
await self.prisma_client.db.litellm_managedobjecttable.update_many(
where={"id": {"in": [job.id for job in completed_jobs]}},
data={"status": "completed"},
)
verbose_proxy_logger.info(
f"Marked {len(completed_jobs)} response jobs as completed"
where={"id": {"in": job_ids}},
data={"status": status},
)
verbose_proxy_logger.info(f"Marked {len(job_ids)} response jobs as {status}")

View file

@ -1490,6 +1490,15 @@ MANAGED_OBJECT_STALENESS_CUTOFF_DAYS = max(
STALE_OBJECT_CLEANUP_BATCH_SIZE = max(
1, int(os.getenv("STALE_OBJECT_CLEANUP_BATCH_SIZE", 1000))
)
STALE_OBJECT_CLEANUP_MAX_BATCHES_PER_POLL_CYCLE = max(
1, int(os.getenv("STALE_OBJECT_CLEANUP_MAX_BATCHES_PER_POLL_CYCLE", 10))
)
RESPONSES_COST_POLL_RETRY_ATTEMPTS = max(
1, int(os.getenv("RESPONSES_COST_POLL_RETRY_ATTEMPTS", 3))
)
RESPONSES_COST_POLL_RETRY_DELAY_SECONDS = max(
0.0, float(os.getenv("RESPONSES_COST_POLL_RETRY_DELAY_SECONDS", 0.2))
)
# Set PROXY_BATCH_POLLING_ENABLED=false to disable the CheckBatchCost and
# CheckResponsesCost background polling jobs entirely (e.g. to avoid DB load on
# installations with large numbers of stale managed objects).

View file

@ -189,12 +189,12 @@ class TestCheckResponsesCost:
await check_responses_cost_instance.check_responses_cost()
# update_many should only contain the job completion call
# update_many should only contain the failed-status call
calls = (
mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
)
assert len(calls) == 1
assert calls[0][1]["data"]["status"] == "completed"
assert calls[0][1]["data"]["status"] == "failed"
@pytest.mark.asyncio
async def test_check_responses_cost_with_cancelled_response(
@ -232,12 +232,12 @@ class TestCheckResponsesCost:
await check_responses_cost_instance.check_responses_cost()
# update_many should only contain the job completion call
# update_many should only contain the cancelled-status call
calls = (
mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
)
assert len(calls) == 1
assert calls[0][1]["data"]["status"] == "completed"
assert calls[0][1]["data"]["status"] == "cancelled"
@pytest.mark.asyncio
async def test_check_responses_cost_with_in_progress_response(
@ -364,6 +364,72 @@ class TestCheckResponsesCost:
# Stale cleanup still ran via _expire_stale_rows
check_responses_cost_instance._expire_stale_rows.assert_called_once()
@pytest.mark.asyncio
async def test_check_responses_cost_retries_transient_error_then_succeeds(
self, check_responses_cost_instance, mock_prisma_client
):
mock_job = MagicMock()
mock_job.unified_object_id = "resp_test_retry"
mock_job.created_by = "test-user"
mock_job.id = "job-retry"
mock_job.file_object = {"model": "gpt-4o", "id": "resp_test_retry"}
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_response = ResponsesAPIResponse(
id="resp_retry",
object="response",
status="completed",
created_at=int(datetime.now().timestamp()),
output=[],
usage=ResponseAPIUsage(
input_tokens=10,
output_tokens=10,
total_tokens=20,
),
)
with (
patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget,
patch(
"litellm_enterprise.proxy.common_utils.check_responses_cost.asyncio.sleep",
new_callable=AsyncMock,
),
):
mock_aget.side_effect = [Exception("temp failure"), mock_response]
await check_responses_cost_instance.check_responses_cost()
assert mock_aget.await_count == 2
calls = (
mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
)
assert len(calls) == 1
assert calls[0][1]["data"]["status"] == "completed"
@pytest.mark.asyncio
async def test_cleanup_stale_managed_objects_runs_multiple_batches(
self, check_responses_cost_instance, mock_prisma_client
):
from litellm.constants import STALE_OBJECT_CLEANUP_BATCH_SIZE
check_responses_cost_instance._expire_stale_rows = AsyncMock(
side_effect=[
STALE_OBJECT_CLEANUP_BATCH_SIZE,
3,
]
)
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[]
)
await check_responses_cost_instance.check_responses_cost()
assert check_responses_cost_instance._expire_stale_rows.await_count == 2
@pytest.mark.asyncio
async def test_check_responses_cost_multiple_jobs(
self, check_responses_cost_instance, mock_prisma_client

View file

@ -308,6 +308,7 @@ class TestCheckResponsesCost:
prisma_client=mock_prisma_client,
llm_router=mock_llm_router,
)
checker._expire_stale_rows = AsyncMock(return_value=0)
assert checker.proxy_logging_obj == mock_proxy_logging_obj
assert checker.prisma_client == mock_prisma_client
@ -332,12 +333,13 @@ class TestCheckResponsesCost:
prisma_client=mock_prisma_client,
llm_router=mock_llm_router,
)
checker._expire_stale_rows = AsyncMock(return_value=0)
# Should not raise any errors
await checker.check_responses_cost()
# Verify find_many was called with correct parameters (includes pagination)
mock_prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with(
mock_prisma_client.db.litellm_managedobjecttable.find_many.assert_any_call(
where={
"status": {"in": ["queued", "in_progress"]},
"file_purpose": "response",
@ -388,6 +390,7 @@ class TestCheckResponsesCost:
prisma_client=mock_prisma_client,
llm_router=mock_llm_router,
)
checker._expire_stale_rows = AsyncMock(return_value=0)
# Mock litellm.aget_responses to return completed response
with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget:
@ -444,14 +447,15 @@ class TestCheckResponsesCost:
prisma_client=mock_prisma_client,
llm_router=mock_llm_router,
)
checker._expire_stale_rows = AsyncMock(return_value=0)
with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget:
mock_aget.return_value = failed_response
await checker.check_responses_cost()
# Verify job was marked as completed even though it failed
# (stale cleanup also calls update_many, so check the specific completion call)
# Verify job was marked as failed
# (stale cleanup also calls update_many, so check the specific status call)
update_many_calls = (
mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
)
@ -461,6 +465,8 @@ class TestCheckResponsesCost:
if c.kwargs.get("where", {}).get("id") is not None
]
assert len(completion_calls) == 1
assert completion_calls[0].kwargs["where"]["id"]["in"] == ["job-456"]
assert completion_calls[0].kwargs["data"]["status"] == "failed"
@pytest.mark.asyncio
async def test_check_responses_cost_with_in_progress_job(
@ -497,6 +503,7 @@ class TestCheckResponsesCost:
prisma_client=mock_prisma_client,
llm_router=mock_llm_router,
)
checker._expire_stale_rows = AsyncMock(return_value=0)
with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget:
mock_aget.return_value = in_progress_response
@ -540,6 +547,7 @@ class TestCheckResponsesCost:
prisma_client=mock_prisma_client,
llm_router=mock_llm_router,
)
checker._expire_stale_rows = AsyncMock(return_value=0)
# Mock litellm.aget_responses to raise an exception
with patch(