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