mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
c746e5720b
commit
45b7509692
2 changed files with 39 additions and 14 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue