fix: avoid duplicate cost logging in fallback path; guard integer constants against zero/negative values

This commit is contained in:
Ishaan Jaffer 2026-03-12 18:00:55 -07:00
parent 52c0574cc5
commit a14b951248
3 changed files with 62 additions and 9 deletions

View file

@ -97,7 +97,11 @@ class CheckBatchCost:
await self._cleanup_stale_managed_objects()
# Look for all batches that have not yet been processed by CheckBatchCost
# Look for all batches that have not yet been processed by CheckBatchCost.
# _has_batch_processed_column tracks whether the column exists so the
# completion update can omit it on older schemas (avoiding a silent failure
# that would cause infinite reprocessing and duplicate cost logging).
_has_batch_processed_column = True
try:
jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many(
where={
@ -110,6 +114,7 @@ class CheckBatchCost:
)
except Exception:
# Fallback: batch_processed column may not exist on older schemas
_has_batch_processed_column = False
verbose_proxy_logger.warning(
"CheckBatchCost: batch_processed column not found, querying without it"
)
@ -288,13 +293,15 @@ class CheckBatchCost:
# mark the job as complete
try:
update_data: dict = {
"status": "complete",
"file_object": response.model_dump_json(),
}
if _has_batch_processed_column:
update_data["batch_processed"] = True
await self.prisma_client.db.litellm_managedobjecttable.update(
where={"id": job.id},
data={
"batch_processed": True,
"status": "complete",
"file_object": response.model_dump_json(),
},
data=update_data,
)
except Exception as db_err:
verbose_proxy_logger.error(

View file

@ -1351,9 +1351,9 @@ PROXY_BUDGET_RESCHEDULER_MIN_TIME = int(
os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597)
)
PROXY_BATCH_POLLING_INTERVAL = int(os.getenv("PROXY_BATCH_POLLING_INTERVAL", 3600))
MAX_OBJECTS_PER_POLL_CYCLE = int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50))
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS = int(
os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7)
MAX_OBJECTS_PER_POLL_CYCLE = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50)))
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS = max(
1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7))
)
# Set PROXY_BATCH_POLLING_ENABLED=false to disable the CheckBatchCost and
# CheckResponsesCost background polling jobs entirely (e.g. to avoid DB load on

View file

@ -110,3 +110,49 @@ class TestCheckBatchCost:
assert "stale_expired" in fallback_where["status"]["not_in"]
# Fallback must still paginate
assert calls[1][1]["take"] == MAX_OBJECTS_PER_POLL_CYCLE
@pytest.mark.asyncio
async def test_fallback_completion_update_omits_batch_processed(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
):
"""When batch_processed column is absent, completion update must not include it.
If it did, the update would fail silently, the job would never be marked done,
and every subsequent poll cycle would re-log the cost (duplicate billing).
"""
from unittest.mock import patch
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_job = MagicMock()
mock_job.id = "job-fallback-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA==" # base64-looking value
mock_job.created_by = "user-1"
# Primary query fails → fallback path
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
side_effect=[Exception("column batch_processed does not exist"), [mock_job]]
)
# Stub out the heavy per-job processing so we reach the update()
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
return_value=None, # causes "not a valid unified object id" early-continue
),
):
await check_batch_cost_instance.check_batch_cost()
# Even though the job was skipped (invalid ID), confirm the fallback path was taken
# by checking the find_many calls
find_calls = mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args_list
assert len(find_calls) == 2
fallback_where = find_calls[1][1]["where"]
assert "batch_processed" not in fallback_where
# If a completion update were issued, it must not contain batch_processed
for call in mock_prisma_client.db.litellm_managedobjecttable.update.call_args_list:
assert "batch_processed" not in call[1].get("data", {})