From 21055e3fd840375f11a9f90d8b8e1c176284c79b Mon Sep 17 00:00:00 2001 From: Techboy bebop <142545999+kumarpriyanshu09@users.noreply.github.com> Date: Sun, 27 Sep 2026 00:15:12 -0400 Subject: [PATCH 1/4] fix(tools): salvage concatenated JSON tool call arguments (#43260) * fix(tools): salvage concatenated JSON tool call arguments * fix(tools): harden concatenated tool-call salvage for review findings Skip non-dict JSON during split so salvage cannot emit empty tool calls. Collapse srvtoolu_ expansions to the first object so server results stay paired. Allocate __concat_n ids that cannot collide with sibling tool call ids. Propagate cache_control onto every expanded Anthropic tool_use block. Rename the XML invoke loop variable so the key-leak gate no longer flags {args} * test(tools): cover concat id bump and srvtoolu array keep Only collapse srvtoolu_ when concatenated salvage expanded; a valid JSON array argument stays one server tool input * revert(anthropic): drop concat expansion from pass-through adapter Co-authored-by: Techboy bebop * revert(tools): keep concat salvage out of request-side tool converters Co-authored-by: Techboy bebop * fix(tools): expand strictly salvaged concatenated tool arguments in normalized tool calls Co-authored-by: Techboy bebop * fix(tools): retain at most the salvage cap while validating concatenated arguments Co-authored-by: Techboy bebop * test(tools): assert concat sibling ids unique after sanitization A sibling id that only collides after colon-to-underscore sanitization must force the next concat suffix Co-authored-by: Techboy bebop * refactor(tools): drop unused strict mode from split_concatenated_json_objects Strict mode had no production caller. Rejection cases now sit on salvage, and split matches upstream main Co-authored-by: Techboy bebop --------- Co-authored-by: Techboy bebop --- .../prompt_templates/common_utils.py | 55 ++ .../prompt_templates/factory.py | 210 +++++-- ...ore_utils_prompt_templates_common_utils.py | 40 +- ...llm_core_utils_prompt_templates_factory.py | 540 +++++++++--------- 4 files changed, 527 insertions(+), 318 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 41563d501a9..e555d7e8ec0 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -2503,6 +2503,61 @@ def split_concatenated_json_objects(raw: str) -> list[dict[str, object]]: return results +MAX_SALVAGED_TOOL_ARGUMENT_OBJECTS: Final = 8 + + +def salvage_concatenated_tool_arguments(raw: str) -> tuple[dict[str, object], ...]: + """Return complete concatenated JSON objects that are safe to expand. + + Identical objects collapse to the first one and are not capped. More than + ``MAX_SALVAGED_TOOL_ARGUMENT_OBJECTS`` objects that are not all identical + returns an empty tuple. Anything that is not a full concatenation of JSON + objects returns an empty tuple. Repeated copies of the first object are not + retained, and once the cap is passed the rest of the string is only checked. + """ + stripped: Final = raw.strip() + if not stripped: + return () + decoder: Final = json.JSONDecoder() + length: Final = len(stripped) + idx = 0 # rebind-ok: cursor walks the concatenated JSON string + count = 0 # rebind-ok: counts complete objects without retaining duplicates + kept = () # rebind-ok: holds at most one object past the salvage cap + exceeded = False # rebind-ok: cap already passed, the tail is only validated + while idx < length: + while idx < length and stripped[idx] in " \t\n\r": + idx += 1 + if idx >= length: + break + try: + obj, end_idx = decoder.raw_decode(stripped, idx) + except json.JSONDecodeError: + return () + if not isinstance(obj, dict): + return () + idx = end_idx + if exceeded: + continue + count += 1 + if not kept: + kept = (obj,) + continue + if obj == kept[0] and len(kept) == 1: + continue + if len(kept) == 1 and count > 2 and count - 1 > MAX_SALVAGED_TOOL_ARGUMENT_OBJECTS: + exceeded = True + continue + if len(kept) == 1 and count > 2: + kept = (kept[0],) * (count - 1) + if len(kept) >= MAX_SALVAGED_TOOL_ARGUMENT_OBJECTS: + exceeded = True + continue + kept = (*kept, obj) + if exceeded: + return () + return kept + + def text_completion_prompt_to_messages(prompt: object) -> tuple[AllMessageValues, ...]: """ Wrap an OpenAI ``/v1/completions`` ``prompt`` into Chat Completion messages. diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 7b12d1e939f..c4e242fd360 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1,12 +1,14 @@ import base64 import copy import hashlib +import itertools import json import mimetypes import re import xml.etree.ElementTree as ET from collections.abc import Iterator, Mapping, Sequence from enum import Enum +from types import MappingProxyType from typing import Any, Final, TypeAlias, TypedDict, cast, overload from jinja2.sandbox import ImmutableSandboxedEnvironment @@ -52,6 +54,7 @@ from .common_utils import ( is_non_content_values_set, is_unsignable_thinking_block, parse_tool_call_arguments, + salvage_concatenated_tool_arguments, ) from .image_handling import convert_url_to_base64 @@ -5381,80 +5384,167 @@ class NormalizedToolCall(TypedDict): arguments: dict[str, object] -def _parse_tool_call_arguments(raw: object, tool_name: str | None, context: str) -> dict[str, object]: +_ArgumentObjects: TypeAlias = tuple[dict[str, object], ...] +_ParsedToolCall: TypeAlias = tuple[str | None, str | None, _ArgumentObjects] + + +def _optional_call_id(value: object) -> str | None: + if isinstance(value, str) and value: + return value + return None + + +def _optional_tool_name(value: object) -> str | None: + if isinstance(value, str): + return value + return None + + +def _split_tool_call_ids(calls: Sequence[tuple[str | None, int]]) -> tuple[tuple[str | None, ...], ...]: + taken: Final = frozenset(_sanitize_anthropic_tool_use_id(call_id) for call_id, _ in calls if call_id) + + def fresh(call_id: str) -> Iterator[str]: + return filter( + lambda candidate: _sanitize_anthropic_tool_use_id(candidate) not in taken, + (f"{call_id}__concat_{n}" for n in itertools.count(1)), + ) + + suffixes: Final = MappingProxyType( + {_sanitize_anthropic_tool_use_id(call_id): fresh(call_id) for call_id, count in calls if call_id and count > 1} + ) + return tuple( + ( + call_id, + *(next(suffixes[_sanitize_anthropic_tool_use_id(call_id)]) for _ in range(count - 1)), + ) + if call_id + else (None,) * count + for call_id, count in calls + ) + + +def _parse_tool_call_arguments(raw: object, tool_name: str | None, context: str) -> _ArgumentObjects: # Anthropic's tool_use blocks already carry a parsed dict in "input"; # chat completions and the Responses API carry a JSON string that may be # truncated by the model, so route those through the repair-aware parser. if isinstance(raw, dict): - return raw + return (raw,) if not isinstance(raw, str): - return {} + return ({},) normalized_raw: Final = "{}" if raw == REDACTED_BY_LITELLM else raw - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - parse_tool_call_arguments, - ) - try: parsed: Final = parse_tool_call_arguments(normalized_raw, tool_name=tool_name, context=context) except ValueError as e: + salvaged: Final = salvage_concatenated_tool_arguments(normalized_raw) + if salvaged: + verbose_logger.warning( + "Recovered %d tool call(s) from concatenated JSON arguments for tool '%s' (%s)", + len(salvaged), + tool_name or "", + context, + ) + return salvaged verbose_logger.warning("Failed to parse tool call arguments: %s", e) - return {} - return parsed if isinstance(parsed, dict) else {} + return ({},) + return (parsed,) if isinstance(parsed, dict) else ({},) + + +def _choice_tool_calls(choice: object) -> tuple[object, ...]: + message: Final = get_attribute_or_key(choice, "message", None) + tool_calls: Final = get_attribute_or_key(message, "tool_calls", None) if message is not None else None + if isinstance(tool_calls, list): + return tuple(tool_calls) + return () + + +def _selected_choices(response: object, include_all_choices: bool) -> tuple[object, ...]: + choices: Final = get_attribute_or_key(response, "choices", None) + if not isinstance(choices, list) or not choices: + return () + if include_all_choices: + return tuple(choices) + return (choices[0],) + + +def _parsed_chat_tool_call(tool_call: object) -> _ParsedToolCall | None: + function: Final = get_attribute_or_key(tool_call, "function", None) + if function is None: + return None + name: Final = _optional_tool_name(get_attribute_or_key(function, "name")) + return ( + _optional_call_id(get_attribute_or_key(tool_call, "id")), + name, + _parse_tool_call_arguments( + get_attribute_or_key(function, "arguments", "{}"), + tool_name=name, + context="chat completions", + ), + ) + + +def _parsed_calls_in_choice(choice: object) -> tuple[_ParsedToolCall, ...]: + return tuple( + parsed for tool_call in _choice_tool_calls(choice) if (parsed := _parsed_chat_tool_call(tool_call)) is not None + ) + + +def _parsed_chat_tool_calls(response: object, include_all_choices: bool) -> tuple[_ParsedToolCall, ...]: + grouped: Final = tuple( + _parsed_calls_in_choice(choice) for choice in _selected_choices(response, include_all_choices) + ) + return tuple(itertools.chain.from_iterable(grouped)) + + +def _normalized_tool_calls_for_parse( + name: str | None, + call_ids: tuple[str | None, ...], + arguments: _ArgumentObjects, +) -> tuple[NormalizedToolCall, ...]: + return tuple( + NormalizedToolCall(id=call_id, name=name, arguments=argument) + for call_id, argument in zip(call_ids, arguments, strict=True) + ) + + +def _normalized_tool_calls_from_parses(parses: Sequence[_ParsedToolCall]) -> tuple[NormalizedToolCall, ...]: + id_groups: Final = _split_tool_call_ids(tuple((call_id, len(arguments)) for call_id, _, arguments in parses)) + grouped: Final = tuple( + _normalized_tool_calls_for_parse(name, call_ids, arguments) + for (_, name, arguments), call_ids in zip(parses, id_groups, strict=True) + ) + return tuple(itertools.chain.from_iterable(grouped)) def _tool_calls_from_chat_completion_response( response: object, include_all_choices: bool = False -) -> list[NormalizedToolCall]: - choices: Final = get_attribute_or_key(response, "choices", None) - if not (isinstance(choices, list) and choices): - return [] - tool_calls: Final[list[object]] = [] - for choice in choices if include_all_choices else choices[:1]: - message = get_attribute_or_key(choice, "message", None) - choice_tool_calls = get_attribute_or_key(message, "tool_calls", None) if message else None - if isinstance(choice_tool_calls, list): - tool_calls.extend(choice_tool_calls) - result: Final[list[NormalizedToolCall]] = [] - for tc in tool_calls: - fn = get_attribute_or_key(tc, "function", None) - if fn is None: - continue - name = get_attribute_or_key(fn, "name") - result.append( - NormalizedToolCall( - id=get_attribute_or_key(tc, "id"), - name=name, - arguments=_parse_tool_call_arguments( - get_attribute_or_key(fn, "arguments", "{}"), - tool_name=name, - context="chat completions", - ), - ) - ) - return result +) -> tuple[NormalizedToolCall, ...]: + return _normalized_tool_calls_from_parses(_parsed_chat_tool_calls(response, include_all_choices)) -def _tool_calls_from_responses_api_response(response: object) -> list[NormalizedToolCall]: +def _response_function_calls(response: object) -> tuple[object, ...]: output: Final = get_attribute_or_key(response, "output", None) if not isinstance(output, list): - return [] - result: Final[list[NormalizedToolCall]] = [] - for item in output: - if get_attribute_or_key(item, "type") != "function_call": - continue - name = get_attribute_or_key(item, "name") - result.append( - NormalizedToolCall( - id=get_attribute_or_key(item, "call_id") or get_attribute_or_key(item, "id"), - name=name, - arguments=_parse_tool_call_arguments( - get_attribute_or_key(item, "arguments", "{}"), - tool_name=name, - context="responses API", - ), - ) - ) - return result + return () + return tuple(item for item in output if get_attribute_or_key(item, "type") == "function_call") + + +def _parsed_response_tool_call(item: object) -> _ParsedToolCall: + name: Final = _optional_tool_name(get_attribute_or_key(item, "name")) + raw_id: Final = get_attribute_or_key(item, "call_id") or get_attribute_or_key(item, "id") + return ( + _optional_call_id(raw_id), + name, + _parse_tool_call_arguments( + get_attribute_or_key(item, "arguments", "{}"), + tool_name=name, + context="responses API", + ), + ) + + +def _tool_calls_from_responses_api_response(response: object) -> tuple[NormalizedToolCall, ...]: + parses: Final = tuple(_parsed_response_tool_call(item) for item in _response_function_calls(response)) + return _normalized_tool_calls_from_parses(parses) def _tool_calls_from_anthropic_messages_response(response: object) -> list[NormalizedToolCall]: @@ -5494,16 +5584,18 @@ def get_tool_calls_from_response(response: object, include_all_choices: bool = F Callers that only care about a specific tool should filter the result by ``name`` themselves -- this returns every tool call found. """ - chat_tool_calls = _tool_calls_from_chat_completion_response(response, include_all_choices=include_all_choices) + chat_tool_calls: Final = _tool_calls_from_chat_completion_response( + response, include_all_choices=include_all_choices + ) if chat_tool_calls: - return chat_tool_calls + return list(chat_tool_calls) for extractor in ( _tool_calls_from_responses_api_response, _tool_calls_from_anthropic_messages_response, ): tool_calls = extractor(response) if tool_calls: - return tool_calls + return list(tool_calls) return [] diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py index 79c50bf2369..45fc93f04c1 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py @@ -20,7 +20,9 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( hoist_images_from_tool_messages, is_encrypted_reasoning_block, merge_consecutive_system_messages, + parse_tool_call_arguments, responses_reasoning_items_from_thinking_blocks, + salvage_concatenated_tool_arguments, split_concatenated_json_objects, strip_encrypted_reasoning_from_messages, system_messages_first, @@ -269,6 +271,40 @@ def test_split_concatenated_json_salvages_prefix_before_truncated_tail(): assert result == [{"a": 1}, {"b": 2}] +def test_parse_tool_call_arguments_rejects_concatenated_json() -> None: + with pytest.raises(ValueError, match="Failed to parse tool call arguments"): + parse_tool_call_arguments('{"a":1}{"b":2}') + + +def _distinct_json_objects(count: int) -> str: + return "".join(json.dumps({"n": index}, separators=(",", ":")) for index in range(count)) + + +@pytest.mark.parametrize( + ("raw", "expected"), + ( + ('{"a":1}{"b":2}', ({"a": 1}, {"b": 2})), + ('{"a":1}{"a":1}{"a":1}', ({"a": 1},)), + ('{"a":1}{"a":1}{"b":2}', ({"a": 1}, {"a": 1}, {"b": 2})), + (_distinct_json_objects(8), tuple({"n": index} for index in range(8))), + (_distinct_json_objects(9), ()), + (_distinct_json_objects(9) + " junk", ()), + ('{"a":1}' * 7 + '{"b":2}', tuple({"a": 1} for _ in range(7)) + ({"b": 2},)), + ('{"a":1}' * 8 + '{"b":2}', ()), + ('{"a":1}' * 5000, ({"a": 1},)), + ('{"a":1}' * 20, ({"a": 1},)), + ('{"a":1}{"b":', ()), + ('0{"x":1}', ()), + ('{"x":1}0', ()), + ('[1]{"x":1}', ()), + ('{"a":1}{"b":2}}', ()), + ('{"a":1} junk', ()), + ), +) +def test_salvage_concatenated_tool_arguments(raw: str, expected: tuple[dict[str, object], ...]) -> None: + assert salvage_concatenated_tool_arguments(raw) == expected + + # --------------------------------------------------------------------------- # Regression tests for non-OpenAI file content blocks. # @@ -1949,6 +1985,8 @@ class TestMergeConsecutiveSystemMessages: assert merged == [{"role": "system", "content": expected_content}, {"role": "user", "content": "Hello"}] def test_keeps_the_first_message_when_no_system_message_in_the_run_has_content(self): - merged = merge_consecutive_system_messages([{"role": "system"}, {"role": "system"}, {"role": "user", "content": "Hi"}]) + merged = merge_consecutive_system_messages( + [{"role": "system"}, {"role": "system"}, {"role": "user", "content": "Hi"}] + ) assert merged == [{"role": "system"}, {"role": "user", "content": "Hi"}] diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 0c08c5dfd85..1e12a973cdb 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -1,4 +1,5 @@ import base64 +import json import logging import os import re @@ -16,10 +17,12 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( _bedrock_tools_pt, _rename_duplicate_bedrock_document_names, _convert_to_bedrock_tool_call_invoke, + _sanitize_anthropic_tool_use_id, _convert_to_bedrock_tool_call_result, anthropic_messages_pt, convert_to_anthropic_tool_result, convert_to_gemini_tool_call_result, + get_tool_calls_from_response, make_valid_bedrock_tool_name, ollama_pt, sanitize_messages_for_tool_calling, @@ -31,9 +34,7 @@ def _get_gemini_function_response_inline_data_parts(result): assert isinstance(result, list), "expected Gemini parts list" assert len(result) == 1, "multimodal function responses should stay in one part" function_response_part = result[0] - assert ( - "inline_data" not in function_response_part - ), "inline_data should be nested under function_response.parts" + assert "inline_data" not in function_response_part, "inline_data should be nested under function_response.parts" function_response = function_response_part["function_response"] nested_parts = function_response["parts"] return [part["inline_data"] for part in nested_parts if "inline_data" in part] @@ -49,7 +50,9 @@ def test_ollama_pt_simple_messages(): result = ollama_pt(model="llama2", messages=messages) - expected_prompt = "### System:\nYou are a helpful assistant\n\n### Assistant:\nHow can I help you?\n\n### User:\nHello\n\n" + expected_prompt = ( + "### System:\nYou are a helpful assistant\n\n### Assistant:\nHow can I help you?\n\n### User:\nHello\n\n" + ) assert isinstance(result, dict) assert result["prompt"] == expected_prompt assert result["images"] == [] @@ -104,10 +107,7 @@ async def test_anthropic_bedrock_thinking_blocks_with_none_content(): # verify the result assert len(result) == 2 - assert ( - result[1]["content"][0]["reasoningContent"]["reasoningText"]["text"] - == "This is a test thinking block" - ) + assert result[1]["content"][0]["reasoningContent"]["reasoningText"]["text"] == "This is a test thinking block" def test_bedrock_converse_assistant_with_empty_thinking_block_and_tool_calls(): @@ -175,11 +175,7 @@ def test_bedrock_converse_assistant_with_empty_thinking_block_and_tool_calls(): assert len(assistant_blocks) == 1 for block in assistant_blocks[0]["content"]: if "text" in block: - assert block[ - "text" - ].strip(), ( - f"Bedrock Converse rejects blank-text ContentBlocks; got {block!r}" - ) + assert block["text"].strip(), f"Bedrock Converse rejects blank-text ContentBlocks; got {block!r}" # toolUse blocks must still be present tool_use_blocks = [b for b in assistant_blocks[0]["content"] if "toolUse" in b] assert len(tool_use_blocks) == 2 @@ -220,19 +216,16 @@ def test_anthropic_messages_pt_drops_unsignable_thinking_block(thinking_block): {"role": "user", "content": "Now what is 3+3?"}, ] - result = anthropic_messages_pt( - messages=messages, model="claude-sonnet-4-6", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages=messages, model="claude-sonnet-4-6", llm_provider="anthropic") assistant = next(m for m in result if m["role"] == "assistant") content = assistant["content"] - assert all( - block.get("type") not in ("thinking", "redacted_thinking") for block in content - ), f"unsignable thinking block must be dropped, got {content!r}" - assert any( - block.get("type") == "text" and block.get("text") == "2+2 equals 4." - for block in content - ), f"assistant answer text must be preserved, got {content!r}" + assert all(block.get("type") not in ("thinking", "redacted_thinking") for block in content), ( + f"unsignable thinking block must be dropped, got {content!r}" + ) + assert any(block.get("type") == "text" and block.get("text") == "2+2 equals 4." for block in content), ( + f"assistant answer text must be preserved, got {content!r}" + ) def test_anthropic_messages_pt_keeps_signed_thinking_block(): @@ -255,9 +248,7 @@ def test_anthropic_messages_pt_keeps_signed_thinking_block(): {"role": "user", "content": "Now what is 3+3?"}, ] - result = anthropic_messages_pt( - messages=messages, model="claude-sonnet-4-6", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages=messages, model="claude-sonnet-4-6", llm_provider="anthropic") assistant = next(m for m in result if m["role"] == "assistant") thinking_blocks = [b for b in assistant["content"] if b.get("type") == "thinking"] @@ -373,9 +364,7 @@ def test_bedrock_get_document_format_fallback_mimes(): """ # Test DOCX fallback - docx_mime = ( - "application/vnd.openxmlformats-officedocument.wordprocessingml.document" - ) + docx_mime = "application/vnd.openxmlformats-officedocument.wordprocessingml.document" supported_formats = ["pdf", "docx", "xlsx", "csv"] # Mock mimetypes.guess_all_extensions to return empty list (simulating Docker container scenario) @@ -399,15 +388,11 @@ def test_bedrock_get_document_format_mimetypes_success(): """ Test the _get_document_format method when mimetypes.guess_all_extensions works normally. """ - docx_mime = ( - "application/vnd.openxmlformats-officedocument.wordprocessingml.document" - ) + docx_mime = "application/vnd.openxmlformats-officedocument.wordprocessingml.document" supported_formats = ["pdf", "docx", "xlsx", "csv"] # Test normal mimetypes behavior (should not hit fallback) - result = BedrockImageProcessor._get_document_format( - mime_type=docx_mime, supported_doc_formats=supported_formats - ) + result = BedrockImageProcessor._get_document_format(mime_type=docx_mime, supported_doc_formats=supported_formats) assert result == "docx", f"Expected 'docx', got '{result}'" @@ -623,9 +608,7 @@ async def test_bedrock_process_image_async_factory(): image_url = "data:application/pdf; qs=0.001;base64,JVBERi0xLjQKJcOkw7zDtsOfCjIgMCBvYmoKPDwvTGVuZ3RoIDMgMCBSL0ZpbHRlci9GbGF0ZURlY29kZT4" - content_block = await BedrockImageProcessor.process_image_async( - image_url=image_url, format=None - ) + content_block = await BedrockImageProcessor.process_image_async(image_url=image_url, format=None) print(f"content_block: {content_block}") @@ -668,9 +651,7 @@ def test_unpack_defs_resolves_nested_ref_inside_anyof_items(): items_schema = schema["properties"]["vatAmounts"]["anyOf"][0]["items"] # Assertions: items_schema should now be the resolved object, not an empty dict - assert isinstance( - items_schema, dict - ), "Items schema should be a dict after unpacking" + assert isinstance(items_schema, dict), "Items schema should be a dict after unpacking" assert items_schema.get("type") == "object" # Ensure essential properties are present assert set(items_schema.get("properties", {}).keys()) == {"vatRate", "vatAmount"} @@ -861,9 +842,7 @@ def test_convert_gemini_tool_call_result_with_multiple_anthropic_image_blocks(): last_message_with_tool_calls=last_message_with_tool_calls, ) inline_parts = _get_gemini_function_response_inline_data_parts(result) - assert ( - len(inline_parts) == 2 - ), f"expected 2 inline_data parts, got {len(inline_parts)}" + assert len(inline_parts) == 2, f"expected 2 inline_data parts, got {len(inline_parts)}" mime_types = {p["mime_type"] for p in inline_parts} assert mime_types == {"image/png", "image/jpeg"} @@ -899,9 +878,7 @@ def test_convert_gemini_tool_call_result_with_data_url_string(): last_message_with_tool_calls=last_message_with_tool_calls, ) inline_parts = _get_gemini_function_response_inline_data_parts(result) - assert ( - len(inline_parts) == 1 - ), "data-URL image string was not converted to inline_data" + assert len(inline_parts) == 1, "data-URL image string was not converted to inline_data" assert inline_parts[0]["mime_type"] == "image/png" assert inline_parts[0]["data"] == tiny_png_b64 @@ -937,9 +914,9 @@ def test_convert_gemini_tool_call_result_with_data_url_extra_params(): ) inline_parts = _get_gemini_function_response_inline_data_parts(result) assert len(inline_parts) == 1 - assert ( - inline_parts[0]["mime_type"] == "image/png" - ), f"expected clean 'image/png', got '{inline_parts[0]['mime_type']}'" + assert inline_parts[0]["mime_type"] == "image/png", ( + f"expected clean 'image/png', got '{inline_parts[0]['mime_type']}'" + ) def test_bedrock_tools_unpack_defs(): @@ -1036,9 +1013,7 @@ def test_bedrock_tools_pt_strict_parameter(): }, } ] - result = _bedrock_tools_pt( - tools_with_strict, model="anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result = _bedrock_tools_pt(tools_with_strict, model="anthropic.claude-sonnet-4-5-20250929-v1:0") assert result[0]["toolSpec"]["strict"] is True assert result[0]["toolSpec"]["inputSchema"]["json"]["additionalProperties"] is False @@ -1060,9 +1035,7 @@ def test_bedrock_tools_pt_strict_parameter(): }, } ] - result = _bedrock_tools_pt( - tools_without_strict, model="anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result = _bedrock_tools_pt(tools_without_strict, model="anthropic.claude-sonnet-4-5-20250929-v1:0") assert "strict" not in result[0]["toolSpec"] assert "additionalProperties" not in result[0]["toolSpec"]["inputSchema"]["json"] @@ -1085,9 +1058,7 @@ def test_bedrock_image_processor_content_type_fallback_url_extension(): # Test with .png URL image_url = "https://example.com/test-image.png" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/png" assert base64_bytes == base64.b64encode(png_content).decode("utf-8") @@ -1111,9 +1082,7 @@ def test_bedrock_image_processor_content_type_fallback_binary_detection(): # Test with URL without extension image_url = "https://example.com/test-image-without-extension" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/jpeg" assert base64_bytes == base64.b64encode(jpeg_content).decode("utf-8") @@ -1136,9 +1105,7 @@ def test_bedrock_image_processor_content_type_fallback_application_octet_stream( # Test with .gif URL image_url = "https://s3.amazonaws.com/bucket/image.gif" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/gif" assert base64_bytes == base64.b64encode(gif_content).decode("utf-8") @@ -1161,9 +1128,7 @@ def test_bedrock_image_processor_content_type_with_query_params(): # Test with URL containing query parameters (common in S3 signed URLs) image_url = "https://s3.amazonaws.com/bucket/image.webp?AWSAccessKeyId=123&Expires=456&Signature=789" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/webp" assert base64_bytes == base64.b64encode(webp_content).decode("utf-8") @@ -1185,9 +1150,7 @@ def test_bedrock_image_processor_content_type_normal_header(): mock_response.content = png_content image_url = "https://example.com/test-image.png" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/png" assert base64_bytes == base64.b64encode(png_content).decode("utf-8") @@ -1207,7 +1170,7 @@ def test_bedrock_image_processor_content_type_fallback_failure(): # Test with URL without recognizable extension image_url = "https://example.com/unknown-file" - with pytest.raises(ValueError, match='Unable to determine content type from URL: https') as excinfo: + with pytest.raises(ValueError, match="Unable to determine content type from URL: https") as excinfo: BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert "Unable to determine content type" in str(excinfo.value) @@ -1227,16 +1190,12 @@ def test_bedrock_image_processor_content_type_jpeg_variants(): # Test with .jpg extension image_url_jpg = "https://example.com/photo.jpg" - _, content_type_jpg = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url_jpg - ) + _, content_type_jpg = BedrockImageProcessor._post_call_image_processing(mock_response, image_url_jpg) assert content_type_jpg == "image/jpeg" # Test with .jpeg extension image_url_jpeg = "https://example.com/photo.jpeg" - _, content_type_jpeg = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url_jpeg - ) + _, content_type_jpeg = BedrockImageProcessor._post_call_image_processing(mock_response, image_url_jpeg) assert content_type_jpeg == "image/jpeg" @@ -1258,9 +1217,7 @@ def test_bedrock_image_processor_content_type_pdf_document(): # Test with .pdf URL pdf_url = "https://s3.amazonaws.com/bucket/document.pdf" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, pdf_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, pdf_url) assert content_type == "application/pdf" assert base64_bytes == base64.b64encode(pdf_content).decode("utf-8") @@ -1293,12 +1250,8 @@ def test_bedrock_image_processor_content_type_document_formats(): ] for url, expected_mime in test_cases: - _, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, url - ) - assert ( - content_type == expected_mime - ), f"Expected {expected_mime} for {url}, got {content_type}" + _, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, url) + assert content_type == expected_mime, f"Expected {expected_mime} for {url}, got {content_type}" def test_bedrock_image_processor_content_type_s3_pdf_with_query(): @@ -1317,9 +1270,7 @@ def test_bedrock_image_processor_content_type_s3_pdf_with_query(): # S3 signed URL with query parameters s3_url = "https://my-bucket.s3.us-east-1.amazonaws.com/documents/report.pdf?AWSAccessKeyId=AKIAIOSFODNN7EXAMPLE&Expires=1234567890&Signature=abcdef123456" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, s3_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, s3_url) assert content_type == "application/pdf" assert base64_bytes == base64.b64encode(pdf_content).decode("utf-8") @@ -1428,12 +1379,8 @@ def test_bedrock_create_bedrock_block_normalized_base64(): base64_content = base64.b64encode(pdf_content).decode("utf-8") # Create versions with different whitespace - base64_with_newlines = "\n".join( - [base64_content[i : i + 64] for i in range(0, len(base64_content), 64)] - ) - base64_with_spaces = " ".join( - [base64_content[i : i + 32] for i in range(0, len(base64_content), 32)] - ) + base64_with_newlines = "\n".join([base64_content[i : i + 64] for i in range(0, len(base64_content), 64)]) + base64_with_spaces = " ".join([base64_content[i : i + 32] for i in range(0, len(base64_content), 32)]) # Create blocks block1 = BedrockImageProcessor._create_bedrock_block( @@ -1565,9 +1512,7 @@ def test_bedrock_create_bedrock_block_document_name_format(): # Check format: DocumentPDFmessages_{16_hex_chars}_{format} pattern = r"^DocumentPDFmessages_[0-9a-f]{16}_pdf$" - assert re.match( - pattern, document_name - ), f"Document name format mismatch: {document_name}" + assert re.match(pattern, document_name), f"Document name format mismatch: {document_name}" def test_bedrock_create_bedrock_block_different_document_formats(): @@ -1620,9 +1565,7 @@ def test_bedrock_nova_web_search_options_mapping(): assert system_tool["name"] == "nova_grounding" # Test with search_context_size (should be ignored for Nova) - result2 = config._map_web_search_options( - {"search_context_size": "high"}, "us.amazon.nova-premier-v1:0" - ) + result2 = config._map_web_search_options({"search_context_size": "high"}, "us.amazon.nova-premier-v1:0") assert result2 is not None system_tool2 = result2.get("systemTool") @@ -1688,9 +1631,7 @@ def test_bedrock_tools_pt_drops_unmappable_responses_builtin_tools(): {"type": "custom", "name": "free_form"}, ] - result = _bedrock_tools_pt( - tools=tools, model="anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result = _bedrock_tools_pt(tools=tools, model="anthropic.claude-sonnet-4-5-20250929-v1:0") names = [block["toolSpec"]["name"] for block in result if "toolSpec" in block] assert names == ["noop"] @@ -1720,9 +1661,7 @@ def test_bedrock_tools_pt_keeps_anthropic_input_schema_tools(): }, ] - result = _bedrock_tools_pt( - tools=tools, model="anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result = _bedrock_tools_pt(tools=tools, model="anthropic.claude-sonnet-4-5-20250929-v1:0") names = [block["toolSpec"]["name"] for block in result if "toolSpec" in block] assert names == ["lookup"] @@ -1924,9 +1863,7 @@ def test_anthropic_messages_pt_server_tool_use_passthrough(): "tool_use_id": "srvtoolu_01ABC123", "content": { "type": "tool_search_tool_search_result", - "tool_references": [ - {"type": "tool_reference", "tool_name": "get_time"} - ], + "tool_references": [{"type": "tool_reference", "tool_name": "get_time"}], }, }, {"type": "text", "text": "I found the time tool. How can I help you?"}, @@ -1954,20 +1891,14 @@ def test_anthropic_messages_pt_server_tool_use_passthrough(): # Verify server_tool_use block is preserved assert "server_tool_use" in content_types - server_tool_use_block = next( - b for b in assistant_msg["content"] if b.get("type") == "server_tool_use" - ) + server_tool_use_block = next(b for b in assistant_msg["content"] if b.get("type") == "server_tool_use") assert server_tool_use_block["id"] == "srvtoolu_01ABC123" assert server_tool_use_block["name"] == "tool_search_tool_regex" assert server_tool_use_block["input"] == {"query": ".*time.*"} # Verify tool_search_tool_result block is preserved assert "tool_search_tool_result" in content_types - tool_result_block = next( - b - for b in assistant_msg["content"] - if b.get("type") == "tool_search_tool_result" - ) + tool_result_block = next(b for b in assistant_msg["content"] if b.get("type") == "tool_search_tool_result") assert tool_result_block["tool_use_id"] == "srvtoolu_01ABC123" assert tool_result_block["content"]["type"] == "tool_search_tool_search_result" assert tool_result_block["content"]["tool_references"][0]["tool_name"] == "get_time" @@ -2019,9 +1950,7 @@ def test_bedrock_tools_unpack_defs_no_oom_with_nested_refs(): "anyOf": [ {"$ref": "#/$defs/Literal"}, {"$ref": "#/$defs/FieldRef"}, - { - "$ref": "#/$defs/Expression" - }, # Circular: Operand -> Expression -> Operand + {"$ref": "#/$defs/Expression"}, # Circular: Operand -> Expression -> Operand ], }, "Literal": { @@ -2155,9 +2084,7 @@ def test_anthropic_messages_pt_file_block_cache_control_with_explicit_provider() file_block = content_blocks[0] assert file_block["type"] == "document" - assert ( - "cache_control" in file_block - ), "cache_control should be preserved on file/document content blocks" + assert "cache_control" in file_block, "cache_control should be preserved on file/document content blocks" assert file_block["cache_control"]["type"] == "ephemeral" text_block = content_blocks[1] @@ -2365,22 +2292,16 @@ def test_bedrock_tool_call_invoke_concatenated_json(): # First block keeps original tool id assert result[0]["toolUse"]["toolUseId"] == "tooluse_L7I3TewYAUhoheJZQEuwVN" assert result[0]["toolUse"]["name"] == "shell" - assert result[0]["toolUse"]["input"] == { - "command": ["curl", "-i", "http://localhost:9009", "-m", "10"] - } + assert result[0]["toolUse"]["input"] == {"command": ["curl", "-i", "http://localhost:9009", "-m", "10"]} # Subsequent blocks get suffixed ids assert result[1]["toolUse"]["toolUseId"] == "tooluse_L7I3TewYAUhoheJZQEuwVN_1" assert result[1]["toolUse"]["name"] == "shell" - assert result[1]["toolUse"]["input"] == { - "command": ["curl", "-i", "http://localhost:9009/robots.txt", "-m", "5"] - } + assert result[1]["toolUse"]["input"] == {"command": ["curl", "-i", "http://localhost:9009/robots.txt", "-m", "5"]} assert result[2]["toolUse"]["toolUseId"] == "tooluse_L7I3TewYAUhoheJZQEuwVN_2" assert result[2]["toolUse"]["name"] == "shell" - assert result[2]["toolUse"]["input"] == { - "command": ["curl", "-i", "http://localhost:9009/sitemap.xml", "-m", "5"] - } + assert result[2]["toolUse"]["input"] == {"command": ["curl", "-i", "http://localhost:9009/sitemap.xml", "-m", "5"]} def test_bedrock_tool_call_invoke_concatenated_json_with_cache_control(): @@ -2535,9 +2456,7 @@ def test_bedrock_tool_call_invoke_unconvertible_raises_non_retryable_bad_request def test_make_valid_bedrock_tool_name_preserves_hyphens(): assert make_valid_bedrock_tool_name("my-tool") == "my-tool" assert ( - make_valid_bedrock_tool_name( - "CreateCaseKnowledgeArticle_foTWsqR6yDt-OnSsvR5e6Q" - ) + make_valid_bedrock_tool_name("CreateCaseKnowledgeArticle_foTWsqR6yDt-OnSsvR5e6Q") == "CreateCaseKnowledgeArticle_foTWsqR6yDt-OnSsvR5e6Q" ) @@ -2564,9 +2483,7 @@ def test_bedrock_tool_name_sanitized_consistently_in_tools_and_tool_use(): "function": {"name": raw_name, "arguments": "{}"}, } ] - tool_use_name = _convert_to_bedrock_tool_call_invoke(tool_calls)[0]["toolUse"][ - "name" - ] + tool_use_name = _convert_to_bedrock_tool_call_invoke(tool_calls)[0]["toolUse"]["name"] assert tool_spec_name == "foo_bar" assert tool_use_name == tool_spec_name @@ -2589,15 +2506,8 @@ def test_bedrock_converse_messages_pt_tool_use_matches_tool_spec_hyphen_name(): ], }, ] - translated = _bedrock_converse_messages_pt( - messages=messages, model="", llm_provider="" - ) - tool_use_blocks = [ - block - for msg in translated - for block in msg.get("content", []) - if "toolUse" in block - ] + translated = _bedrock_converse_messages_pt(messages=messages, model="", llm_provider="") + tool_use_blocks = [block for msg in translated for block in msg.get("content", []) if "toolUse" in block] assert len(tool_use_blocks) == 1 assert tool_use_blocks[0]["toolUse"]["name"] == tool_name @@ -2694,11 +2604,7 @@ def test_sanitize_messages_deduplicates_tool_results(): result = sanitize_messages_for_tool_calling(messages) # Count tool messages with this ID — should be exactly 1 - tool_results = [ - m - for m in result - if m.get("role") == "tool" and m.get("tool_call_id") == "call_abc123" - ] + tool_results = [m for m in result if m.get("role") == "tool" and m.get("tool_call_id") == "call_abc123"] assert len(tool_results) == 1 # Should keep the LAST occurrence (most complete) assert tool_results[0]["content"] == '{"temperature": 72, "condition": "sunny"}' @@ -2833,11 +2739,7 @@ def test_sanitize_messages_dedup_scoped_per_turn_preserves_cross_turn(): result = sanitize_messages_for_tool_calling(messages) # Both tool results must survive — one per turn - tool_results = [ - m - for m in result - if m.get("role") == "tool" and m.get("tool_call_id") == "call_X" - ] + tool_results = [m for m in result if m.get("role") == "tool" and m.get("tool_call_id") == "call_X"] assert len(tool_results) == 2, ( f"Expected 2 tool results (one per turn), got {len(tool_results)}. " "Dedup may be global instead of per-turn scoped." @@ -2891,32 +2793,26 @@ def test_sanitize_messages_combined_case_a_and_case_d(): tool_results = [m for m in result if m.get("role") in ("tool", "function")] # Case A: call_missing should have a dummy result injected - missing_results = [ - m for m in tool_results if m.get("tool_call_id") == "call_missing" - ] - assert ( - len(missing_results) == 1 - ), f"Expected 1 dummy result for call_missing (Case A), got {len(missing_results)}" + missing_results = [m for m in tool_results if m.get("tool_call_id") == "call_missing"] + assert len(missing_results) == 1, ( + f"Expected 1 dummy result for call_missing (Case A), got {len(missing_results)}" + ) # Case D: call_duped should have exactly 1 result (the fresh one) - duped_results = [ - m for m in tool_results if m.get("tool_call_id") == "call_duped" - ] - assert ( - len(duped_results) == 1 - ), f"Expected 1 result for call_duped after dedup (Case D), got {len(duped_results)}" - assert ( - duped_results[0]["content"] == "fresh_result" - ), f"Expected last-wins 'fresh_result', got '{duped_results[0]['content']}'" + duped_results = [m for m in tool_results if m.get("tool_call_id") == "call_duped"] + assert len(duped_results) == 1, ( + f"Expected 1 result for call_duped after dedup (Case D), got {len(duped_results)}" + ) + assert duped_results[0]["content"] == "fresh_result", ( + f"Expected last-wins 'fresh_result', got '{duped_results[0]['content']}'" + ) # Verify tool results immediately follow the assistant message asst_idx = next(i for i, m in enumerate(result) if m.get("role") == "assistant") - tool_msgs_after_asst = [ - m for m in result[asst_idx + 1 :] if m.get("role") in ("tool", "function") - ] - assert ( - len(tool_msgs_after_asst) == 2 - ), f"Expected 2 tool results after assistant, got {len(tool_msgs_after_asst)}" + tool_msgs_after_asst = [m for m in result[asst_idx + 1 :] if m.get("role") in ("tool", "function")] + assert len(tool_msgs_after_asst) == 2, ( + f"Expected 2 tool results after assistant, got {len(tool_msgs_after_asst)}" + ) # Both tool_call_ids should be present (order may vary) tool_ids = {m["tool_call_id"] for m in tool_msgs_after_asst} assert tool_ids == { @@ -2958,9 +2854,7 @@ def test_anthropic_messages_pt_file_block_preserves_cache_control(): } ] - result = anthropic_messages_pt( - messages, model="claude-sonnet-4-20250514", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages, model="claude-sonnet-4-20250514", llm_provider="anthropic") content_blocks = result[0]["content"] assert len(content_blocks) == 2 @@ -2968,9 +2862,7 @@ def test_anthropic_messages_pt_file_block_preserves_cache_control(): # Document block (from file) should preserve cache_control doc_block = content_blocks[0] assert doc_block["type"] == "document" - assert ( - "cache_control" in doc_block - ), "cache_control was dropped from file/document block" + assert "cache_control" in doc_block, "cache_control was dropped from file/document block" assert doc_block["cache_control"]["type"] == "ephemeral" # Text block should also preserve cache_control @@ -3013,9 +2905,7 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5(monkeypatch): } # Claude 4.5 model: ttl should be preserved - result = add_cache_point_tool_block( - tool_with_1h, model="jp.anthropic.claude-opus-4-7" - ) + result = add_cache_point_tool_block(tool_with_1h, model="jp.anthropic.claude-opus-4-7") assert result is not None assert result["cachePoint"]["type"] == "default" assert result["cachePoint"]["ttl"] == "1h" @@ -3024,16 +2914,12 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5(monkeypatch): tool_with_5m = { "cache_control": {"type": "ephemeral", "ttl": "5m"}, } - result_5m = add_cache_point_tool_block( - tool_with_5m, model="jp.anthropic.claude-opus-4-7" - ) + result_5m = add_cache_point_tool_block(tool_with_5m, model="jp.anthropic.claude-opus-4-7") assert result_5m is not None assert result_5m["cachePoint"]["ttl"] == "5m" # Older model: ttl should be stripped - result_old = add_cache_point_tool_block( - tool_with_1h, model="anthropic.claude-3-5-sonnet-20241022-v2:0" - ) + result_old = add_cache_point_tool_block(tool_with_1h, model="anthropic.claude-3-5-sonnet-20241022-v2:0") assert result_old is not None assert result_old["cachePoint"]["type"] == "default" assert "ttl" not in result_old["cachePoint"] @@ -3052,9 +2938,7 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5(monkeypatch): # cache_control without ttl: returns default cachePoint (unchanged behavior) tool_no_ttl = {"cache_control": {"type": "ephemeral"}} - result_no_ttl = add_cache_point_tool_block( - tool_no_ttl, model="us.anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result_no_ttl = add_cache_point_tool_block(tool_no_ttl, model="us.anthropic.claude-sonnet-4-5-20250929-v1:0") assert result_no_ttl is not None assert result_no_ttl["cachePoint"]["type"] == "default" assert "ttl" not in result_no_ttl["cachePoint"] @@ -3127,9 +3011,7 @@ def test_bedrock_tools_pt_passes_ttl_for_claude_4_5(monkeypatch): assert cache_blocks[0]["cachePoint"]["ttl"] == "1h" # Older model: cachePoint should not have ttl - result_old = _bedrock_tools_pt( - tools, model="anthropic.claude-3-5-sonnet-20241022-v2:0" - ) + result_old = _bedrock_tools_pt(tools, model="anthropic.claude-3-5-sonnet-20241022-v2:0") cache_blocks_old = [b for b in result_old if "cachePoint" in b] assert len(cache_blocks_old) == 1 assert "ttl" not in cache_blocks_old[0]["cachePoint"] @@ -3204,9 +3086,7 @@ def test_bedrock_converse_messages_pt_document_various_formats(): } ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") doc_block = result[0]["content"][0] assert doc_block["document"]["format"] == expected_format, ( @@ -3233,12 +3113,8 @@ def test_bedrock_converse_messages_pt_document_deterministic_name(): } ] - result1 = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) - result2 = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) + result1 = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") + result2 = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") name1 = result1[0]["content"][0]["document"]["name"] name2 = result2[0]["content"][0]["document"]["name"] @@ -3272,34 +3148,18 @@ def test_bedrock_converse_messages_pt_renames_duplicate_document_names(): }, ] - result1 = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) - result2 = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) + result1 = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") + result2 = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") - names1 = [ - block["document"]["name"] - for message in result1 - for block in message["content"] - if "document" in block - ] - names2 = [ - block["document"]["name"] - for message in result2 - for block in message["content"] - if "document" in block - ] + names1 = [block["document"]["name"] for message in result1 for block in message["content"] if "document" in block] + names2 = [block["document"]["name"] for message in result2 for block in message["content"] if "document" in block] assert len(names1) == 2 assert len(set(names1)) == 2 assert names1[1] == f"{names1[0]}_2" assert names1 == names2 - single_turn = _bedrock_converse_messages_pt( - [messages[0]], "anthropic.claude-sonnet-4-6", "bedrock" - ) + single_turn = _bedrock_converse_messages_pt([messages[0]], "anthropic.claude-sonnet-4-6", "bedrock") assert names1[0] == single_turn[0]["content"][0]["document"]["name"] @@ -3321,14 +3181,10 @@ def test_rename_duplicate_bedrock_document_names_skips_organic_suffixes(): def _names(contents): return [block["document"]["name"] for block in contents[0]["content"]] - organic_first = _rename_duplicate_bedrock_document_names( - _contents(["report", "report_2", "report"]) - ) + organic_first = _rename_duplicate_bedrock_document_names(_contents(["report", "report_2", "report"])) assert _names(organic_first) == ["report", "report_2", "report_3"] - organic_last = _rename_duplicate_bedrock_document_names( - _contents(["report", "report", "report_2"]) - ) + organic_last = _rename_duplicate_bedrock_document_names(_contents(["report", "report", "report_2"])) assert _names(organic_last) == ["report", "report_3", "report_2"] @@ -3350,18 +3206,11 @@ def test_bedrock_converse_messages_pt_document_rejects_url_source(): ] with pytest.raises(ValueError, match="only supports base64-encoded"): - _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) + _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") def _collect_cache_points(blocks): - return [ - block["cachePoint"] - for message in blocks - for block in message["content"] - if "cachePoint" in block - ] + return [block["cachePoint"] for message in blocks for block in message["content"] if "cachePoint" in block] @pytest.mark.parametrize( @@ -3527,6 +3376,189 @@ def test_get_tool_calls_from_response_warns_for_malformed_arguments(caplog): assert "Failed to parse tool call arguments" in caplog.text +def _concatenated_json(*payloads: dict[str, object]) -> str: + return "".join(json.dumps(payload, separators=(",", ":")) for payload in payloads) + + +def _function_tool_call(call_id: str | None, name: str, arguments: str) -> dict[str, object]: + return {"id": call_id, "function": {"name": name, "arguments": arguments}} + + +def _chat_tool_response(*tool_calls: dict[str, object]) -> dict[str, object]: + return {"choices": [{"message": {"tool_calls": list(tool_calls)}}]} + + +def test_get_tool_calls_from_response_expands_distinct_concatenated_arguments(caplog): + raw = '{"flag":true}{"box":"A","limit":50}' + response: Final = _chat_tool_response(_function_tool_call("call_move", "move", raw)) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + tool_calls: Final = get_tool_calls_from_response(response) + + assert tool_calls == [ + {"id": "call_move", "name": "move", "arguments": {"flag": True}}, + {"id": "call_move__concat_1", "name": "move", "arguments": {"box": "A", "limit": 50}}, + ] + assert "Recovered 2 tool call(s)" in caplog.text + assert "move" in caplog.text + assert "flag" not in caplog.text + + +def test_get_tool_calls_from_response_expands_responses_api_concatenated_arguments(): + response: Final = { + "output": [ + { + "type": "function_call", + "call_id": "call_move", + "name": "move", + "arguments": '{"flag":true}{"box":"A","limit":50}', + } + ] + } + + assert get_tool_calls_from_response(response) == [ + {"id": "call_move", "name": "move", "arguments": {"flag": True}}, + {"id": "call_move__concat_1", "name": "move", "arguments": {"box": "A", "limit": 50}}, + ] + + +def test_get_tool_calls_from_response_collapses_identical_concatenated_arguments(): + response: Final = _chat_tool_response(_function_tool_call("call_move", "move", '{"flag":true}' * 3)) + + assert get_tool_calls_from_response(response) == [ + {"id": "call_move", "name": "move", "arguments": {"flag": True}}, + ] + + +def test_get_tool_calls_from_response_does_not_expand_a_valid_json_array(): + response: Final = _chat_tool_response(_function_tool_call("call_batch", "batch", '[{"a":1},{"b":2}]')) + + assert get_tool_calls_from_response(response) == [ + {"id": "call_batch", "name": "batch", "arguments": {}}, + ] + + +@pytest.mark.parametrize("arguments", ('{"a":1}{"b":', '0{"x":1}')) +def test_get_tool_calls_from_response_drops_partial_concatenated_arguments(arguments: str, caplog): + response: Final = _chat_tool_response(_function_tool_call("call_move", "move", arguments)) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + tool_calls: Final = get_tool_calls_from_response(response) + + assert tool_calls == [{"id": "call_move", "name": "move", "arguments": {}}] + assert "Failed to parse tool call arguments" in caplog.text + + +@pytest.mark.parametrize(("count", "expands"), ((8, True), (9, False))) +def test_get_tool_calls_from_response_caps_distinct_concatenated_arguments(count: int, expands: bool): + raw = _concatenated_json(*({"n": index} for index in range(count))) + response: Final = _chat_tool_response(_function_tool_call("call", "move", raw)) + + tool_calls: Final = get_tool_calls_from_response(response) + + if expands: + assert [call["id"] for call in tool_calls] == ["call", *(f"call__concat_{index}" for index in range(1, count))] + assert [call["arguments"] for call in tool_calls] == [{"n": index} for index in range(count)] + return + assert tool_calls == [{"id": "call", "name": "move", "arguments": {}}] + + +def test_get_tool_calls_from_response_skips_concat_ids_taken_by_a_sibling(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response( + _function_tool_call("call", "move", raw), + _function_tool_call("call__concat_1", "look", '{"x":1}'), + ) + + assert [call["id"] for call in get_tool_calls_from_response(response)] == [ + "call", + "call__concat_2", + "call__concat_1", + ] + + +def test_get_tool_calls_from_response_keeps_sanitized_concat_ids_distinct(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response( + _function_tool_call("a:b", "move", raw), + _function_tool_call("a_b__concat_1", "look", '{"x":1}'), + ) + + ids: Final = [call["id"] for call in get_tool_calls_from_response(response)] + sanitized: Final = [_sanitize_anthropic_tool_use_id(call_id) for call_id in ids if isinstance(call_id, str)] + + assert len(sanitized) == len(set(sanitized)) + assert ids == ["a:b", "a:b__concat_2", "a_b__concat_1"] + + +def test_get_tool_calls_from_response_bumps_suffix_when_sibling_sanitizes_onto_it(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response( + _function_tool_call("a_b", "move", raw), + _function_tool_call("a:b__concat_1", "look", '{"x":1}'), + ) + + ids: Final = [call["id"] for call in get_tool_calls_from_response(response)] + sanitized: Final = [_sanitize_anthropic_tool_use_id(call_id) for call_id in ids if isinstance(call_id, str)] + + assert len(ids) == len(sanitized) + assert len(sanitized) == len(set(sanitized)) + assert ids == ["a_b", "a_b__concat_2", "a:b__concat_1"] + + +def test_get_tool_calls_from_response_continues_concat_suffixes_per_sanitized_base(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response( + _function_tool_call("x", "move", raw), + _function_tool_call("x", "move", raw), + ) + + assert [call["id"] for call in get_tool_calls_from_response(response)] == [ + "x", + "x__concat_1", + "x", + "x__concat_2", + ] + + +def test_get_tool_calls_from_response_skips_a_run_of_reserved_concat_ids(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + siblings: Final = tuple(_function_tool_call(f"call__concat_{index}", "look", '{"x":1}') for index in range(1, 51)) + response: Final = _chat_tool_response(_function_tool_call("call", "move", raw), *siblings) + + ids: Final = [call["id"] for call in get_tool_calls_from_response(response)] + + assert ids[0] == "call" + assert ids[1] == "call__concat_51" + + +def test_get_tool_calls_from_response_reserves_concat_ids_across_choices(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = { + "choices": [ + {"message": {"tool_calls": [_function_tool_call("call", "move", raw)]}}, + {"message": {"tool_calls": [_function_tool_call("call__concat_1", "look", '{"x":1}')]}}, + ] + } + + assert [call["id"] for call in get_tool_calls_from_response(response, include_all_choices=True)] == [ + "call", + "call__concat_2", + "call__concat_1", + ] + + +def test_get_tool_calls_from_response_does_not_invent_ids_for_a_missing_call_id(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response(_function_tool_call(None, "move", raw)) + + tool_calls: Final = get_tool_calls_from_response(response) + + assert len(tool_calls) == 2 + assert all(call["id"] is None for call in tool_calls) + assert [call["arguments"] for call in tool_calls] == [{"a": 1}, {"b": 2}] + + def test_group_tool_exchanges_pairs_assistant_with_its_tool_rows(): from litellm.litellm_core_utils.prompt_templates.factory import group_tool_exchanges @@ -3625,9 +3657,7 @@ def test_bedrock_converse_pdf_only_user_message_gets_text_block(): } ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-haiku-4-5", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-haiku-4-5", "bedrock") assert len(result) == 1 assert any("document" in block for block in result[0]["content"]) @@ -3645,9 +3675,7 @@ def test_bedrock_converse_document_with_text_gets_no_extra_text_block(): } ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-haiku-4-5", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-haiku-4-5", "bedrock") assert _text_blocks(result[0]) == ["summarize this"] @@ -3660,9 +3688,7 @@ def test_bedrock_converse_image_only_user_message_gets_no_text_block(): } ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-haiku-4-5", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-haiku-4-5", "bedrock") assert any("image" in block for block in result[0]["content"]) assert _text_blocks(result[0]) == [] @@ -3705,9 +3731,7 @@ def test_bedrock_converse_tool_round_trip_document_injects_text_before_cache_poi }, ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-haiku-4-5", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-haiku-4-5", "bedrock") assert _text_blocks(result[0]) == ["read the pdf"] document_message = result[-1] From 829cba1bf18c22593ddf65737e30d1b905651259 Mon Sep 17 00:00:00 2001 From: Thippaluri Yaseen Basha Date: Sun, 27 Sep 2026 09:48:33 +0530 Subject: [PATCH 2/4] fix(gemini): forward seed to the Gemini API instead of rejecting it (#43197) * fix(gemini): forward seed to the Gemini API instead of rejecting it The gemini/ provider left seed out of its supported params, so requests with seed failed with UnsupportedParamsError, or lost the seed silently when drop_params was on. The Gemini API accepts generationConfig.seed and the inherited mapping already translates it, so adding it to the allowlist is enough * test(gemini): assert the forwarded seed without mutating shared state --- litellm/llms/gemini/chat/transformation.py | 1 + ...test_vertex_and_google_ai_studio_gemini.py | 28 +++++++++++++++++++ 2 files changed, 29 insertions(+) diff --git a/litellm/llms/gemini/chat/transformation.py b/litellm/llms/gemini/chat/transformation.py index 285350aecba..cae0ba49c3f 100644 --- a/litellm/llms/gemini/chat/transformation.py +++ b/litellm/llms/gemini/chat/transformation.py @@ -96,6 +96,7 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig): "logprobs", "frequency_penalty", "presence_penalty", + "seed", "modalities", "parallel_tool_calls", "web_search_options", diff --git a/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 7548f3c2daa..fd735afb16e 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -3413,6 +3413,34 @@ def test_google_ai_studio_presence_penalty_supported(): assert "presence_penalty" in supported_params +@pytest.mark.asyncio +@pytest.mark.parametrize("drop_params", [False, True]) +async def test_google_ai_studio_forwards_seed_to_generation_config(drop_params: bool): + def echo_seed_sent_upstream(request: httpx.Request) -> httpx.Response: + seed_sent: Final = json.loads(request.content).get("generationConfig", {}).get("seed") + return httpx.Response( + 200, + json={ + "candidates": [ + {"content": {"parts": [{"text": f"seed={seed_sent}"}], "role": "model"}, "finishReason": "STOP"} + ], + "usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }, + request=request, + ) + + response: Final = await litellm.acompletion( + model="gemini/gemini-3.8-flash", + messages=[{"role": "user", "content": "hi"}], + seed=42, + drop_params=drop_params, + api_key="fake-gemini-key", + client=AsyncHTTPHandler(transport=httpx.MockTransport(echo_seed_sent_upstream)), + ) + + assert response.choices[0].message.content == "seed=42" + + # ==================== Tool Type Separation Tests ==================== # These tests verify that each Tool object contains exactly one type per Vertex AI API spec # Ref: https://cloud.google.com/vertex-ai/generative-ai/docs/reference/rest/v1beta1/Tool From 2101c860c25546250fa55c5a67653093a9d28873 Mon Sep 17 00:00:00 2001 From: Stewart Park <388348+stewartpark@users.noreply.github.com> Date: Sat, 26 Sep 2026 21:26:38 -0700 Subject: [PATCH 3/4] fix(vertex_ai): make Gemma fake streams work with traced Responses (#43147) * test(vertex_ai): reproduce traced Gemma Responses stream failure * fix(vertex_ai): wrap Gemma fake streams for Responses tracing * test(vertex_ai): cover Gemma traced streams and usage options * test(vertex_ai): inject gemma test deps and assert hidden usage accounting Replace class-level patches in the Vertex AI shard test with the provider's documented dependency-injection seams (httpx.MockTransport client + credential cache), and pin the default/omit-usage trace behavior: LiteLLM still accounts all tokens; ddtrace's metric is absent by design, asserted rather than silent. Mutation-checked: commenting out CustomStreamWrapper chunk accumulation turns the new assertions red; restoring them turns green. * test(vertex_ai): drop explanatory comment from usage-option assertions --- .../vertex_gemma_models/transformation.py | 30 ++++- .../test_vertex_gemma_transformation.py | 93 ++++++++++++++ .../test_vertex_gemma_transformation.py | 117 +++++++++++++++++- 3 files changed, 229 insertions(+), 11 deletions(-) create mode 100644 tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py index ea97f0a0a9a..33922e38674 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py @@ -28,8 +28,8 @@ from litellm.types.utils import ModelResponse if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer - from litellm.llms.base_llm.base_model_iterator import MockResponseIterator def parse_vertex_gemma_container_error(predictions: object) -> VertexGemmaContainerError | None: @@ -73,7 +73,9 @@ class VertexGemmaConfig(OpenAIGPTConfig): self, model_response: ModelResponse, stream: bool, - ) -> "ModelResponse | MockResponseIterator": + model: str, + logging_obj: "LiteLLMLoggingObj", + ) -> "ModelResponse | CustomStreamWrapper": """ Helper method to return fake stream iterator if streaming is requested. @@ -82,12 +84,18 @@ class VertexGemmaConfig(OpenAIGPTConfig): stream: Whether streaming was requested Returns: - MockResponseIterator if stream=True, otherwise the model_response + CustomStreamWrapper if stream=True, otherwise the model_response """ if stream: + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.llms.base_llm.base_model_iterator import MockResponseIterator - return MockResponseIterator(model_response=model_response) + return CustomStreamWrapper( + completion_stream=MockResponseIterator(model_response=model_response), + model=model, + custom_llm_provider="vertex_ai", + logging_obj=logging_obj, + ) return model_response def transform_request( @@ -373,7 +381,12 @@ class VertexGemmaConfig(OpenAIGPTConfig): ) # Return fake stream iterator if streaming was requested - return self._handle_fake_stream_response(model_response=model_response, stream=stream) + return self._handle_fake_stream_response( + model_response=model_response, + stream=stream, + model=model, + logging_obj=logging_obj, + ) async def _async_completion( self, @@ -463,4 +476,9 @@ class VertexGemmaConfig(OpenAIGPTConfig): ) # Return fake stream iterator if streaming was requested - return self._handle_fake_stream_response(model_response=model_response, stream=stream) + return self._handle_fake_stream_response( + model_response=model_response, + stream=stream, + model=model, + logging_obj=logging_obj, + ) diff --git a/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py new file mode 100644 index 00000000000..294d26b2e58 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -0,0 +1,93 @@ +import json +from collections.abc import AsyncIterator +from types import SimpleNamespace +from typing import Any, cast + +import httpx +import pytest + +import litellm +from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +from litellm.main import vertex_gemma_chat_completion +from litellm.types.llms.openai import OutputTextDeltaEvent, ResponseCompletedEvent, ResponsesAPIStreamingResponse + +_VERTEX_URL = "https://example.invalid/v1/projects/test/locations/us-central1/endpoints/test:predict" +_MESSAGES = [{"role": "user", "content": "Reply exactly READY"}] +_FAKE_CREDENTIALS = "gemma-test-credentials" + + +def _vertex_response(): + return { + "predictions": { + "id": "chatcmpl-stream-test", + "created": 1759863903, + "model": "google/gemma-3-12b-it", + "object": "chat.completion", + "choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "READY"}}], + "usage": {"prompt_tokens": 14, "completion_tokens": 1, "total_tokens": 15}, + } + } + + +@pytest.fixture(autouse=True) +def _cached_access_token(): + """Serve a fake token from the handler's credential cache so no auth round-trip runs.""" + cache = vertex_gemma_chat_completion._credentials_project_mapping + key = (_FAKE_CREDENTIALS, "test") + cache[key] = (SimpleNamespace(token="fake-token", expired=False), "test") + yield + cache.pop(key, None) + + +def test_sync_gemma_stream(): + captured: dict[str, Any] = {} + + def handle(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_vertex_response()) + + stream = litellm.completion( + model="vertex_ai/gemma/test-model", + messages=_MESSAGES, + stream=True, + api_base=_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_CREDENTIALS, + client=httpx.Client(transport=httpx.MockTransport(handle)), + ) + + assert isinstance(stream, CustomStreamWrapper) + chunks = list(stream) + + assert "stream" not in captured["body"]["instances"][0] + assert len(chunks) == 2 + assert chunks[0].choices[0].delta.content == "READY" + assert chunks[1].choices[0].finish_reason == "stop" + + +@pytest.mark.asyncio +async def test_async_gemma_responses_stream(): + captured: dict[str, Any] = {} + + def handle(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_vertex_response()) + + response = await litellm.aresponses( + model="vertex_ai/gemma/test-model", + input="Reply exactly READY", + stream=True, + api_base=_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_CREDENTIALS, + client=httpx.AsyncClient(transport=httpx.MockTransport(handle)), + ) + events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] + + assert "stream" not in captured["body"]["instances"][0] + assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) + assert isinstance(events[-1], ResponseCompletedEvent) + assert events[-1].response.usage is not None + assert events[-1].response.usage.total_tokens == 15 diff --git a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py index e5ca31833ce..97f4f290958 100644 --- a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -5,11 +5,18 @@ Maps to: litellm/llms/vertex_ai/vertex_gemma_models/transformation.py """ import json +from collections.abc import AsyncIterator +from typing import cast from unittest.mock import AsyncMock, Mock, patch import pytest import litellm +from litellm.types.llms.openai import ( + OutputTextDeltaEvent, + ResponseCompletedEvent, + ResponsesAPIStreamingResponse, +) @pytest.fixture(autouse=True) @@ -439,8 +446,9 @@ class TestVertexGemmaCompletion: Verifies: 1. Request body does NOT include 'stream' parameter (model doesn't support it) - 2. Response returns a MockResponseIterator that yields chunks + 2. Response wraps a MockResponseIterator and yields chunks """ + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.llms.base_llm.base_model_iterator import MockResponseIterator # Mock Vertex response @@ -502,8 +510,8 @@ class TestVertexGemmaCompletion: vertex_location="us-central1", ) - # Verify the response is a MockResponseIterator - assert isinstance(response, MockResponseIterator), f"Expected MockResponseIterator, got {type(response)}" + assert isinstance(response, CustomStreamWrapper) + assert isinstance(response.completion_stream, MockResponseIterator) # Verify the request sent to Vertex does NOT include 'stream' call_args = mock_client.post.call_args @@ -520,8 +528,9 @@ class TestVertexGemmaCompletion: async for chunk in response: chunks.append(chunk) - # Should get exactly one chunk (fake streaming) - assert len(chunks) == 1, f"Expected 1 chunk from fake stream, got {len(chunks)}" + assert len(chunks) == 2 + assert chunks[1].choices[0].finish_reason == "stop" + assert all(getattr(chunk, "usage", None) is None for chunk in chunks) # Verify the chunk has the expected content chunk = chunks[0] @@ -529,6 +538,104 @@ class TestVertexGemmaCompletion: assert len(chunk.choices) > 0 assert chunk.choices[0].delta.content == "Streaming test response" + @pytest.mark.asyncio + async def test_aresponses_streams_vertex_gemma_with_llm_tracing(self): + pytest.importorskip("ddtrace") + from ddtrace.contrib.internal.litellm.patch import patch as patch_litellm + from ddtrace.contrib.internal.litellm.patch import unpatch as unpatch_litellm + from ddtrace.llmobs._integrations.base_stream_handler import TracedAsyncStream + + from litellm.responses.litellm_completion_transformation.streaming_iterator import ( + LiteLLMCompletionStreamingIterator, + ) + + reply = Mock(status_code=200) + reply.json.return_value = _make_gemma_vertex_response(content="READY") + client = Mock() + client.post = AsyncMock(return_value=reply) + + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=client), + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "test-project"), + ), + ): + patch_litellm() + try: + response = await litellm.aresponses( + model="vertex_ai/gemma/test-model", + input="Reply exactly READY", + stream=True, + api_base="https://example.invalid/v1/projects/test-project/locations/us-central1/endpoints/test:predict", + vertex_project="test-project", + vertex_location="us-central1", + ) + bridge = cast(LiteLLMCompletionStreamingIterator, response) + traced_stream = bridge.litellm_custom_stream_wrapper + assert isinstance(traced_stream, TracedAsyncStream) + events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] + span = traced_stream.handler.primary_span + assert span.finished + assert span.get_tag("_dd.llmobs.span_kind") == "llm" + assert span.get_metric("_dd.llmobs.total_tokens") == 114 + finally: + unpatch_litellm() + + assert "stream" not in client.post.call_args.kwargs["json"]["instances"][0] + assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) + assert isinstance(events[-1], ResponseCompletedEvent) + assert events[-1].response.usage.total_tokens == 114 + + @pytest.mark.asyncio + @pytest.mark.parametrize("stream_options", [None, {"include_usage": False}, {"include_usage": True}]) + async def test_acompletion_stream_respects_usage_option_with_llm_tracing(self, stream_options): + pytest.importorskip("ddtrace") + from ddtrace.contrib.internal.litellm.patch import patch as patch_litellm + from ddtrace.contrib.internal.litellm.patch import unpatch as unpatch_litellm + + reply = Mock(status_code=200) + reply.json.return_value = _make_gemma_vertex_response(content="READY") + client = Mock(post=AsyncMock(return_value=reply)) + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=client), + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "test-project"), + ), + ): + patch_litellm() + try: + stream = await litellm.acompletion( + model="vertex_ai/gemma/test-model", + messages=[{"role": "user", "content": "Reply exactly READY"}], + stream=True, + **({"stream_options": stream_options} if stream_options is not None else {}), + api_base="https://example.invalid/v1/projects/test-project/locations/us-central1/endpoints/test:predict", + vertex_project="test-project", + vertex_location="us-central1", + ) + chunks = [chunk async for chunk in stream] + span = stream.handler.primary_span + assert span.finished + assert span.get_tag("_dd.llmobs.span_kind") == "llm" + finally: + unpatch_litellm() + + assert len(chunks) == (3 if stream_options and stream_options["include_usage"] else 2) + assert chunks[0].choices[0].delta.content == "READY" + assert chunks[1].choices[0].finish_reason == "stop" + if stream_options and stream_options["include_usage"]: + assert chunks[-1].choices[0].delta.content is None + assert chunks[-1].usage.total_tokens == 114 + assert span.get_metric("_dd.llmobs.total_tokens") == 114 + else: + from litellm.litellm_core_utils.streaming_handler import calculate_total_usage + + assert all(getattr(chunk, "usage", None) is None for chunk in chunks) + assert calculate_total_usage(chunks=stream.chunks).total_tokens == 114 + assert span.get_metric("_dd.llmobs.total_tokens") is None + @pytest.mark.asyncio async def test_acompletion_filters_stream_and_stream_options(self): """ From 491d454826342aa8b53aa69edd0242ba3b6f8b4d Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 04:51:56 +0000 Subject: [PATCH 4/4] fix(responses): emit the reasoning item on streaming /v1/responses for signature-only thinking (#43414) * fix(responses): emit the reasoning item on streaming /v1/responses for signature-only thinking Anthropic models return thinking blocks with empty text and the reasoning carried in the signature: Claude Fable 5.1 and Claude Opus 5.5 by default, and Bedrock adaptive thinking with or without an effort. On streaming /v1/responses the chat->Responses bridge opened a reasoning output item only on reasoning_content text (LiteLLMCompletionStreamingIterator._ensure_output_item_for_chunk), and ChunkProcessor.get_combined_thinking_content kept an assembled thinking block only when it had thinking text. Such a response emitted no reasoning item mid-stream and none in response.completed, so a streaming Responses client could not replay the reasoning even though the reasoning tokens were billed. Non-streaming /v1/responses was unaffected. Open the reasoning item when the delta carries a signed or redacted thinking block, and keep a signed block through stream assembly even when its thinking text is empty. Unsigned text-only fragments are still dropped. The reasoning-text path is unchanged. (cherry picked from commit bc9b6f8a5c3ac9a2b46e3f9f01f7c2c5f9b688e7) * test(vertex_ai): move orphaned gemma streaming tests into the llm-vertex-ai shard PR #43147 left a copy of the Gemma streaming tests under tests/test_litellm/llms, a tree no CI shard claims, which broke assert-ci-coverage and assert-shard-coverage on main. Fold the two streaming tests into the existing tests/unit/llms/vertex_ai file so the llm-vertex-ai shard runs them Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Chloe Lu Co-authored-by: Krrish Dholakia Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../streaming_chunk_builder_utils.py | 2 +- .../streaming_iterator.py | 7 +- .../test_vertex_gemma_transformation.py | 93 ------------------- .../test_streaming_chunk_builder_utils.py | 25 +++++ .../test_vertex_gemma_transformation.py | 78 ++++++++++++++++ .../test_streaming_iterator_transformation.py | 40 ++++++++ 6 files changed, 150 insertions(+), 95 deletions(-) delete mode 100644 tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index d975c3551f3..67684a230e3 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -685,7 +685,7 @@ class ChunkProcessor: def _flush_thinking_block() -> None: nonlocal current_thinking_text_parts, current_signature - if len(current_thinking_text_parts) > 0 and current_signature: + if current_signature: thinking_blocks.append( ChatCompletionThinkingBlock( type="thinking", diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index 5173cd04a89..21a33c17ab8 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -73,6 +73,11 @@ def _output_items_with_id(items: tuple[Any, ...], item_type: str, item_id: str | ) +def _delta_has_signed_thinking_block(delta: object) -> bool: + blocks: Final = getattr(delta, "thinking_blocks", None) or () + return any(isinstance(b, dict) and (b.get("signature") or b.get("data")) for b in blocks) + + class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): """ Async iterator for processing streaming responses from the Responses API. @@ -936,7 +941,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self.sent_output_item_added_event = True # Reasoning-first - if hasattr(delta, "reasoning_content") and delta.reasoning_content: + if (hasattr(delta, "reasoning_content") and delta.reasoning_content) or _delta_has_signed_thinking_block(delta): self._reasoning_active = True if self._cached_reasoning_item_id is None: self._cached_reasoning_item_id = f"rs_{uuid.uuid4()}" diff --git a/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py deleted file mode 100644 index 294d26b2e58..00000000000 --- a/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ /dev/null @@ -1,93 +0,0 @@ -import json -from collections.abc import AsyncIterator -from types import SimpleNamespace -from typing import Any, cast - -import httpx -import pytest - -import litellm -from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper -from litellm.main import vertex_gemma_chat_completion -from litellm.types.llms.openai import OutputTextDeltaEvent, ResponseCompletedEvent, ResponsesAPIStreamingResponse - -_VERTEX_URL = "https://example.invalid/v1/projects/test/locations/us-central1/endpoints/test:predict" -_MESSAGES = [{"role": "user", "content": "Reply exactly READY"}] -_FAKE_CREDENTIALS = "gemma-test-credentials" - - -def _vertex_response(): - return { - "predictions": { - "id": "chatcmpl-stream-test", - "created": 1759863903, - "model": "google/gemma-3-12b-it", - "object": "chat.completion", - "choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "READY"}}], - "usage": {"prompt_tokens": 14, "completion_tokens": 1, "total_tokens": 15}, - } - } - - -@pytest.fixture(autouse=True) -def _cached_access_token(): - """Serve a fake token from the handler's credential cache so no auth round-trip runs.""" - cache = vertex_gemma_chat_completion._credentials_project_mapping - key = (_FAKE_CREDENTIALS, "test") - cache[key] = (SimpleNamespace(token="fake-token", expired=False), "test") - yield - cache.pop(key, None) - - -def test_sync_gemma_stream(): - captured: dict[str, Any] = {} - - def handle(request: httpx.Request) -> httpx.Response: - captured["body"] = json.loads(request.content) - return httpx.Response(200, json=_vertex_response()) - - stream = litellm.completion( - model="vertex_ai/gemma/test-model", - messages=_MESSAGES, - stream=True, - api_base=_VERTEX_URL, - vertex_project="test", - vertex_location="us-central1", - vertex_credentials=_FAKE_CREDENTIALS, - client=httpx.Client(transport=httpx.MockTransport(handle)), - ) - - assert isinstance(stream, CustomStreamWrapper) - chunks = list(stream) - - assert "stream" not in captured["body"]["instances"][0] - assert len(chunks) == 2 - assert chunks[0].choices[0].delta.content == "READY" - assert chunks[1].choices[0].finish_reason == "stop" - - -@pytest.mark.asyncio -async def test_async_gemma_responses_stream(): - captured: dict[str, Any] = {} - - def handle(request: httpx.Request) -> httpx.Response: - captured["body"] = json.loads(request.content) - return httpx.Response(200, json=_vertex_response()) - - response = await litellm.aresponses( - model="vertex_ai/gemma/test-model", - input="Reply exactly READY", - stream=True, - api_base=_VERTEX_URL, - vertex_project="test", - vertex_location="us-central1", - vertex_credentials=_FAKE_CREDENTIALS, - client=httpx.AsyncClient(transport=httpx.MockTransport(handle)), - ) - events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] - - assert "stream" not in captured["body"]["instances"][0] - assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) - assert isinstance(events[-1], ResponseCompletedEvent) - assert events[-1].response.usage is not None - assert events[-1].response.usage.total_tokens == 15 diff --git a/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py b/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py index af763da2d87..aaf877df364 100644 --- a/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py +++ b/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py @@ -236,6 +236,31 @@ def test_get_combined_thinking_content_preserves_interleaved_blocks(): assert result[2]["signature"] == "sig_block2" +def test_get_combined_thinking_content_keeps_signed_block_without_thinking_text(): + chunks: Final = [ + ModelResponseStream( + id="chatcmpl-123", + object="chat.completion.chunk", + created=1234567890, + model="claude-sonnet-4-20250514", + choices=[ + StreamingChoices( + index=0, + delta=Delta(thinking_blocks=[{"type": "thinking", "thinking": "", "signature": "sig_only"}]), + finish_reason=None, + ) + ], + ) + ] + + result: Final = ChunkProcessor(chunks=chunks).get_combined_thinking_content(chunks) + + assert result is not None + assert [(block["type"], block["thinking"], block["signature"]) for block in result] == [ + ("thinking", "", "sig_only") + ] + + def test_cache_read_input_tokens_retained(): chunk1 = ModelResponseStream( id="chatcmpl-95aabb85-c39f-443d-ae96-0370c404d70c", diff --git a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py index 97f4f290958..92e684e42c8 100644 --- a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -1303,3 +1303,81 @@ class TestVertexGemmaCompletion: mock_async_post.assert_awaited_once() assert mock_async_post.call_args.kwargs["client"] is None assert response.choices[0].message.content == "default async handler fallback" + + +_GEMMA_VERTEX_URL = "https://example.invalid/v1/projects/test/locations/us-central1/endpoints/test:predict" +_FAKE_GEMMA_CREDENTIALS = "gemma-test-credentials" + + +@pytest.fixture +def _gemma_cached_access_token(): + """Serve a fake token from the handler's credential cache so no auth round-trip runs.""" + from types import SimpleNamespace + + from litellm.main import vertex_gemma_chat_completion + + cache = vertex_gemma_chat_completion._credentials_project_mapping + key = (_FAKE_GEMMA_CREDENTIALS, "test") + cache[key] = (SimpleNamespace(token="fake-token", expired=False), "test") + yield + cache.pop(key, None) + + +def test_sync_gemma_stream(_gemma_cached_access_token): + import httpx + + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper + + captured = {} + + def handle(request): + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_make_gemma_vertex_response(content="READY", total_tokens=15)) + + stream = litellm.completion( + model="vertex_ai/gemma/test-model", + messages=[{"role": "user", "content": "Reply exactly READY"}], + stream=True, + api_base=_GEMMA_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_GEMMA_CREDENTIALS, + client=httpx.Client(transport=httpx.MockTransport(handle)), + ) + + assert isinstance(stream, CustomStreamWrapper) + chunks = list(stream) + + assert "stream" not in captured["body"]["instances"][0] + assert len(chunks) == 2 + assert chunks[0].choices[0].delta.content == "READY" + assert chunks[1].choices[0].finish_reason == "stop" + + +@pytest.mark.asyncio +async def test_async_gemma_responses_stream(_gemma_cached_access_token): + import httpx + + captured = {} + + def handle(request): + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_make_gemma_vertex_response(content="READY", total_tokens=15)) + + response = await litellm.aresponses( + model="vertex_ai/gemma/test-model", + input="Reply exactly READY", + stream=True, + api_base=_GEMMA_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_GEMMA_CREDENTIALS, + client=httpx.AsyncClient(transport=httpx.MockTransport(handle)), + ) + events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] + + assert "stream" not in captured["body"]["instances"][0] + assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) + assert isinstance(events[-1], ResponseCompletedEvent) + assert events[-1].response.usage is not None + assert events[-1].response.usage.total_tokens == 15 diff --git a/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py b/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py index 8fbba0dbf87..041bcf1b6d7 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py +++ b/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py @@ -978,6 +978,25 @@ def _reasoning_chunk(reasoning: str, finish_reason: str | None = None) -> ModelR ) +def _signature_only_thinking_chunk(signature: str) -> ModelResponseStream: + return ModelResponseStream( + id=CHAT_COMPLETION_ID, + created=1748575031, + model="claude-haiku-4-5", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + index=0, + delta=Delta( + role="assistant", + thinking_blocks=[{"type": "thinking", "thinking": "", "signature": signature}], + ), + finish_reason=None, + ) + ], + ) + + async def _collect_events( iterator: LiteLLMCompletionStreamingIterator, sync_mode: bool ) -> list[BaseLiteLLMOpenAIResponseObject]: @@ -1015,6 +1034,27 @@ async def test_tool_only_stream_emits_no_message_item_events(sync_mode: bool): assert any(getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED for event in events) +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_signature_only_thinking_streams_a_replayable_reasoning_item(sync_mode: bool): + iterator: Final = _build_iterator([_signature_only_thinking_chunk("sig_only"), _chunk("4", finish_reason="stop")]) + + events: Final = await _collect_events(iterator, sync_mode) + + added_item_types: Final = [ + event.item.type + for event in events + if getattr(event, "type", None) == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + ] + completed: Final = next( + event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ) + reasoning_items: Final = [item for item in completed.response.output if getattr(item, "type", None) == "reasoning"] + assert added_item_types[0] == "reasoning" + assert len(reasoning_items) == 1 + assert json.loads(reasoning_items[0].encrypted_content)[0]["signature"] == "sig_only" + + @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio async def test_reasoning_then_text_announces_message_item_before_text_events(sync_mode: bool):