mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
chore: sync main to pick up the CI test discovery fix
This commit is contained in:
commit
ba0e4e2d23
20 changed files with 1169 additions and 350 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -96,6 +96,7 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
|
|||
"logprobs",
|
||||
"frequency_penalty",
|
||||
"presence_penalty",
|
||||
"seed",
|
||||
"modalities",
|
||||
"parallel_tool_calls",
|
||||
"web_search_options",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"}]
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue