From b4e598fe364c4f618c4f20880d132124e0acb74d Mon Sep 17 00:00:00 2001 From: harish-berri Date: Mon, 11 May 2026 22:45:52 +0000 Subject: [PATCH] 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. --- .../common_utils/check_responses_cost.py | 163 +++++++++++------- litellm/constants.py | 9 + .../test_check_responses_cost.py | 74 +++++++- .../test_responses_background_cost.py | 14 +- 4 files changed, 195 insertions(+), 65 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py index dc0168683c8..07fca882926 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -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}") diff --git a/litellm/constants.py b/litellm/constants.py index 3a42ced3def..02c68095599 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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). diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index 4c0ca94df48..b75185b774f 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -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 diff --git a/tests/test_litellm/integrations/test_responses_background_cost.py b/tests/test_litellm/integrations/test_responses_background_cost.py index 0d4218f2137..6ddd7251853 100644 --- a/tests/test_litellm/integrations/test_responses_background_cost.py +++ b/tests/test_litellm/integrations/test_responses_background_cost.py @@ -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(