fix(proxy): bill cancelled and failed batches that still produced an output file

This commit is contained in:
mateo-berri 2026-08-17 12:23:59 -07:00
parent badd737526
commit 21984101e5
2 changed files with 34 additions and 47 deletions

View file

@ -24,12 +24,16 @@ if TYPE_CHECKING:
CHECK_BATCH_COST_USER_AGENT = "LiteLLM Proxy/CheckBatchCost"
TERMINAL_MANAGED_OBJECT_STATUSES: Final[Tuple[str, ...]] = (
PROVIDER_TERMINAL_BATCH_STATUSES: Final[Tuple[str, ...]] = (
"completed",
"complete",
"failed",
"expired",
"cancelled",
)
TERMINAL_MANAGED_OBJECT_STATUSES: Final[Tuple[str, ...]] = (
*PROVIDER_TERMINAL_BATCH_STATUSES,
"stale_expired",
)
@ -796,7 +800,7 @@ class CheckBatchCost:
## RETRIEVE THE BATCH JOB OUTPUT FILE
if (
response.status in ("completed", "complete", "expired")
response.status in PROVIDER_TERMINAL_BATCH_STATUSES
and response.output_file_id is not None
):
try:
@ -837,13 +841,7 @@ class CheckBatchCost:
f"CheckBatchCost: failed to mark job {job.id} complete in DB: {db_err}"
)
elif response.status in (
"completed",
"complete",
"failed",
"expired",
"cancelled",
):
elif response.status in PROVIDER_TERMINAL_BATCH_STATUSES:
try:
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,

View file

