chore: sync main to pick up the CI test discovery fix

This commit is contained in:
Tin Chi Lo 2026-09-26 22:04:08 -07:00
commit ba0e4e2d23
20 changed files with 1169 additions and 350 deletions

View file

@ -57,7 +57,7 @@
"mcp-servers-2025-12-04": null,
"output-128k-2025-02-19": null,
"structured-output-2024-03-01": null,
"per-turn-control-2026-07-01": null,
"per-turn-control-2026-07-01": "per-turn-control-2026-07-01",
"prompt-caching-scope-2026-01-05": "prompt-caching-scope-2026-01-05",
"skills-2025-10-02": "skills-2025-10-02",
"structured-outputs-2025-11-13": "structured-outputs-2025-11-13",

View file

@ -1992,11 +1992,14 @@ def is_encrypted_reasoning_block(block: object) -> bool:
def is_unsignable_thinking_block(block: object) -> bool:
"""A thinking block Anthropic cannot accept on input.
Anthropic verifies the thinking signature cryptographically, so a block whose
signature is null, empty, or missing (e.g. from an open-source reasoning model)
is rejected with a 400 and must be dropped rather than blanked or repaired, and
so is a block whose signature or data carries another provider's encrypted
reasoning. A `redacted_thinking` block Anthropic minted is always kept.
Anthropic verifies the signature cryptographically, so a block with a null,
empty, or missing signature (e.g. from an open-source reasoning model) is
rejected with a 400, and so is a block whose signature or data carries
another provider's encrypted reasoning. It also rejects a `thinking` block
whose text is empty or whitespace-only ("each thinking block must contain
thinking"), regardless of signature, e.g. when a `thinking_blocks` history
item from a non-Anthropic reasoning provider is replayed through this path.
`redacted_thinking` blocks carry no signature and are always kept.
"""
if is_encrypted_reasoning_block(block):
return True
@ -2006,7 +2009,10 @@ def is_unsignable_thinking_block(block: object) -> bool:
if mapping.get("type") != "thinking":
return False
signature: Final = mapping.get("signature")
return not (isinstance(signature, str) and len(signature) > 0)
if not (isinstance(signature, str) and len(signature) > 0):
return True
thinking_text: Final = mapping.get("thinking")
return not (isinstance(thinking_text, str) and len(thinking_text.strip()) > 0)
def strip_encrypted_reasoning_from_messages(messages: object) -> None:
@ -2497,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.

View file

@ -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 "<unknown>",
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 []

View file

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

View file

@ -96,6 +96,7 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
"logprobs",
"frequency_penalty",
"presence_penalty",
"seed",
"modalities",
"parallel_tool_calls",
"web_search_options",

View file

