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>
This commit is contained in:
yucheng 2026-09-23 13:12:25 +00:00
parent 6cbad5cbab
commit a760991a7a
5 changed files with 115 additions and 33 deletions

View file

@ -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,

View file

@ -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."""

View file

@ -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):

View file

@ -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

View file

@ -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