From a760991a7a7041c68775008168af847589e629d3 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 23 Sep 2026 13:12:25 +0000 Subject: [PATCH] fix(batches): pair bedrock lines by recordId and type them from the reconstructed response Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/batches/batch_line_item_logging.py | 91 +++++++++++++------ .../llms/base_llm/batches/transformation.py | 11 ++- .../llms/bedrock/batches/transformation.py | 5 + .../batches/test_batch_line_item_logging.py | 17 +++- .../bedrock/batches/test_transformation.py | 24 ++++- 5 files changed, 115 insertions(+), 33 deletions(-) diff --git a/litellm/batches/batch_line_item_logging.py b/litellm/batches/batch_line_item_logging.py index 02dd9cceaeb..3d54625d3fe 100644 --- a/litellm/batches/batch_line_item_logging.py +++ b/litellm/batches/batch_line_item_logging.py @@ -71,15 +71,16 @@ def _output_entries(file_content: bytes) -> Iterator[Mapping[str, object]]: yield mapping +def _line_id(entry: Mapping[str, object]) -> str | None: + line_id: Final = entry.get("custom_id") or entry.get("recordId") + return line_id if isinstance(line_id, str) and line_id else None + + def _requests_by_custom_id(input_file_content: bytes) -> Mapping[str, Mapping[str, object]]: - """Parse the batch input JSONL into {custom_id: request line}, skipping - malformed lines and lines without a custom_id.""" + """Parse the batch input JSONL into {line id: request line}, keyed by + custom_id or recordId, skipping malformed lines and lines without either.""" return MappingProxyType( - { - custom_id: entry - for entry in _output_entries(input_file_content) - if isinstance((custom_id := entry.get("custom_id")), str) and custom_id - } + {line_id: entry for entry in _output_entries(input_file_content) if (line_id := _line_id(entry)) is not None} ) @@ -93,6 +94,9 @@ def _request_body_for_entry( params: Final = _as_object_mapping(request_line.get("params")) if params: return params + request_model_input: Final = _as_object_mapping(request_line.get("modelInput")) + if request_model_input: + return request_model_input model_input: Final = _as_object_mapping(entry.get("modelInput")) return model_input if model_input else _EMPTY_BODY @@ -114,6 +118,16 @@ def _call_type_for_request(request_line: Mapping[str, object] | None) -> str: return _CALL_TYPE_BY_BATCH_URL.get(url if isinstance(url, str) else "", "acompletion") +def _call_type_for_line(request_line: Mapping[str, object] | None, result: "_BatchLineResult | None") -> str: + if isinstance(result, EmbeddingResponse): + return "aembedding" + if isinstance(result, ResponsesAPIResponse): + return "aresponses" + if isinstance(result, ModelResponse): + return "acompletion" + return _call_type_for_request(request_line) + + def _line_messages(request_body: Mapping[str, object]) -> object: return request_body.get("messages") or request_body.get("input") or () @@ -138,13 +152,16 @@ def _line_result( model: str, response_body: Mapping[str, object], ) -> _BatchLineResult: - if custom_llm_provider == "bedrock": - from litellm.llms.bedrock.batches.transformation import bedrock_batch_line_to_response + from litellm.types.utils import LlmProviders + from litellm.utils import ProviderConfigManager - bedrock_result: Final = bedrock_batch_line_to_response(response_body, model) - if bedrock_result is None: + provider_config: Final = ProviderConfigManager.get_provider_batches_config(model, LlmProviders(custom_llm_provider)) + if provider_config is not None: + transformed: Final = provider_config.transform_batch_output_line(response_body, model) + if transformed is not None: + return transformed + if custom_llm_provider == "bedrock": raise ValueError(f"unrecognized bedrock batch output line shape. keys={sorted(response_body)}") - return bedrock_result if call_type == "aembedding": return EmbeddingResponse(**response_body) # pyright: ignore[reportArgumentType] # provider output bodies are dicts expanded as response ctor kwargs if call_type == "aresponses": @@ -156,6 +173,24 @@ def _line_result( return ModelResponse(**response_body) # pyright: ignore[reportArgumentType] # same as above +def _line_result_or_none( + request_call_type: str, + custom_llm_provider: _BatchLineProvider, + model: str, + response_body: Mapping[str, object], + custom_id: object, +) -> "_BatchLineResult | None": + try: + return _line_result(request_call_type, custom_llm_provider, model, response_body) + except Exception: # noqa: BLE001 # one unparseable line must not drop the rest of the batch's line events + verbose_logger.warning( + "batch output line could not be reconstructed as a %s response, skipping it. custom_id=%s", + request_call_type, + custom_id, + ) + return None + + def _new_child_logging( parent: "Logging", model: str, @@ -218,17 +253,28 @@ async def _emit_line_event( model_name: str | None, model_info: ModelInfo | None, ) -> bool: - custom_id: Final = entry.get("custom_id") or entry.get("recordId") - request_line: Final = requests_by_id.get(custom_id if isinstance(custom_id, str) else "") + custom_id: Final = _line_id(entry) + request_line: Final = requests_by_id.get(custom_id or "") request_body: Final = _request_body_for_entry(entry, request_line) status_code: Final = _line_status_code(entry, custom_llm_provider) - call_type: Final = _call_type_for_request(request_line) + request_call_type: Final = _call_type_for_request(request_line) response_body: Final = _get_response_from_batch_job_output_file(entry, custom_llm_provider) parent_start_time: Final = parent.start_time # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # Logging.start_time is untyped upstream start_time: Final = parent_start_time if isinstance(parent_start_time, datetime) else datetime.now() # noqa: DTZ005 # naive to match the logging pipeline start_time parent_params: Final = _as_object_mapping(parent.litellm_params) or _EMPTY_BODY # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # Logging.litellm_params is untyped upstream - model: Final = _line_model(response_body, request_body, parent) + + successful: Final = _batch_response_was_successful(entry, custom_llm_provider) + stats: Final = _safe_output_line_stats(entry, custom_llm_provider, model_name, model_info) if successful else None + result: Final[_BatchLineResult | None] = ( + _line_result_or_none(request_call_type, custom_llm_provider, model, response_body, custom_id) + if successful + else None + ) + if successful and result is None: + return False + + call_type: Final = _call_type_for_line(request_line, result) child: Final = _new_child_logging( parent=parent, model=model, @@ -248,7 +294,7 @@ async def _emit_line_event( ) now: Final = datetime.now() # noqa: DTZ005 # naive to match the logging pipeline start_time - if not _batch_response_was_successful(entry, custom_llm_provider): + if result is None: exception: Final = _BatchLineFailure( entry.get("error") or entry.get("response") or {} # mutable-ok: fallback payload dict passed to Exception ) @@ -261,17 +307,6 @@ async def _emit_line_event( ) return True - stats: Final = _safe_output_line_stats(entry, custom_llm_provider, model_name, model_info) - try: - result: Final = _line_result(call_type, custom_llm_provider, model, response_body) - except Exception: # noqa: BLE001 # one unparseable line must not drop the rest of the batch's line events - verbose_logger.warning( - "batch output line could not be reconstructed as a %s response, skipping it. custom_id=%s", - call_type, - custom_id, - ) - return False - result._hidden_params = _line_hidden_params( # pyright: ignore[reportPrivateUsage] # same hidden_params channel the aggregate batch event uses batch, custom_id, diff --git a/litellm/llms/base_llm/batches/transformation.py b/litellm/llms/base_llm/batches/transformation.py index 34c622d4cf6..904eae2f4f2 100644 --- a/litellm/llms/base_llm/batches/transformation.py +++ b/litellm/llms/base_llm/batches/transformation.py @@ -1,5 +1,6 @@ import types from abc import ABC, abstractmethod +from collections.abc import Mapping from typing import TYPE_CHECKING, Any import httpx @@ -9,7 +10,7 @@ from litellm.types.llms.openai import ( AllMessageValues, CreateBatchRequest, ) -from litellm.types.utils import LiteLLMBatch, LlmProviders +from litellm.types.utils import EmbeddingResponse, LiteLLMBatch, LlmProviders, ModelResponse if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj @@ -39,6 +40,14 @@ class BaseBatchesConfig(ABC): def custom_llm_provider(self) -> LlmProviders: """Return the LLM provider type for this configuration.""" + def transform_batch_output_line( + self, model_output: Mapping[str, object], model: str + ) -> ModelResponse | EmbeddingResponse | None: + """Reconstruct one provider batch output line into a litellm response, or + None when the line is OpenAI-shaped and the caller can use the generic + reconstruction.""" + return None + @classmethod def get_config(cls): """Get configuration dictionary for this class.""" diff --git a/litellm/llms/bedrock/batches/transformation.py b/litellm/llms/bedrock/batches/transformation.py index 6d43e9ee7f9..f8d71bab2c7 100644 --- a/litellm/llms/bedrock/batches/transformation.py +++ b/litellm/llms/bedrock/batches/transformation.py @@ -144,6 +144,11 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): def custom_llm_provider(self) -> LlmProviders: return LlmProviders.BEDROCK + def transform_batch_output_line( + self, model_output: Mapping[str, object], model: str + ) -> ModelResponse | EmbeddingResponse | None: + return bedrock_batch_line_to_response(model_output, model) + @classmethod def _get_bare_model_name_from_s3_key(cls, object_key: str) -> str | None: if not object_key.startswith(BEDROCK_MANAGED_S3_BATCH_PREFIX): 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 0be7f76b594..1f37969bcfb 100644 --- a/tests/test_litellm/batches/test_batch_line_item_logging.py +++ b/tests/test_litellm/batches/test_batch_line_item_logging.py @@ -625,16 +625,27 @@ async def test_line_items_bedrock_shapes(recorder): by_id = {_hidden(e).get("batch_custom_id"): e for e in recorder.success_events} - anthropic_line = _payload(by_id["br-anth"]) + anthropic_event = by_id["br-anth"] + assert anthropic_event["call_type"] == "acompletion" + assert _hidden(anthropic_event)["batch_custom_id"] == "br-anth" + anthropic_line = _payload(anthropic_event) + assert any(m.get("content") == "hi br-anth" for m in anthropic_line["messages"]) assert anthropic_line["response"]["choices"][0]["message"]["content"] == "bedrock claude hi" assert anthropic_line["prompt_tokens"] == 2 assert anthropic_line["completion_tokens"] == 3 - titan_line = _payload(by_id["br-titan"]) + titan_event = by_id["br-titan"] + assert titan_event["call_type"] == "aembedding" + assert _hidden(titan_event)["batch_custom_id"] == "br-titan" + titan_line = _payload(titan_event) + assert titan_event["optional_params"]["inputText"] == "embed me" assert titan_line["response"]["data"][0]["embedding"] == [0.1, 0.2] assert titan_line["prompt_tokens"] == 4 - converse_line = _payload(by_id["br-conv"]) + converse_event = by_id["br-conv"] + assert converse_event["call_type"] == "acompletion" + converse_line = _payload(converse_event) + assert any(m.get("content") == "hi br-conv" for m in converse_line["messages"]) assert converse_line["response"]["choices"][0]["message"]["content"] == "converse hi" assert converse_line["prompt_tokens"] == 5 assert converse_line["completion_tokens"] == 6 diff --git a/tests/test_litellm/llms/bedrock/batches/test_transformation.py b/tests/test_litellm/llms/bedrock/batches/test_transformation.py index 347c459a369..4bce3c37504 100644 --- a/tests/test_litellm/llms/bedrock/batches/test_transformation.py +++ b/tests/test_litellm/llms/bedrock/batches/test_transformation.py @@ -20,7 +20,7 @@ import httpx import pytest from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig -from litellm.types.utils import LlmProviders +from litellm.types.utils import EmbeddingResponse, LlmProviders, ModelResponse # AWS JobStatus -> OpenAI BatchJobStatus, exactly as encoded in transformation.py # (both transform_create_batch_response and transform_retrieve_batch_response). @@ -941,3 +941,25 @@ def test_retrieve_request_accepts_partition_arns(config: BedrockBatchesConfig, a batch_id=arn, optional_params={}, litellm_params={} ) assert result["url"].startswith(expected_prefix) + + +def test_transform_batch_output_line_dispatches_on_shape(config: BedrockBatchesConfig) -> None: + titan = config.transform_batch_output_line( + {"embedding": [0.1, 0.2], "inputTextTokenCount": 4}, model="amazon.titan-embed" + ) + assert isinstance(titan, EmbeddingResponse) + assert titan.data[0]["embedding"] == [0.1, 0.2] + assert titan.usage.prompt_tokens == 4 + + converse = config.transform_batch_output_line( + { + "output": {"message": {"role": "assistant", "content": [{"text": "hi"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 5, "outputTokens": 6, "totalTokens": 11}, + }, + model="amazon.nova-lite", + ) + assert isinstance(converse, ModelResponse) + assert converse.choices[0].message.content == "hi" + + assert config.transform_batch_output_line({"foo": 1}, model="x") is None