diff --git a/litellm/batches/batch_line_item_logging.py b/litellm/batches/batch_line_item_logging.py index b91ef3dccd3..a414c419710 100644 --- a/litellm/batches/batch_line_item_logging.py +++ b/litellm/batches/batch_line_item_logging.py @@ -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, diff --git a/tests/test_litellm/batches/test_batch_line_item_logging.py b/tests/test_litellm/batches/test_batch_line_item_logging.py index 62c6bd359f2..1c59eac24bb 100644 --- a/tests/test_litellm/batches/test_batch_line_item_logging.py +++ b/tests/test_litellm/batches/test_batch_line_item_logging.py @@ -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": ""}], @@ -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