fix(batches): reconstruct bedrock and mistral batch lines through their provider transformations

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-23 10:40:51 +00:00
parent 45b7509692
commit bef2f0489f
6 changed files with 229 additions and 30 deletions

View file

@ -25,7 +25,9 @@ from litellm.types.utils import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging
_BatchLineProvider: TypeAlias = Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"]
_BatchLineProvider: TypeAlias = Literal[
"openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"
]
_SUPPORTED_LINE_PROVIDERS: Final = frozenset(get_args(_BatchLineProvider))
@ -131,14 +133,24 @@ _BatchLineResult: TypeAlias = "ModelResponse | EmbeddingResponse | ResponsesAPIR
def _line_result(
call_type: str, custom_llm_provider: _BatchLineProvider, response_body: Mapping[str, object]
call_type: str,
custom_llm_provider: _BatchLineProvider,
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
bedrock_result: Final = bedrock_batch_line_to_response(response_body, model)
if bedrock_result is None:
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":
return ResponsesAPIResponse(**response_body) # pyright: ignore[reportArgumentType] # same as above
if custom_llm_provider == "anthropic":
from litellm.litellm_core_utils.litellm_logging import anthropic_message_to_model_response
from litellm.llms.anthropic.chat.transformation import anthropic_message_to_model_response
return anthropic_message_to_model_response(response_body, None)
return ModelResponse(**response_body) # pyright: ignore[reportArgumentType] # same as above
@ -216,9 +228,10 @@ async def _emit_line_event(
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)
child: Final = _new_child_logging(
parent=parent,
model=_line_model(response_body, request_body, parent),
model=model,
messages=_line_messages(request_body),
call_type=call_type,
start_time=start_time,
@ -250,7 +263,7 @@ async def _emit_line_event(
stats: Final = _safe_output_line_stats(entry, custom_llm_provider, model_name, model_info)
try:
result: Final = _line_result(call_type, custom_llm_provider, response_body)
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",

View file

@ -53,7 +53,7 @@ def batch_cost_is_final(batch: Batch) -> bool:
async def calculate_batch_cost_and_usage(
file_content_dictionary: list[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
model_name: str | None = None,
model_info: ModelInfo | None = None,
) -> BatchCostUsageResult:
@ -83,7 +83,7 @@ async def calculate_batch_cost_and_usage(
async def _handle_completed_batch(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
model_name: str | None = None,
litellm_params: dict | None = None,
model_info: ModelInfo | None = None,
@ -449,7 +449,7 @@ def _provider_output_file_id(output_file_id: str) -> str:
async def _fetch_batch_managed_file_content(
file_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"] = "openai",
litellm_params: dict | None = None,
) -> bytes:
"""
@ -479,7 +479,7 @@ async def _fetch_batch_managed_file_content(
async def _fetch_batch_output_file_content(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"] = "openai",
litellm_params: dict | None = None,
) -> bytes:
"""
@ -501,7 +501,7 @@ async def _fetch_batch_output_file_content(
async def count_error_file_failed_requests(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
litellm_params: dict | None,
) -> int:
"""Count failed requests reported only in the batch's separate error file.

View file

@ -535,24 +535,6 @@ def _provider_response_id(source: object) -> str | None:
return candidate if isinstance(candidate, str) and candidate else None
def anthropic_message_to_model_response(result: Mapping[str, object], speed: str | None) -> ModelResponse:
import httpx
from litellm.types.llms.anthropic import AnthropicResponse
pydantic_result: Final = AnthropicResponse.model_validate(result)
return litellm.AnthropicConfig().transform_parsed_response(
completion_response=pydantic_result.model_dump(),
raw_response=httpx.Response(
status_code=200,
headers={},
),
model_response=litellm.ModelResponse(id=_provider_response_id(result)),
json_mode=None,
speed=speed,
)
def mask_api_base_credentials(api_base: str) -> str:
if "key=" not in api_base:
return api_base
@ -4255,6 +4237,8 @@ class Logging(LiteLLMLoggingBaseClass):
litellm_params={},
)
else:
from litellm.llms.anthropic.chat.transformation import anthropic_message_to_model_response
result = anthropic_message_to_model_response(
cast(Mapping[str, object], result), # cast-ok: handler result is typed Any upstream
speed=self.optional_params.get("speed") if self.optional_params else None,

View file

@ -48,6 +48,7 @@ from litellm.types.llms.anthropic import (
AnthropicMessagesToolChoice,
AnthropicOutputSchema,
AnthropicOutputTokensDetails,
AnthropicResponse,
AnthropicSystemMessageContent,
AnthropicThinkingParam,
AnthropicWebSearchTool,
@ -2789,3 +2790,15 @@ def _valid_user_id(user_id: str) -> bool:
return False
return True
def anthropic_message_to_model_response(result: Mapping[str, object], speed: str | None) -> ModelResponse:
pydantic_result: Final = AnthropicResponse.model_validate(result)
result_id: Final = result.get("id")
return AnthropicConfig().transform_parsed_response(
completion_response=pydantic_result.model_dump(),
raw_response=httpx.Response(status_code=200, headers={}),
model_response=ModelResponse(id=result_id if isinstance(result_id, str) and result_id else None),
json_mode=None,
speed=speed,
)

View file

@ -27,7 +27,7 @@ from litellm.types.llms.openai import (
AllMessageValues,
CreateBatchRequest,
)
from litellm.types.utils import LiteLLMBatch, LlmProviders, Usage
from litellm.types.utils import EmbeddingResponse, LiteLLMBatch, LlmProviders, ModelResponse, Usage
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import (
@ -96,6 +96,41 @@ def titan_embedding_usage_from_batch_output(model_output: Mapping[str, object])
)
def bedrock_batch_line_to_response(
model_output: Mapping[str, object], model: str
) -> ModelResponse | EmbeddingResponse | None:
"""Reconstruct a Bedrock batch output line (the ``modelOutput`` object) into
the litellm response type its shape implies, or None when the shape is
unrecognized."""
if "embedding" in model_output:
embedding: Final = model_output.get("embedding")
return EmbeddingResponse(
model=model,
data=[{"object": "embedding", "index": 0, "embedding": embedding if isinstance(embedding, list) else []}],
usage=titan_embedding_usage_from_batch_output(model_output),
)
if "output" in model_output:
from ..chat.converse_transformation import AmazonConverseConfig
return AmazonConverseConfig()._transform_response( # pyright: ignore[reportPrivateUsage] # same reconstruction the converse chat path performs on the live response
model=model,
response=Response(200, json=dict(model_output)),
model_response=ModelResponse(),
stream=False,
logging_obj=None,
optional_params={},
api_key=None,
data="",
messages=[],
encoding=None,
)
if "content" in model_output:
from litellm.llms.anthropic.chat.transformation import anthropic_message_to_model_response
return anthropic_message_to_model_response(model_output, None)
return None
class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
"""
Config for Bedrock Batches - handles batch job creation and management for Bedrock

View file

@ -228,7 +228,7 @@ 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())
await _log_completed_batch(_parent_logging(custom_llm_provider="cohere"), _batch())
file_mock.assert_not_awaited()
assert len(recorder.success_events) == 1
@ -391,6 +391,94 @@ ANTHROPIC_OUTPUT_JSONL = json.dumps(
).encode()
BEDROCK_INPUT_JSONL = b"\n".join(
[
json.dumps(
{
"recordId": "br-anth",
"modelInput": {"messages": [{"role": "user", "content": "hi br-anth"}]},
}
).encode(),
json.dumps(
{
"recordId": "br-titan",
"modelInput": {"inputText": "embed me"},
}
).encode(),
json.dumps(
{
"recordId": "br-conv",
"modelInput": {"messages": [{"role": "user", "content": "hi br-conv"}]},
}
).encode(),
]
)
BEDROCK_OUTPUT_JSONL = b"\n".join(
[
json.dumps(
{
"recordId": "br-anth",
"modelOutput": {
"id": "msg_br",
"type": "message",
"role": "assistant",
"content": [{"type": "text", "text": "bedrock claude hi"}],
"model": "anthropic.claude-3",
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 2, "output_tokens": 3},
},
}
).encode(),
json.dumps(
{
"recordId": "br-titan",
"modelOutput": {"embedding": [0.1, 0.2], "inputTextTokenCount": 4},
}
).encode(),
json.dumps(
{
"recordId": "br-conv",
"modelOutput": {
"output": {"message": {"role": "assistant", "content": [{"text": "converse hi"}]}},
"stopReason": "end_turn",
"usage": {"inputTokens": 5, "outputTokens": 6, "totalTokens": 11},
"metrics": {"latencyMs": 12},
},
}
).encode(),
]
)
MISTRAL_INPUT_JSONL = json.dumps(
{
"custom_id": "mis",
"body": {"model": "mistral-small", "messages": [{"role": "user", "content": "hi mis"}]},
}
).encode()
MISTRAL_OUTPUT_JSONL = json.dumps(
{
"custom_id": "mis",
"response": {
"status_code": 200,
"body": {
"id": "mis-1",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "mistral hi"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5},
},
},
}
).encode()
def _edge_file_content(file_id: str, **_kwargs):
return SimpleNamespace(
content={
@ -398,6 +486,10 @@ def _edge_file_content(file_id: str, **_kwargs):
"output-2": EDGE_OUTPUT_JSONL,
"input-anth": ANTHROPIC_INPUT_JSONL,
"output-anth": ANTHROPIC_OUTPUT_JSONL,
"input-bed": BEDROCK_INPUT_JSONL,
"output-bed": BEDROCK_OUTPUT_JSONL,
"input-mist": MISTRAL_INPUT_JSONL,
"output-mist": MISTRAL_OUTPUT_JSONL,
}[file_id]
)
@ -502,3 +594,65 @@ async def test_line_items_anthropic_shapes(recorder):
assert payload["response"]["choices"][0]["message"]["content"] == "hello b2"
assert payload["prompt_tokens"] == 1
assert payload["completion_tokens"] == 2
def _provider_batch(batch_id: str, input_file_id: str, output_file_id: str) -> LiteLLMBatch:
return LiteLLMBatch(
id=batch_id,
object="batch",
endpoint="/v1/chat/completions",
input_file_id=input_file_id,
output_file_id=output_file_id,
error_file_id=None,
status="completed",
completion_window="24h",
created_at=1,
)
@pytest.mark.asyncio
async def test_line_items_bedrock_shapes(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=_edge_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
patch("litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02)), # test-quality-ok: the pricing table boundary, same seam existing batch_utils tests patch
):
await _log_completed_batch(
_parent_logging(custom_llm_provider="bedrock"),
_provider_batch("batch_bed", "input-bed", "output-bed"),
)
by_id = {_hidden(e).get("batch_custom_id"): e for e in recorder.success_events}
anthropic_line = _payload(by_id["br-anth"])
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"])
assert titan_line["response"]["data"][0]["embedding"] == [0.1, 0.2]
assert titan_line["prompt_tokens"] == 4
converse_line = _payload(by_id["br-conv"])
assert converse_line["response"]["choices"][0]["message"]["content"] == "converse hi"
assert converse_line["prompt_tokens"] == 5
assert converse_line["completion_tokens"] == 6
@pytest.mark.asyncio
async def test_line_items_mistral_shape(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=_edge_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
patch("litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02)), # test-quality-ok: the pricing table boundary, same seam existing batch_utils tests patch
):
await _log_completed_batch(
_parent_logging(custom_llm_provider="mistral"),
_provider_batch("batch_mist", "input-mist", "output-mist"),
)
line = next(e for e in recorder.success_events if _hidden(e).get("batch_custom_id") == "mis")
assert _hidden(line)["batch_line_status_code"] == 200
assert _payload(line)["response"]["choices"][0]["message"]["content"] == "mistral hi"