mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: avoid duplicate cost logging in fallback path; guard integer constants against zero/negative values
This commit is contained in:
parent
52c0574cc5
commit
a14b951248
3 changed files with 62 additions and 9 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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", {})
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue