From d9e377f129d9b002ea6ed81b1e2d1be2d0824930 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 15 Aug 2026 12:56:53 -0700 Subject: [PATCH] fix(batches): confirm poller batch_processed support at startup so no retrieve accounts inline before the first poll cycle Probe the column before the scheduler registers CheckBatchCost, closing the window where a retrieve that decided the poller was inactive billed a batch the first poll cycle then billed again. Also drop narration docstrings and section banners from the new tests. --- .../proxy/common_utils/check_batch_cost.py | 28 +++++++++++- litellm/proxy/proxy_server.py | 1 + .../proxy_unit_tests/test_check_batch_cost.py | 43 +++++++++++++++++-- .../proxy/batches_endpoints/test_endpoints.py | 13 ------ .../test_files_common_utils.py | 31 ------------- tests/test_litellm/proxy/test_proxy_server.py | 38 ++++++++++++++++ 6 files changed, 105 insertions(+), 49 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 5f9f7a32a5b..7c88b78ad4e 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -53,6 +53,32 @@ class CheckBatchCost: self._has_batch_processed_column: bool = True self.batch_processed_support_confirmed: bool = False + @staticmethod + def _is_missing_batch_processed_column_error(err: Exception) -> bool: + message: Final = str(err).lower() + return "batch_processed" in message or "unknown column" in message or "does not exist" in message + + async def confirm_batch_processed_support(self) -> None: + """ + Probe the batch_processed column before the proxy serves traffic, so the retrieve + path never sees an unconfirmed poller on a schema that has the column and accounts + inline for a batch the first poll cycle then accounts again. + """ + try: + await self.prisma_client.db.litellm_managedobjecttable.find_first( + where={"file_purpose": "batch", "batch_processed": False} + ) + except Exception as probe_err: + if not self._is_missing_batch_processed_column_error(probe_err): + verbose_proxy_logger.debug( + f"CheckBatchCost: batch_processed probe failed, the poll cycle will confirm support: {probe_err}" + ) + return + self._has_batch_processed_column = False + verbose_proxy_logger.warning("CheckBatchCost: batch_processed column not found, querying without it") + return + self.batch_processed_support_confirmed = True + async def _get_user_info(self, batch_id: str, user_id: Optional[str]) -> Dict[str, Any]: """ Look up user email and key alias by user_id for enriching the S3 callback metadata. @@ -725,7 +751,7 @@ class CheckBatchCost: ) self.batch_processed_support_confirmed = True except Exception as query_err: - if "batch_processed" not in str(query_err).lower() and "unknown column" not in str(query_err).lower() and "does not exist" not in str(query_err).lower(): + if not self._is_missing_batch_processed_column_error(query_err): raise # Permanent schema gap — cache the result so future cycles skip straight to fallback self._has_batch_processed_column = False diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 10ee2e10a39..bda6fc25499 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -8986,6 +8986,7 @@ class ProxyStartupEvent: llm_router=llm_router, track_unmanaged_batch_cost=general_settings.get("track_unmanaged_batch_cost", False), ) + await check_batch_cost_job.confirm_batch_processed_support() scheduler.add_job( check_batch_cost_job.check_batch_cost, "interval", diff --git a/tests/proxy_unit_tests/test_check_batch_cost.py b/tests/proxy_unit_tests/test_check_batch_cost.py index 298701b767c..cf9c7c5c0c5 100644 --- a/tests/proxy_unit_tests/test_check_batch_cost.py +++ b/tests/proxy_unit_tests/test_check_batch_cost.py @@ -114,6 +114,45 @@ class TestCheckBatchCost: assert "stale_expired" in where["status"]["not_in"] assert "created_at" in where + @pytest.mark.asyncio + async def test_startup_probe_confirms_batch_processed_support( + self, check_batch_cost_instance, mock_prisma_client + ): + mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=None) + + await check_batch_cost_instance.confirm_batch_processed_support() + + probe_where = mock_prisma_client.db.litellm_managedobjecttable.find_first.call_args[1]["where"] + assert probe_where["batch_processed"] is False + assert check_batch_cost_instance.batch_processed_support_confirmed is True + assert check_batch_cost_instance._has_batch_processed_column is True + + @pytest.mark.asyncio + async def test_startup_probe_marks_column_absent( + self, check_batch_cost_instance, mock_prisma_client + ): + mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( + side_effect=Exception("column batch_processed does not exist") + ) + + await check_batch_cost_instance.confirm_batch_processed_support() + + assert check_batch_cost_instance.batch_processed_support_confirmed is False + assert check_batch_cost_instance._has_batch_processed_column is False + + @pytest.mark.asyncio + async def test_startup_probe_transient_error_defers_to_poll_cycle( + self, check_batch_cost_instance, mock_prisma_client + ): + mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( + side_effect=Exception("connection reset by peer") + ) + + await check_batch_cost_instance.confirm_batch_processed_support() + + assert check_batch_cost_instance.batch_processed_support_confirmed is False + assert check_batch_cost_instance._has_batch_processed_column is True + @pytest.mark.asyncio async def test_find_many_uses_pagination_and_excludes_stale( self, check_batch_cost_instance, mock_prisma_client @@ -143,10 +182,6 @@ class TestCheckBatchCost: assert "complete" not in not_in assert "completed" not in not_in assert find_call[1]["where"]["batch_processed"] is False - # A successful filtered query is the only proof the column exists. The retrieve - # path reads this to decide whether handing accounting to the poller is safe: - # without the column the poller's fallback query excludes complete/completed, so - # a batch already marked complete would never be accounted by anyone. assert check_batch_cost_instance.batch_processed_support_confirmed is True @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 6d00e56030f..80d52de0012 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -2408,18 +2408,10 @@ async def test_cancel__unified_batch_id_allowed_when_managed_files_required(canc assert cancel_harness.router_acancel.call_count == 1 -# =========================================================================== # -# Retrieve - who accounts for a managed batch's cost. Retrieving a batch and -# the CheckBatchCost poller both computed it, so whichever observed completion -# first won and the other either double counted or was locked out. -# =========================================================================== # @pytest.mark.asyncio async def test_retrieve__managed_batch_defers_cost_to_the_poller_when_it_is_running(retrieve_harness): - """With the poller running it is the single accountant, so the retrieve must not also - record cost. Without this the same batch is billed once per retrieve, and a caller - polling its own batch inflates spend by however many times it looked.""" with patch.object(endpoints, "batch_cost_poller_is_active", MagicMock(return_value=True)): await call_retrieve(retrieve_harness, _unified_batch_id()) @@ -2430,9 +2422,6 @@ async def test_retrieve__managed_batch_defers_cost_to_the_poller_when_it_is_runn @pytest.mark.asyncio async def test_retrieve__managed_batch_still_accounts_inline_without_a_poller(retrieve_harness): - """No poller means nothing else will ever account for this batch, so the retrieve has - to keep doing it. Suppressing here unconditionally would lose batch cost entirely on - any proxy running with batch polling disabled.""" with patch.object(endpoints, "batch_cost_poller_is_active", MagicMock(return_value=False)): await call_retrieve(retrieve_harness, _unified_batch_id()) @@ -2443,8 +2432,6 @@ async def test_retrieve__managed_batch_still_accounts_inline_without_a_poller(re @pytest.mark.asyncio async def test_retrieve__raw_batch_id_is_untouched_by_the_poller_handoff(retrieve_harness): - """An unmanaged batch has no managed object row and so no poller queue entry. It must - keep accounting inline whatever the poller is doing.""" with patch.object(endpoints, "batch_cost_poller_is_active", MagicMock(return_value=True)): await call_retrieve(retrieve_harness, "batch-raw-xyz") diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py index c16450446a5..590f74d946f 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py @@ -138,12 +138,6 @@ def _job_for(poller): ], ) def test_batch_cost_poller_is_active(monkeypatch, polling_enabled, job, expected): - """The predicate must only claim the poller when it can actually be relied on, so a - proxy with polling switched off, without the enterprise job, or whose poller has not - confirmed batch_processed support keeps accounting for batch cost on the retrieve - path. The unconfirmed case is the one that matters for legacy schemas: without the - column the poller falls back to a query excluding terminal statuses, so a batch the - retrieve path already marked complete would never be accounted by anyone.""" import litellm.constants import litellm.proxy.proxy_server as proxy_server_module from litellm.proxy.openai_files_endpoints.common_utils import ( @@ -212,11 +206,6 @@ async def _run_update(monkeypatch, poller_active: bool) -> dict: @pytest.mark.asyncio async def test_retrieving_a_completed_batch_leaves_batch_processed_to_the_cost_poller(monkeypatch): - """batch_processed is what removes a batch from CheckBatchCost's queue, which selects - batch_processed=False. Retrieving a batch records no cost when the poller is active, so - setting the flag here retired the poller on behalf of work nobody had done: a cost - callback that then failed lost the batch's cost permanently with no retry left. The - status update must still happen so callers see the terminal state.""" data = await _run_update(monkeypatch, poller_active=True) assert "batch_processed" not in data @@ -225,8 +214,6 @@ async def test_retrieving_a_completed_batch_leaves_batch_processed_to_the_cost_p @pytest.mark.asyncio async def test_retrieving_a_completed_batch_still_marks_processed_without_a_cost_poller(monkeypatch): - """With no poller to hand off to, this path is the only accountant, so it keeps setting - the flag. Otherwise a proxy with polling disabled would never unblock file deletion.""" data = await _run_update(monkeypatch, poller_active=False) assert data["batch_processed"] is True @@ -234,9 +221,6 @@ async def test_retrieving_a_completed_batch_still_marks_processed_without_a_cost def test_batch_cost_poller_is_active_is_false_when_the_job_has_no_bound_poller(monkeypatch): - """A scheduler that hands back a plain function rather than a bound method leaves no - poller to interrogate, so the predicate stays conservative instead of assuming the - column is supported.""" import litellm.constants import litellm.proxy.proxy_server as proxy_server_module from litellm.proxy.openai_files_endpoints.common_utils import ( @@ -256,8 +240,6 @@ def test_batch_cost_poller_is_active_is_false_when_the_job_has_no_bound_poller(m def test_batch_cost_poller_is_active_is_false_when_get_job_raises(monkeypatch): - """Scheduler backends raise varied types; an unreadable scheduler must not be read as - a working poller.""" import litellm.constants import litellm.proxy.proxy_server as proxy_server_module from litellm.proxy.openai_files_endpoints.common_utils import ( @@ -277,8 +259,6 @@ def test_batch_cost_poller_is_active_is_false_when_get_job_raises(monkeypatch): @pytest.mark.asyncio async def test_retrieving_a_batch_whose_status_is_unchanged_writes_nothing(monkeypatch): - """A caller polling an already-complete batch must not write at all, so repeated polls - cannot flip batch_processed or disturb whichever component owns accounting.""" import litellm.proxy.openai_files_endpoints.common_utils as cu monkeypatch.setattr(cu, "batch_cost_poller_is_active", lambda: False) @@ -307,8 +287,6 @@ async def test_retrieving_a_batch_whose_status_is_unchanged_writes_nothing(monke @pytest.mark.asyncio async def test_update_batch_in_database_is_a_noop_for_unmanaged_batches(monkeypatch): - """Batches with no managed object row have neither the flag nor a poller queue entry, so - this path must leave them alone entirely.""" import litellm.proxy.openai_files_endpoints.common_utils as cu prisma_client = MagicMock() @@ -330,14 +308,8 @@ async def test_update_batch_in_database_is_a_noop_for_unmanaged_batches(monkeypa @pytest.mark.asyncio async def test_the_caller_s_accounting_decision_wins_over_a_later_poller_transition(monkeypatch): - """The ownership decision is made before the provider retrieval and acted on there, so - re-deciding afterwards can observe a poller that only just became usable. That split - left the retrieve accounting inline while the row stayed unmarked, so the poller - accounted for the same batch again and billed it twice. Passing the decision through - makes both halves agree even when the poller transitions mid-flight.""" import litellm.proxy.openai_files_endpoints.common_utils as cu - # The predicate now reports an active poller, i.e. it flipped during the retrieval. monkeypatch.setattr(cu, "batch_cost_poller_is_active", lambda: True) monkeypatch.setattr(cu, "ensure_batch_response_managed_file_ids", AsyncMock()) @@ -366,9 +338,6 @@ async def test_the_caller_s_accounting_decision_wins_over_a_later_poller_transit @pytest.mark.asyncio async def test_a_caller_that_handed_off_accounting_still_leaves_the_marker_alone(monkeypatch): - """The mirror case: a caller that suppressed its own accounting must leave the marker - for the poller even if the predicate has since stopped reporting one, otherwise the - batch is retired without anyone having accounted for it.""" import litellm.proxy.openai_files_endpoints.common_utils as cu monkeypatch.setattr(cu, "batch_cost_poller_is_active", lambda: False) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 918d39646b0..ba207242e29 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -7070,6 +7070,44 @@ async def test_update_general_settings_store_model_in_db_none_keeps_current(): assert ps.store_model_in_db is False +@pytest.mark.asyncio +async def test_batch_cost_poller_is_confirmed_before_serving(monkeypatch): + monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False) + from litellm.proxy.openai_files_endpoints.common_utils import batch_cost_poller_is_active + from litellm.proxy.proxy_server import ProxyStartupEvent + from litellm.proxy.utils import ProxyLogging + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=None) + mock_proxy_logging = MagicMock(spec=ProxyLogging) + mock_proxy_logging.slack_alerting_instance = MagicMock() + mock_proxy_logging.db_spend_update_writer = MagicMock() + + with ( + patch("litellm.proxy.proxy_server.proxy_config", AsyncMock()), + patch("litellm.proxy.proxy_server.store_model_in_db", False), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.proxy.proxy_server.PROXY_BATCH_POLLING_ENABLED", True), + patch("litellm.constants.PROXY_BATCH_POLLING_ENABLED", True), + patch("litellm.proxy.proxy_server.get_secret_bool", return_value=False), + ): + await ProxyStartupEvent.initialize_scheduled_background_jobs( + general_settings={}, + prisma_client=mock_prisma_client, + proxy_budget_rescheduler_min_time=1, + proxy_budget_rescheduler_max_time=2, + proxy_batch_write_at=5, + proxy_logging_obj=mock_proxy_logging, + ) + + poller = proxy_server_module.scheduler.get_job("check_batch_cost_job").func.__self__ + assert poller.batch_processed_support_confirmed is True + assert batch_cost_poller_is_active() is True + probe_where = mock_prisma_client.db.litellm_managedobjecttable.find_first.call_args[1]["where"] + assert probe_where["batch_processed"] is False + + @pytest.mark.asyncio async def test_store_model_in_db_db_override_when_config_false(): """