@ -322,6 +322,7 @@ class ContextCachingEndpoints(VertexBase):
if not is_prompt_caching_valid_prompt(
model=model,
messages=cached_messages,
tools=optional_params.get("tools"),
custom_llm_provider=custom_llm_provider,
):
verbose_logger.debug(
@ -481,6 +482,7 @@ class ContextCachingEndpoints(VertexBase):
if not is_prompt_caching_valid_prompt(
model=model,
messages=cached_messages,
tools=optional_params.get("tools"),
custom_llm_provider=custom_llm_provider,
):
verbose_logger.debug(

View file

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

View file

@ -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()}"

View file

@ -134,7 +134,7 @@ reporting failures as test errors. Already deleted files and batches that are
terminal are safe to clean up again. Managed batch cancellation polls for up to two minutes
before input deletion. A managed batch still `cancelling` after that is left for the provider to
finish, and its input file is left in place because LiteLLM refuses to delete a file a non-terminal
batch references. Both are reported as `BatchCleanupLeftover` warnings naming their ids rather than
batch references. Both are reported as `UserWarning`s naming their ids rather than
failing the test. Any other status or error still fails
Accepted cancellation may still report validating or in_progress while the provider
updates its state. Raw and model-encoded batches are polled until cancelling or

View file

@ -28,10 +28,6 @@ class BatchCleanupClient(Protocol):
def cancel_batch(self, batch_id: str, *, key: str, provider: str | None = None) -> Result[BatchObject]: ...
class BatchCleanupLeftover(UserWarning):
pass
def cleanup_result[R: BaseModel](
action: Callable[[], Result[R]], *, wait: Callable[[float], None] = sleep
) -> Result[R]:
@ -68,7 +64,7 @@ def cleanup_file(client: BatchCleanupClient, file_id: str, *, key: str, provider
if isinstance(result, UnknownApiError) and result.status_code == 400 and FILE_IN_USE_REFUSAL in result.body:
warnings.warn(
f"Left file {file_id} in place: LiteLLM refused to delete it while a batch still references it",
BatchCleanupLeftover,
UserWarning,
stacklevel=2,
)
return
@ -140,7 +136,7 @@ def cleanup_batch(
)
warnings.warn(
f"Left batch {batch_id} cancelling after {BATCH_CANCEL_TIMEOUT_SECONDS}s for the provider to finish",
BatchCleanupLeftover,
UserWarning,
stacklevel=2,
)
return

View file

@ -7,7 +7,6 @@ import pytest
from batch_cleanup import (
BATCH_CANCEL_TIMEOUT_SECONDS,
CLEANUP_DELAYS,
BatchCleanupLeftover,
cleanup_batch,
cleanup_file,
cleanup_result,
@ -141,7 +140,7 @@ class TestFileCleanup:
calls=ExpectedCalls((f"delete None {MANAGED_FILE_ID}",)),
files=(UnknownApiError(status_code=400, body=IN_USE_REFUSAL),),
)
with pytest.warns(BatchCleanupLeftover, match=MANAGED_FILE_ID):
with pytest.warns(UserWarning, match=MANAGED_FILE_ID):
cleanup_file(client, MANAGED_FILE_ID, key="test-key")
client.calls.assert_done()
@ -241,7 +240,7 @@ class TestBatchCancellation:
key: Final = manager.key()
manager.defer(lambda: cleanup_file(client, MANAGED_FILE_ID, key=key))
manager.defer(lambda: cleanup_batch(client, MANAGED_BATCH_ID, key=key, clock=ticks))
with pytest.warns(BatchCleanupLeftover) as leftovers:
with pytest.warns(UserWarning, match="^Left ") as leftovers:
manager.teardown()
client.calls.assert_done()
messages: Final = tuple(str(warning.message) for warning in leftovers)

View file

@ -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"}]

View file

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

View file

@ -95,13 +95,19 @@ def test_added_per_turn_control_beta_survives_the_anthropic_allowlist():
assert PER_TURN_CONTROL in _betas(filtered)
@pytest.mark.parametrize("provider", ["bedrock", "bedrock_converse", "vertex_ai", "azure_ai", "databricks"])
@pytest.mark.parametrize("provider", ["bedrock", "bedrock_converse", "vertex_ai", "databricks"])
def test_per_turn_control_beta_is_dropped_for_providers_without_it(provider):
filtered = update_headers_with_filtered_beta(headers={"anthropic-beta": PER_TURN_CONTROL}, provider=provider)
assert "anthropic-beta" not in filtered
def test_per_turn_control_beta_is_forwarded_for_azure_ai():
filtered = update_headers_with_filtered_beta(headers={"anthropic-beta": PER_TURN_CONTROL}, provider="azure_ai")
assert _betas(filtered) == {PER_TURN_CONTROL}
def test_json_provider_passthrough_adds_per_turn_control_beta():
config = JSONProviderAnthropicMessagesConfig(
SimpleProviderConfig(

View file

@ -1503,6 +1503,129 @@ class TestContextCachingEndpoints:
# Restart the patcher so teardown_method can stop it cleanly
self._token_check_patcher.start()
@pytest.mark.parametrize("is_async", [False, True])
@pytest.mark.parametrize(
"custom_llm_provider", ["gemini", "vertex_ai"]
)
@pytest.mark.asyncio
async def test_check_and_create_cache_considers_tools_for_min_tokens(
self, custom_llm_provider, is_async
):
"""Test that context caching accounts for tools when validating minimum token count.
Fixes #42804: When messages alone are below the threshold, but tools push the total
over the minimum token count, context caching must proceed and include tools.
"""
self._token_check_patcher.stop()
short_cached_messages = [
{
"role": "system",
"content": "Short system instruction.",
"cache_control": {"type": "ephemeral"},
}
]
non_cached_messages = [
{"role": "user", "content": "Hello world"},
]
all_messages = short_cached_messages + non_cached_messages
large_tools = [
{
"type": "function",
"function": {
"name": f"synthetic_tool_{i}",
"description": "A very descriptive explanation of a synthetic tool designed to add tokens to the prompt cache prefix " * 8,
"parameters": {
"type": "object",
"properties": {
f"arg_{j}": {"type": "string", "description": "Argument description for caching verification " * 4}
for j in range(10)
},
"required": [f"arg_{j}" for j in range(5)],
},
},
}
for i in range(12)
]
optional_params = {
**self.sample_optional_params,
"tools": large_tools,
}
mock_response = MagicMock()
mock_response.json.return_value = {
"name": "cachedContents/test_cache_id",
"model": "gemini-1.5-pro",
}
mock_response.status_code = 200
self.mock_client.post.return_value = mock_response
self.mock_async_client.post = AsyncMock(return_value=mock_response)
with patch.object(
self.context_caching,
"_get_token_and_url_context_caching",
return_value=("fake_token", "https://fake.url/cachedContents"),
), patch.object(
self.context_caching,
"check_cache",
return_value=None,
), patch.object(
self.context_caching,
"async_check_cache",
new_callable=AsyncMock,
return_value=None,
):
if is_async:
result = await self.context_caching.async_check_and_create_cache(
messages=all_messages,
optional_params=optional_params,
api_key="test_key",
api_base=None,
model="gemini-1.5-pro",
client=self.mock_async_client,
timeout=30.0,
logging_obj=self.mock_logging,
cached_content=None,
custom_llm_provider=custom_llm_provider,
vertex_project="test_project",
vertex_location="us-central1",
vertex_auth_header="test_token",
)
else:
result = self.context_caching.check_and_create_cache(
messages=all_messages,
optional_params=optional_params,
api_key="test_key",
api_base=None,
model="gemini-1.5-pro",
client=self.mock_client,
timeout=30.0,
logging_obj=self.mock_logging,
cached_content=None,
custom_llm_provider=custom_llm_provider,
vertex_project="test_project",
vertex_location="us-central1",
vertex_auth_header="test_token",
)
messages, returned_params, returned_cache = result
assert messages == non_cached_messages
assert returned_cache == "cachedContents/test_cache_id"
assert "tools" not in returned_params
post_mock = self.mock_async_client.post if is_async else self.mock_client.post
post_mock.assert_called_once()
call_kwargs = post_mock.call_args.kwargs
assert call_kwargs["json"]["tools"] == large_tools
assert call_kwargs["json"]["contents"] == [
{"role": "user", "parts": [{"text": "Short system instruction."}]}
]
self._token_check_patcher.start()
def _model_turn_final_messages(self, final_cached_role):
tool_call = {
"id": "call_abc123",

View file

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

View file

@ -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):
"""
@ -1196,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"))
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"))
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 == 114

View file

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

View file

@ -603,7 +603,7 @@ async def test_jev_provider_credentials_never_leave_their_own_environment_pair(
def respond(sent: httpx.Request) -> httpx.Response:
assert str(sent.url) == f"{base}/v1/systemone"
assert sent.headers.get("authorization") == (f"Bearer {key}" if key else None)
assert json.loads(sent.content) == request.model_dump(mode="json")
assert TypeAdapter(Mapping[str, object]).validate_json(sent.content) == request.model_dump(mode="json")
return httpx.Response(200, json=answer.model_dump(mode="json"))
handler: Final = AsyncHTTPHandler(transport=httpx.MockTransport(respond))