diff --git a/litellm/batches/batch_line_item_logging.py b/litellm/batches/batch_line_item_logging.py index a414c419710..02dd9cceaeb 100644 --- a/litellm/batches/batch_line_item_logging.py +++ b/litellm/batches/batch_line_item_logging.py @@ -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", diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 7209ac6a1e7..468a0c63eae 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -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. diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index de740fdeca7..ea410d6d6ba 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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, diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index b7c2ce3c568..7f9294c58df 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -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, + ) diff --git a/litellm/llms/bedrock/batches/transformation.py b/litellm/llms/bedrock/batches/transformation.py index e4001566b8c..6d43e9ee7f9 100644 --- a/litellm/llms/bedrock/batches/transformation.py +++ b/litellm/llms/bedrock/batches/transformation.py @@ -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 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 1c59eac24bb..0be7f76b594 100644 --- a/tests/test_litellm/batches/test_batch_line_item_logging.py +++ b/tests/test_litellm/batches/test_batch_line_item_logging.py @@ -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"