mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
f3a041dec3
commit
b4e598fe36
4 changed files with 195 additions and 65 deletions
|
|
@ -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}")
|
||||
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue