fix: count error-file failures in the batch cost poller path

This commit is contained in:
Mateo Wang 2026-08-29 12:06:42 -07:00
parent fee8619708
commit 4ef012627d
2 changed files with 149 additions and 3 deletions

View file

@ -2,9 +2,10 @@
Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if the cost has been tracked.
"""
from dataclasses import replace as dataclasses_replace
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Dict, Final, List, Optional, Tuple
from typing import TYPE_CHECKING, Any, Dict, Final, List, Literal, Optional, Tuple, cast
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
@ -626,6 +627,7 @@ class CheckBatchCost:
later poll.
"""
from litellm.batches.batch_utils import (
_count_error_file_failed_requests,
_get_file_content_as_dictionary,
calculate_batch_cost_and_usage,
)
@ -761,12 +763,31 @@ class CheckBatchCost:
model_id=model_id,
deployment_model=litellm_model_name,
)
batch_result = await calculate_batch_cost_and_usage(
batch_file_provider: Final = cast(
Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"], llm_provider
)
output_file_result: Final = await calculate_batch_cost_and_usage(
file_content_dictionary=file_content_as_dict,
custom_llm_provider=llm_provider, # type: ignore
custom_llm_provider=batch_file_provider,
model_name=model_name,
model_info=deployment_model_info,
)
error_file_failed_requests: Final = await _count_error_file_failed_requests(
response,
custom_llm_provider=batch_file_provider,
litellm_params={
**credentials,
"_litellm_internal_model_credentials": MappingProxyType(dict(credentials)),
},
)
batch_result: Final = (
output_file_result
if not error_file_failed_requests
else dataclasses_replace(
output_file_result,
failed_requests=output_file_result.failed_requests + error_file_failed_requests,
)
)
logging_obj = LiteLLMLogging(
model=batch_result.models[0],
messages=[{"role": "user", "content": "<retrieve_batch>"}],

View file

@ -1373,6 +1373,131 @@ class TestCheckBatchCost:
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_error_file_failures_add_to_failed_request_count(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
):
"""OpenAI-shaped providers report per-request failures only in a separate
error file. The poller prices from the output file, so without also counting
the error file's lines, batch_failed_requests on the spend log undercounts:
regression test for the poller path merging error-file failures.
"""
from unittest.mock import patch
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=1
)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=None
)
mock_job = MagicMock()
mock_job.id = "job-error-file-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = "user-1"
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_response = MagicMock()
mock_response.status = "completed"
mock_response.output_file_id = "file-output-123"
mock_response.error_file_id = "file-error-456"
mock_response.model_dump_json.return_value = (
'{"id":"batch-1","status":"completed"}'
)
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
return_value={"api_key": "sk-test"}
)
mock_deployment = MagicMock()
mock_deployment.litellm_params.custom_llm_provider = "openai"
mock_deployment.litellm_params.model = "gpt-4"
mock_deployment.model_info.model_dump.return_value = {}
mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment)
output_file_content = MagicMock()
output_file_content.content = b'{"id":"req-1"}'
error_file_content = MagicMock()
error_file_content.content = (
b'{"id":"err-1","error":{"message":"rejected"}}\n'
b'{"id":"err-2","error":{"message":"rejected"}}\n\n'
)
def _file_content_for(**kwargs):
if kwargs["file_id"] == "file-error-456":
return error_file_content
return output_file_content
decoded_id = "llm_model_id,model-123;llm_batch_id,batch-456;"
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
side_effect=[decoded_id, None, None],
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id",
return_value="model-123",
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id",
return_value="batch-456",
),
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
side_effect=_file_content_for,
) as mock_afile_content,
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary",
return_value=[{"id": "req-1"}],
),
patch(
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
new_callable=AsyncMock,
return_value=_batch_cost_result(
0.01,
{"prompt_tokens": 10, "completion_tokens": 5},
["gpt-4"],
successful_requests=3,
failed_requests=1,
),
),
patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("gpt-4", "openai", None, None),
),
patch(
"litellm.litellm_core_utils.litellm_logging.Logging"
) as mock_logging_cls,
):
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
mock_logging_cls.return_value = mock_logging_obj
await check_batch_cost_instance.check_batch_cost()
assert mock_afile_content.await_count == 2, (
"the poller must fetch the error file in addition to the output file"
)
fetched_file_ids = {
call.kwargs["file_id"] for call in mock_afile_content.await_args_list
}
assert fetched_file_ids == {"file-output-123", "file-error-456"}
mock_logging_obj.async_success_handler.assert_awaited_once()
handler_kwargs = mock_logging_obj.async_success_handler.await_args.kwargs
assert handler_kwargs["batch_successful_requests"] == 3
assert handler_kwargs["batch_failed_requests"] == 3, (
"2 error-file lines must add to the output file's 1 failed request"
)
assert handler_kwargs["batch_cost"] == 0.01
@pytest.mark.asyncio
async def test_terminal_batch_with_missing_output_file_is_retired_unbilled(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router