From 74a780467bf4f7b8e1dc383b457399c509727570 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Tue, 12 May 2026 00:43:10 +0000 Subject: [PATCH] feat(proxy): enhance stale row expiration logic with terminal status guard - Updated `_expire_stale_rows` method to include a terminal status guard in the `update_many` query, preventing updates on rows with certain statuses. - Adjusted the proxy server to use a calculated polling interval for scheduling the response cost check job, improving scheduling accuracy. - Added a new unit test to verify the correct behavior of the updated `_expire_stale_rows` method, ensuring it respects the terminal status guard. --- .../common_utils/check_responses_cost.py | 14 ++++++- litellm/proxy/proxy_server.py | 8 ++-- .../test_check_responses_cost.py | 42 +++++++++++++++++++ 3 files changed, 60 insertions(+), 4 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 07fca882926..e9ccbe3b3d0 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -65,7 +65,19 @@ class CheckResponsesCost: return 0 await self.prisma_client.db.litellm_managedobjecttable.update_many( - where={"id": {"in": stale_ids}}, + where={ + "id": {"in": stale_ids}, + "status": { + "not_in": [ + "completed", + "complete", + "failed", + "expired", + "cancelled", + "stale_expired", + ] + }, + }, data={"status": "stale_expired"}, ) return len(stale_ids) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 164b0e90426..1117a512118 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -7188,11 +7188,13 @@ class ProxyStartupEvent: prisma_client=prisma_client, llm_router=llm_router, ) + _effective_responses_interval = ( + proxy_responses_polling_interval + random.randint(0, 30) + ) scheduler.add_job( check_responses_cost_job.check_responses_cost, "interval", - seconds=proxy_responses_polling_interval - + random.randint(0, 30), # Add small random offset + seconds=_effective_responses_interval, # REMOVED jitter parameter - major cause of memory leak id="check_responses_cost_job", replace_existing=True, @@ -7200,7 +7202,7 @@ class ProxyStartupEvent: ) verbose_proxy_logger.info( "Responses cost check job scheduled successfully (interval=%ss)", - proxy_responses_polling_interval, + _effective_responses_interval, ) except Exception as e: diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index b75185b774f..907ec64bcf5 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -104,6 +104,48 @@ class TestCheckResponsesCost: call_args = check_responses_cost_instance._expire_stale_rows.call_args assert call_args[0][1] == STALE_OBJECT_CLEANUP_BATCH_SIZE + @pytest.mark.asyncio + async def test_expire_stale_rows_rechecks_terminal_status_on_update( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """update_many should include terminal-status guard to avoid race clobber.""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + check_responses_cost_instance = CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[MagicMock(id="row-1"), MagicMock(id="row-2")] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=0 + ) + + marked = await check_responses_cost_instance._expire_stale_rows( + cutoff=datetime.now(), batch_size=10 + ) + + assert marked == 2 + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_awaited_once() + where = ( + mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args[1][ + "where" + ] + ) + assert where["id"]["in"] == ["row-1", "row-2"] + assert where["status"]["not_in"] == [ + "completed", + "complete", + "failed", + "expired", + "cancelled", + "stale_expired", + ] + @pytest.mark.asyncio async def test_check_responses_cost_with_completed_response( self, check_responses_cost_instance, mock_prisma_client, mock_llm_router