From e7f8cd50d04164a3bfb359436b5e8ba852aba932 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Thu, 12 Mar 2026 13:24:40 -0700 Subject: [PATCH] fix: exclude stale_expired from batch poll queries; fix update_many assertions in tests --- .../proxy/common_utils/check_batch_cost.py | 13 ++- .../test_check_responses_cost.py | 93 +++++++++++-------- 2 files changed, 66 insertions(+), 40 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 605f0cc735a..fe3e8f1402c 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -102,7 +102,7 @@ class CheckBatchCost: where={ "file_purpose": "batch", "batch_processed": False, - "status": {"not_in": ["failed", "expired", "cancelled"]}, + "status": {"not_in": ["failed", "expired", "cancelled", "stale_expired"]}, }, take=MAX_OBJECTS_PER_POLL_CYCLE, order={"created_at": "asc"}, @@ -115,7 +115,16 @@ class CheckBatchCost: jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many( where={ "file_purpose": "batch", - "status": {"not_in": ["failed", "expired", "cancelled", "complete", "completed"]}, + "status": { + "not_in": [ + "failed", + "expired", + "cancelled", + "complete", + "completed", + "stale_expired", + ] + }, }, take=MAX_OBJECTS_PER_POLL_CYCLE, order={"created_at": "asc"}, diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index bbffd6e0e75..d542b8c2181 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -133,8 +133,9 @@ class TestCheckResponsesCost: ), ) - # Mock update_many - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=0 + ) # Run the check with mocked litellm.aget_responses with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: @@ -142,11 +143,12 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - # Verify the job was marked as completed - mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() - call_args = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args - assert call_args[1]["data"]["status"] == "completed" - assert call_args[1]["where"]["id"]["in"] == ["job-123"] + # calls[0] = stale cleanup, calls[1] = job completion + calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list + assert len(calls) == 2 + completion_call = calls[1] + assert completion_call[1]["data"]["status"] == "completed" + assert completion_call[1]["where"]["id"]["in"] == ["job-123"] @pytest.mark.asyncio async def test_check_responses_cost_with_failed_response( @@ -173,8 +175,9 @@ class TestCheckResponsesCost: usage=None, ) - # Mock update_many - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=0 + ) # Run the check with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: @@ -182,10 +185,10 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - # Verify the job was marked as completed (even though response failed) - mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() - call_args = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args - assert call_args[1]["data"]["status"] == "completed" + # calls[0] = stale cleanup, calls[1] = job completion + calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list + assert len(calls) == 2 + assert calls[1][1]["data"]["status"] == "completed" @pytest.mark.asyncio async def test_check_responses_cost_with_cancelled_response( @@ -212,8 +215,9 @@ class TestCheckResponsesCost: usage=None, ) - # Mock update_many - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=0 + ) # Run the check with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: @@ -221,8 +225,10 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - # Verify the job was marked as completed - mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() + # calls[0] = stale cleanup, calls[1] = job completion + calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list + assert len(calls) == 2 + assert calls[1][1]["data"]["status"] == "completed" @pytest.mark.asyncio async def test_check_responses_cost_with_in_progress_response( @@ -249,8 +255,9 @@ class TestCheckResponsesCost: usage=None, ) - # Mock update_many - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=0 + ) # Run the check with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: @@ -258,8 +265,10 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - # Verify no updates were made (response still in progress) - mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_not_called() + # Only the stale-cleanup call should have fired — no completion update + calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list + assert len(calls) == 1 + assert calls[0][1]["data"] == {"status": "stale_expired"} @pytest.mark.asyncio async def test_check_responses_cost_with_queued_response( @@ -286,8 +295,9 @@ class TestCheckResponsesCost: usage=None, ) - # Mock update_many - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=0 + ) # Run the check with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: @@ -295,8 +305,10 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - # Verify no updates were made (response still queued) - mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_not_called() + # Only the stale-cleanup call should have fired — no completion update + calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list + assert len(calls) == 1 + assert calls[0][1]["data"] == {"status": "stale_expired"} @pytest.mark.asyncio async def test_check_responses_cost_with_exception( @@ -313,8 +325,9 @@ class TestCheckResponsesCost: return_value=[mock_job] ) - # Mock update_many - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=0 + ) # Run the check with mocked exception with patch( @@ -325,8 +338,10 @@ class TestCheckResponsesCost: # Should not raise, just skip the job await check_responses_cost_instance.check_responses_cost() - # Verify no updates were made (job was skipped due to error) - mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_not_called() + # Only the stale-cleanup call should have fired — no completion update + calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list + assert len(calls) == 1 + assert calls[0][1]["data"] == {"status": "stale_expired"} @pytest.mark.asyncio async def test_check_responses_cost_multiple_jobs( @@ -389,8 +404,9 @@ class TestCheckResponsesCost: ), ) - # Mock update_many - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=0 + ) # Run the check with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: @@ -398,10 +414,11 @@ class TestCheckResponsesCost: await check_responses_cost_instance.check_responses_cost() - # Verify only the 2 completed jobs were marked as complete - mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() - call_args = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args - assert len(call_args[1]["where"]["id"]["in"]) == 2 - assert "job-1" in call_args[1]["where"]["id"]["in"] - assert "job-3" in call_args[1]["where"]["id"]["in"] - assert "job-2" not in call_args[1]["where"]["id"]["in"] + # calls[0] = stale cleanup, calls[1] = completion of 2 finished jobs + calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list + assert len(calls) == 2 + completion_call = calls[1] + assert len(completion_call[1]["where"]["id"]["in"]) == 2 + assert "job-1" in completion_call[1]["where"]["id"]["in"] + assert "job-3" in completion_call[1]["where"]["id"]["in"] + assert "job-2" not in completion_call[1]["where"]["id"]["in"]