fix(batches): skip line-item callbacks for providers without a line reconstruction

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-23 09:00:38 +00:00
parent c746e5720b
commit 45b7509692
2 changed files with 39 additions and 14 deletions

View file

@ -3,7 +3,7 @@ import uuid
from collections.abc import Iterator, Mapping
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal, TypeAlias
from typing import TYPE_CHECKING, Final, Literal, TypeAlias, cast, get_args
from litellm._logging import verbose_logger
from litellm.batches.batch_utils import (
@ -27,6 +27,15 @@ if TYPE_CHECKING:
_BatchLineProvider: TypeAlias = Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"]
_SUPPORTED_LINE_PROVIDERS: Final = frozenset(get_args(_BatchLineProvider))
def _supported_line_provider(value: str) -> _BatchLineProvider | None:
if value in _SUPPORTED_LINE_PROVIDERS:
return cast("_BatchLineProvider", value) # cast-ok: membership in the literal's args was just checked
return None
_CALL_TYPE_BY_BATCH_URL: Final = MappingProxyType(
{
"/v1/chat/completions": "acompletion",
@ -291,7 +300,7 @@ async def _fetch_managed_file_or_empty(
async def log_batch_line_items(
batch: LiteLLMBatch,
custom_llm_provider: _BatchLineProvider,
custom_llm_provider: str,
parent: "Logging",
model_name: str | None,
litellm_params: dict[str, object] | None, # mutable-ok: the logging object's shared litellm_params dict
@ -303,6 +312,14 @@ async def log_batch_line_items(
aretrieve_batch event still bills the batch, so per-line events carry
``batch_parent_id`` and never update spend themselves. Any failure here
is logged and swallowed: aggregate accounting must be unaffected."""
line_provider: Final = _supported_line_provider(custom_llm_provider)
if line_provider is None:
verbose_logger.warning(
"batch line-item callbacks are not supported for provider %s, skipping. batch_id=%s",
custom_llm_provider,
batch.id,
)
return 0
emitted = 0 # rebind-ok: loop accumulator for emitted line count
try:
internal_credentials: Final = parent._litellm_internal_model_credentials # pyright: ignore[reportPrivateUsage] # declared transport attribute on Logging
@ -313,17 +330,11 @@ async def log_batch_line_items(
else litellm_params
)
input_file_content: Final = await _fetch_managed_file_or_empty(
batch.input_file_id, custom_llm_provider, fetch_params
)
input_file_content: Final = await _fetch_managed_file_or_empty(batch.input_file_id, line_provider, fetch_params)
requests_by_id: Final = _requests_by_custom_id(input_file_content)
output_content: Final = await _fetch_managed_file_or_empty(
batch.output_file_id, custom_llm_provider, fetch_params
)
error_content: Final = await _fetch_managed_file_or_empty(
batch.error_file_id, custom_llm_provider, fetch_params
)
output_content: Final = await _fetch_managed_file_or_empty(batch.output_file_id, line_provider, fetch_params)
error_content: Final = await _fetch_managed_file_or_empty(batch.error_file_id, line_provider, fetch_params)
for content in (output_content, error_content):
for entry in _output_entries(content):
try:
@ -331,7 +342,7 @@ async def log_batch_line_items(
entry=entry,
requests_by_id=requests_by_id,
batch=batch,
custom_llm_provider=custom_llm_provider,
custom_llm_provider=line_provider,
parent=parent,
model_name=model_name,
model_info=model_info,

View file

@ -134,7 +134,7 @@ def recorder():
litellm._async_failure_callback = saved_failure # test-quality-ok: teardown restoring the value set above
def _parent_logging() -> Logging:
def _parent_logging(custom_llm_provider: str = "openai") -> Logging:
logging_obj = Logging(
model="gpt-4o",
messages=[{"role": "user", "content": "<retrieve_batch>"}],
@ -147,7 +147,7 @@ def _parent_logging() -> Logging:
logging_obj.update_environment_variables(
litellm_params={"metadata": {"model_info": {"id": "dep-1"}, "model_group": "gpt-4o"}},
optional_params={},
custom_llm_provider="openai",
custom_llm_provider=custom_llm_provider,
)
return logging_obj
@ -223,6 +223,20 @@ async def test_flag_off_emits_only_aggregate(recorder):
file_mock.assert_not_called()
@pytest.mark.asyncio
async def test_line_items_skipped_for_unsupported_provider(recorder):
litellm.store_batch_line_items_in_callbacks = True # test-quality-ok: the flag under test is a module global; fixture restores it
file_mock: Final = AsyncMock(side_effect=_file_content)
with patch("litellm.files.main.afile_content", file_mock): # test-quality-ok: afile_content is the provider boundary; no injection seam for managed file fetch
await _log_completed_batch(_parent_logging(custom_llm_provider="bedrock"), _batch())
file_mock.assert_not_awaited()
assert len(recorder.success_events) == 1
assert _payload(recorder.success_events[0])["response_cost"] == 1.5
assert "batch_custom_id" not in _hidden(recorder.success_events[0])
assert len(recorder.failure_events) == 0
@pytest.mark.asyncio
async def test_in_progress_batch_poll_emits_no_line_events(recorder):
litellm.store_batch_line_items_in_callbacks = True # test-quality-ok: the flag under test is a module global; fixture restores it