@ -817,12 +817,12 @@ class TestCheckBatchCost:
mock_llm_router,
terminal_status,
):
"""A cancelled/failed batch with provider output files must be persisted with
unified managed file IDs, never raw provider IDs. Raw IDs written here leak
to every later GET /batches/{id} and GET /batches because the terminal row is
final (batch_processed=True) and read paths only resolve, never mint.
(Expired with an output file is billed through the completed path instead,
covered by test_expired_with_output_file_is_billed.)
"""A cancelled/failed batch with a provider error file (and no output file) must
be persisted with unified managed file IDs, never raw provider IDs. Raw IDs
written here leak to every later GET /batches/{id} and GET /batches because the
terminal row is final (batch_processed=True) and read paths only resolve, never
mint. (Any terminal status with an output file is billed through the completed
path instead, covered by test_terminal_status_with_output_file_is_billed.)
"""
import base64
import json
@ -832,15 +832,11 @@ class TestCheckBatchCost:
unified_batch_uid = base64.urlsafe_b64encode(
b"litellm_proxy;model_id:model-123;llm_batch_id:batch-456"
).decode()
raw_output_file_id = "file-terminal-out-abc"
raw_error_file_id = "file-terminal-err-xyz"
raw_input_file_id = "file-terminal-in-123"
unified_input_file_id = base64.urlsafe_b64encode(
b"litellm_proxy:application/octet-stream;unified_id,in-1;target_model_names,gpt-5-batch"
).decode()
unified_output_file_id = base64.urlsafe_b64encode(
f"litellm_proxy:application/octet-stream;unified_id,u-1;llm_output_file_id,{raw_output_file_id}".encode()
).decode()
unified_error_file_id = base64.urlsafe_b64encode(
f"litellm_proxy:application/octet-stream;unified_id,u-2;llm_output_file_id,{raw_error_file_id}".encode()
).decode()
@ -884,16 +880,13 @@ class TestCheckBatchCost:
input_file_id=raw_input_file_id,
object="batch",
status=terminal_status,
output_file_id=raw_output_file_id,
output_file_id=None,
error_file_id=raw_error_file_id,
)
mock_llm_router.aretrieve_batch = AsyncMock(return_value=response)
mock_hook = MagicMock()
mock_hook.get_unified_output_file_id.side_effect = [
unified_output_file_id,
unified_error_file_id,
]
mock_hook.get_unified_output_file_id.side_effect = [unified_error_file_id]
mock_hook.store_unified_file_id = AsyncMock()
check_batch_cost_instance.proxy_logging_obj.get_proxy_hook.return_value = (
mock_hook
@ -901,12 +894,7 @@ class TestCheckBatchCost:
await check_batch_cost_instance.check_batch_cost()
mock_hook.get_unified_output_file_id.assert_any_call(
output_file_id=raw_output_file_id,
model_id="model-123",
model_name="gpt-5-batch",
)
mock_hook.get_unified_output_file_id.assert_any_call(
mock_hook.get_unified_output_file_id.assert_called_once_with(
output_file_id=raw_error_file_id,
model_id="model-123",
model_name="gpt-5-batch",
@ -915,10 +903,7 @@ class TestCheckBatchCost:
next(iter(c.kwargs["model_mappings"].values())): c.kwargs["file_id"]
for c in mock_hook.store_unified_file_id.call_args_list
}
assert stored == {
raw_output_file_id: unified_output_file_id,
raw_error_file_id: unified_error_file_id,
}
assert stored == {raw_error_file_id: unified_error_file_id}
for store_call in mock_hook.store_unified_file_id.call_args_list:
assert store_call.kwargs["user_api_key_dict"].user_id == "user-1"
assert store_call.kwargs["user_api_key_dict"].team_id == "team-1"
@ -932,9 +917,8 @@ class TestCheckBatchCost:
persisted = json.loads(update_data["file_object"])
assert persisted["id"] == unified_batch_uid
assert persisted["input_file_id"] == unified_input_file_id
assert persisted["output_file_id"] == unified_output_file_id
assert persisted["output_file_id"] is None
assert persisted["error_file_id"] == unified_error_file_id
assert raw_output_file_id not in update_data["file_object"]
assert raw_error_file_id not in update_data["file_object"]
@pytest.mark.asyncio
@ -1067,12 +1051,17 @@ class TestCheckBatchCost:
), "a non-terminal batch must not be written back (would stop polling prematurely)"
@pytest.mark.asyncio
async def test_expired_with_output_file_is_billed(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
@pytest.mark.parametrize("terminal_status", ["expired", "cancelled", "failed"])
async def test_terminal_status_with_output_file_is_billed(
self,
check_batch_cost_instance,
mock_prisma_client,
mock_llm_router,
terminal_status,
):
"""An expired batch that still produced an output file served real request lines,
so it must be billed (cost tracked) and then marked processed, not silently
marked terminal without billing.
"""A terminal (expired/cancelled/failed) batch that still produced an output file
served real request lines, so it must be billed (cost tracked) and then marked
processed, not silently marked terminal without billing.
"""
from unittest.mock import patch
@ -1085,7 +1074,7 @@ class TestCheckBatchCost:
)
mock_job = MagicMock()
mock_job.id = "job-expired-with-output-1"
mock_job.id = "job-terminal-with-output-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = "user-1"
@ -1095,10 +1084,10 @@ class TestCheckBatchCost:
)
mock_response = MagicMock()
mock_response.status = "expired"
mock_response.status = terminal_status
mock_response.output_file_id = "file-output-123"
mock_response.model_dump_json.return_value = (
'{"id":"batch-1","status":"expired"}'
f'{{"id":"batch-1","status":"{terminal_status}"}}'
)
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
@ -1164,7 +1153,7 @@ class TestCheckBatchCost:
assert (
mock_afile_content.await_count == 1
), "expired batch with an output file must fetch results and be billed"
), f"{terminal_status} batch with an output file must fetch results and be billed"
mock_logging_obj.async_success_handler.assert_awaited_once()
assert (
mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1
@ -1174,8 +1163,8 @@ class TestCheckBatchCost:
]["data"]
assert update_data["batch_processed"] is True
assert (
update_data["status"] == "expired"
), "billed expired batch must keep its real terminal status in the DB"
update_data["status"] == terminal_status
), f"billed {terminal_status} batch must keep its real terminal status in the DB"
@pytest.mark.asyncio
async def test_raw_output_file_id_converted_to_managed_id(