mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_lit_4913_headroom_streaming_ccr
This commit is contained in:
commit
fb08fc9574
74 changed files with 2642 additions and 525 deletions
|
|
@ -6,10 +6,10 @@
|
|||
"limit": 2564
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 320
|
||||
"limit": 319
|
||||
},
|
||||
"reportAttributeAccessIssue": {
|
||||
"limit": 483
|
||||
"limit": 480
|
||||
},
|
||||
"reportCallIssue": {
|
||||
"limit": 113
|
||||
|
|
@ -30,7 +30,7 @@
|
|||
"limit": 7
|
||||
},
|
||||
"reportGeneralTypeIssues": {
|
||||
"limit": 154
|
||||
"limit": 105
|
||||
},
|
||||
"reportIncompatibleMethodOverride": {
|
||||
"limit": 56
|
||||
|
|
@ -99,19 +99,19 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 44528
|
||||
"limit": 44526
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38804
|
||||
"limit": 38782
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19829
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 30355
|
||||
"limit": 30349
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 117
|
||||
|
|
@ -123,7 +123,7 @@
|
|||
"limit": 5
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 833
|
||||
"limit": 831
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 0
|
||||
|
|
|
|||
|
|
@ -73,6 +73,11 @@ ARRAY_KEYS: dict[str, JsonSchema] = {
|
|||
"description": "Output modalities the model can produce.",
|
||||
"items": {"type": "string", "enum": ["text", "image", "audio", "video", "code"]},
|
||||
},
|
||||
"reasoning_effort_levels": {
|
||||
"type": "array",
|
||||
"description": "Exact reasoning_effort levels this deployment accepts; wins over supports_* flags.",
|
||||
"items": {"type": "string", "enum": ["none", "minimal", "low", "medium", "high", "xhigh", "max"]},
|
||||
},
|
||||
"supported_regions": {
|
||||
"type": "array",
|
||||
"description": "Cloud regions the model is available in ('global' or region ids).",
|
||||
|
|
|
|||
|
|
@ -59,9 +59,11 @@ if TYPE_CHECKING:
|
|||
from litellm.types.llms.openai import (
|
||||
ALL_RESPONSES_API_TOOL_PARAMS,
|
||||
AllMessageValues,
|
||||
ChatCompletionFileObject,
|
||||
ChatCompletionImageObject,
|
||||
ChatCompletionRedactedThinkingBlock,
|
||||
ChatCompletionThinkingBlock,
|
||||
ChatCompletionToolReferenceObject,
|
||||
OpenAIMessageContentListBlock,
|
||||
)
|
||||
from litellm.types.utils import Choices
|
||||
|
|
@ -175,6 +177,16 @@ def _map_incomplete_reason_to_finish_reason(incomplete_reason: str | None) -> Li
|
|||
return "length"
|
||||
|
||||
|
||||
def _input_file_from_file_value(file_value: object) -> dict[str, object]:
|
||||
if not isinstance(file_value, dict):
|
||||
return {"type": "input_file"}
|
||||
file_dict: Final = cast("dict[str, object]", file_value) # cast-ok: runtime dict checked
|
||||
return {
|
||||
"type": "input_file",
|
||||
**{key: file_dict[key] for key in ("file_id", "file_data", "filename") if key in file_dict},
|
||||
}
|
||||
|
||||
|
||||
def _incomplete_reason_from_response_payload(response_payload: object) -> str | None:
|
||||
if not isinstance(response_payload, Mapping):
|
||||
return None
|
||||
|
|
@ -957,7 +969,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
content: str
|
||||
| list[object]
|
||||
| Iterable[
|
||||
Union["OpenAIMessageContentListBlock", "ChatCompletionThinkingBlock", "ChatCompletionRedactedThinkingBlock"]
|
||||
Union[
|
||||
"OpenAIMessageContentListBlock",
|
||||
"ChatCompletionThinkingBlock",
|
||||
"ChatCompletionRedactedThinkingBlock",
|
||||
"ChatCompletionToolReferenceObject",
|
||||
]
|
||||
]
|
||||
| None,
|
||||
role: str,
|
||||
|
|
@ -1006,17 +1023,15 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
result.append(converted)
|
||||
verbose_logger.debug("Chat provider: image -> %s", converted)
|
||||
elif item_type == "file":
|
||||
# Map Chat Completion file to Responses API input_file
|
||||
# {"type": "file", "file": {"file_data": "...", "filename": "..."}}
|
||||
# -> {"type": "input_file", "file_data": "...", "filename": "..."}
|
||||
file_data = item.get("file", {})
|
||||
converted = {"type": "input_file"}
|
||||
if isinstance(file_data, dict):
|
||||
for key in ["file_id", "file_data", "filename"]:
|
||||
if key in file_data:
|
||||
converted[key] = file_data[key]
|
||||
converted = _input_file_from_file_value(
|
||||
cast("ChatCompletionFileObject", item).get("file"), # cast-ok: type tag checked
|
||||
)
|
||||
result.append(converted)
|
||||
verbose_logger.debug("Chat provider: file -> %s", converted)
|
||||
elif item_type == "tool_reference":
|
||||
verbose_logger.debug(
|
||||
"Chat provider: tool_reference has no responses API equivalent; skipped"
|
||||
)
|
||||
elif item_type in [
|
||||
"input_text",
|
||||
"input_image",
|
||||
|
|
|
|||
|
|
@ -76,7 +76,10 @@ from litellm.llms.perplexity.cost_calculator import (
|
|||
from litellm.llms.tencent.cost_calculator import (
|
||||
cost_per_token as tencent_cost_per_token,
|
||||
)
|
||||
from litellm.llms.together_ai.cost_calculator import get_model_params_and_category
|
||||
from litellm.llms.together_ai.cost_calculator import (
|
||||
get_model_params_and_category,
|
||||
has_together_registry_pricing,
|
||||
)
|
||||
from litellm.llms.vertex_ai.cost_calculator import (
|
||||
cost_per_character as google_cost_per_character,
|
||||
)
|
||||
|
|
@ -1569,10 +1572,9 @@ def completion_cost(
|
|||
|
||||
return MCPCostCalculator.calculate_mcp_tool_call_cost(litellm_logging_obj=litellm_logging_obj)
|
||||
# Calculate cost based on prompt_tokens, completion_tokens
|
||||
if "togethercomputer" in model or "together_ai" in model or custom_llm_provider == "together_ai":
|
||||
# together ai prices based on size of llm
|
||||
# get_model_params_and_category takes a model name and returns the category of LLM size it is in model_prices_and_context_window.json
|
||||
|
||||
if (
|
||||
"togethercomputer" in model or "together_ai" in model or custom_llm_provider == "together_ai"
|
||||
) and not has_together_registry_pricing(model, litellm.model_cost):
|
||||
model = get_model_params_and_category(model, call_type=CallTypes(call_type))
|
||||
|
||||
# replicate llms are calculate based on time for request running
|
||||
|
|
|
|||
|
|
@ -376,7 +376,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
# 2. list of objects - only apply to last item per Anthropic spec
|
||||
elif isinstance(message_content, list):
|
||||
if len(message_content) > 0 and isinstance(message_content[-1], dict):
|
||||
message_content[-1]["cache_control"] = control
|
||||
message_content[-1]["cache_control"] = control # pyright: ignore[reportGeneralTypeIssues] # loose runtime dict
|
||||
return message
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
193
litellm/litellm_core_utils/audio_utils/subtitle_utils.py
Normal file
193
litellm/litellm_core_utils/audio_utils/subtitle_utils.py
Normal file
|
|
@ -0,0 +1,193 @@
|
|||
"""Provider-agnostic SRT/WebVTT subtitle synthesis from timestamped transcription tokens."""
|
||||
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
from itertools import accumulate, chain
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
|
||||
CUE_MAX_TOKENS: Final = 15
|
||||
CUE_MAX_DURATION_MS: Final = 5000
|
||||
|
||||
SRT_RESPONSE_FORMAT: Final = "srt"
|
||||
VTT_RESPONSE_FORMAT: Final = "vtt"
|
||||
SUBTITLE_RESPONSE_FORMATS: Final = frozenset((SRT_RESPONSE_FORMAT, VTT_RESPONSE_FORMAT))
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SubtitleToken:
|
||||
text: str
|
||||
start_ms: int | None = None
|
||||
end_ms: int | None = None
|
||||
speaker: str | int | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SubtitleCue:
|
||||
start_ms: int
|
||||
end_ms: int
|
||||
text: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _CueAccumulator:
|
||||
texts: tuple[str, ...] = ()
|
||||
start_ms: int | None = None
|
||||
end_ms: int | None = None
|
||||
speaker: str | int | None = None
|
||||
|
||||
|
||||
def _completed_cue(accumulator: _CueAccumulator) -> tuple[SubtitleCue, ...]:
|
||||
if not accumulator.texts or accumulator.start_ms is None:
|
||||
return ()
|
||||
text: Final = "".join(accumulator.texts).strip()
|
||||
if not text:
|
||||
return ()
|
||||
end_ms: Final = accumulator.end_ms if accumulator.end_ms is not None else accumulator.start_ms
|
||||
return (SubtitleCue(start_ms=accumulator.start_ms, end_ms=end_ms, text=text),)
|
||||
|
||||
|
||||
def _cue_break_reached(accumulator: _CueAccumulator, token: SubtitleToken) -> bool:
|
||||
if len(accumulator.texts) >= CUE_MAX_TOKENS:
|
||||
return True
|
||||
return (
|
||||
accumulator.start_ms is not None
|
||||
and token.start_ms is not None
|
||||
and token.start_ms - accumulator.start_ms >= CUE_MAX_DURATION_MS
|
||||
)
|
||||
|
||||
|
||||
_AbsorbStep = tuple[tuple[SubtitleCue, ...], _CueAccumulator]
|
||||
|
||||
|
||||
def _absorb_token(accumulator: _CueAccumulator, token: SubtitleToken) -> _AbsorbStep:
|
||||
if token.start_ms is None and accumulator.start_ms is None:
|
||||
return (), accumulator
|
||||
if token.speaker is not None and token.speaker != accumulator.speaker:
|
||||
return _completed_cue(accumulator), _CueAccumulator(
|
||||
texts=(token.text,),
|
||||
start_ms=token.start_ms,
|
||||
end_ms=token.end_ms,
|
||||
speaker=token.speaker,
|
||||
)
|
||||
if _cue_break_reached(accumulator, token):
|
||||
return _completed_cue(accumulator), _CueAccumulator(
|
||||
texts=(token.text,),
|
||||
start_ms=token.start_ms,
|
||||
end_ms=token.end_ms,
|
||||
speaker=accumulator.speaker,
|
||||
)
|
||||
return (), _CueAccumulator(
|
||||
texts=(*accumulator.texts, token.text),
|
||||
start_ms=accumulator.start_ms if accumulator.start_ms is not None else token.start_ms,
|
||||
end_ms=token.end_ms if token.end_ms is not None else accumulator.end_ms,
|
||||
speaker=accumulator.speaker,
|
||||
)
|
||||
|
||||
|
||||
def _absorb_step(carry: _AbsorbStep, token: SubtitleToken) -> _AbsorbStep:
|
||||
return _absorb_token(carry[1], token)
|
||||
|
||||
|
||||
def group_subtitle_tokens_into_cues(tokens: Sequence[SubtitleToken]) -> tuple[SubtitleCue, ...]:
|
||||
steps: Final = tuple(accumulate(tokens, _absorb_step, initial=((), _CueAccumulator())))
|
||||
completed: Final = chain.from_iterable(emitted for emitted, _ in steps)
|
||||
return (*completed, *_completed_cue(steps[-1][1]))
|
||||
|
||||
|
||||
def _format_timestamp(total_ms: int, millis_separator: str) -> str:
|
||||
clamped: Final = max(total_ms, 0)
|
||||
hours, hour_remainder = divmod(clamped, 3_600_000)
|
||||
minutes, minute_remainder = divmod(hour_remainder, 60_000)
|
||||
seconds, millis = divmod(minute_remainder, 1_000)
|
||||
return f"{hours:02d}:{minutes:02d}:{seconds:02d}{millis_separator}{millis:03d}"
|
||||
|
||||
|
||||
def _render_srt(cues: Sequence[SubtitleCue]) -> str:
|
||||
lines: Final = tuple(
|
||||
line
|
||||
for index, cue in enumerate(cues, start=1)
|
||||
for line in (
|
||||
str(index),
|
||||
f"{_format_timestamp(cue.start_ms, ',')} --> {_format_timestamp(cue.end_ms, ',')}",
|
||||
cue.text,
|
||||
"",
|
||||
)
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _render_vtt(cues: Sequence[SubtitleCue]) -> str:
|
||||
cue_lines: Final = tuple(
|
||||
line
|
||||
for cue in cues
|
||||
for line in (
|
||||
f"{_format_timestamp(cue.start_ms, '.')} --> {_format_timestamp(cue.end_ms, '.')}",
|
||||
cue.text,
|
||||
"",
|
||||
)
|
||||
)
|
||||
return "\n".join(("WEBVTT", "", *cue_lines))
|
||||
|
||||
|
||||
def render_subtitle_tokens_as_srt(tokens: Sequence[SubtitleToken]) -> str:
|
||||
"""Render tokens as an SRT document; empty string when no token has timestamp data."""
|
||||
cues: Final = group_subtitle_tokens_into_cues(tokens)
|
||||
if not cues:
|
||||
return ""
|
||||
return _render_srt(cues)
|
||||
|
||||
|
||||
def render_subtitle_tokens_as_vtt(tokens: Sequence[SubtitleToken]) -> str:
|
||||
"""Render tokens as a WebVTT document; the WEBVTT header is emitted even without cues."""
|
||||
return _render_vtt(group_subtitle_tokens_into_cues(tokens))
|
||||
|
||||
|
||||
class TranscriptionWordTiming(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
word: str = ""
|
||||
start: float | None = None
|
||||
end: float | None = None
|
||||
speaker: str | None = None
|
||||
|
||||
|
||||
_WORD_TIMINGS_ADAPTER: Final = TypeAdapter(tuple[TranscriptionWordTiming, ...])
|
||||
|
||||
|
||||
def _seconds_to_ms(seconds: float | None) -> int | None:
|
||||
if seconds is None:
|
||||
return None
|
||||
return round(seconds * 1000)
|
||||
|
||||
|
||||
def _word_to_subtitle_token(word: TranscriptionWordTiming) -> SubtitleToken:
|
||||
return SubtitleToken(
|
||||
text=f"{word.word} ",
|
||||
start_ms=_seconds_to_ms(word.start),
|
||||
end_ms=_seconds_to_ms(word.end),
|
||||
speaker=word.speaker,
|
||||
)
|
||||
|
||||
|
||||
def _parse_word_timings(words: object) -> tuple[TranscriptionWordTiming, ...]:
|
||||
try:
|
||||
return _WORD_TIMINGS_ADAPTER.validate_python(words)
|
||||
except ValidationError:
|
||||
return ()
|
||||
|
||||
|
||||
def synthesize_subtitle_document(words: object, response_format: str) -> str | None:
|
||||
"""
|
||||
Build an SRT/VTT document from OpenAI verbose_json-style word dicts
|
||||
(word/start/end in float seconds, optional speaker). Returns None when the
|
||||
format is not a subtitle format or the words carry no usable timestamps.
|
||||
"""
|
||||
if response_format not in SUBTITLE_RESPONSE_FORMATS:
|
||||
return None
|
||||
tokens: Final = tuple(_word_to_subtitle_token(word) for word in _parse_word_timings(words))
|
||||
cues: Final = group_subtitle_tokens_into_cues(tokens)
|
||||
if not cues:
|
||||
return None
|
||||
return _render_srt(cues) if response_format == SRT_RESPONSE_FORMAT else _render_vtt(cues)
|
||||
|
|
@ -1747,6 +1747,46 @@ def hoist_images_from_tool_messages(
|
|||
]
|
||||
|
||||
|
||||
def _is_tool_reference_part(part: object) -> bool:
|
||||
return isinstance(part, dict) and part.get("type") == "tool_reference"
|
||||
|
||||
|
||||
def _tool_message_carries_tool_reference(message: AllMessageValues) -> bool:
|
||||
if message.get("role") != "tool":
|
||||
return False
|
||||
content = message.get("content")
|
||||
return isinstance(content, list) and any(_is_tool_reference_part(part) for part in content)
|
||||
|
||||
|
||||
def _drop_tool_reference_parts(message: AllMessageValues) -> AllMessageValues:
|
||||
if not _tool_message_carries_tool_reference(message):
|
||||
return message
|
||||
content = cast(list, message.get("content")) # cast-ok: shape checked by _tool_message_carries_tool_reference
|
||||
remaining_parts = [ # mutable-ok: tool message content must stay a json list
|
||||
part for part in content if not _is_tool_reference_part(part)
|
||||
]
|
||||
new_content = remaining_parts if remaining_parts else ""
|
||||
rewritten = {**message, "content": new_content} # mutable-ok: chat messages are plain json dicts
|
||||
return cast(AllMessageValues, rewritten) # cast-ok: dict spread keeps keys like cache_control
|
||||
|
||||
|
||||
def drop_tool_reference_parts_from_tool_messages(
|
||||
messages: list[AllMessageValues], # mutable-ok: message pipelines type messages as mutable lists
|
||||
) -> list[AllMessageValues]: # mutable-ok: message pipelines type messages as mutable lists
|
||||
"""
|
||||
Remove tool_reference content parts from role:"tool" messages.
|
||||
|
||||
The OpenAI chat spec only accepts text in tool messages, so a tool_reference
|
||||
part carried through the Anthropic adapter makes strict providers reject the
|
||||
request. The reference names an already-declared tool rather than carrying
|
||||
content, so it is dropped; a reference-only result keeps its tool message with
|
||||
empty text so the preceding tool_call stays answered.
|
||||
"""
|
||||
if not any(_tool_message_carries_tool_reference(message) for message in messages):
|
||||
return messages
|
||||
return [_drop_tool_reference_parts(message) for message in messages] # mutable-ok: pipelines mutate message lists
|
||||
|
||||
|
||||
def _attempt_json_repair(s: str) -> Any | None:
|
||||
"""
|
||||
Attempt to repair truncated JSON produced by LLM tool calls.
|
||||
|
|
|
|||
|
|
@ -1412,7 +1412,7 @@ def convert_to_gemini_tool_call_result(
|
|||
)
|
||||
except Exception as e:
|
||||
verbose_logger.warning("Failed to process image in tool response: %s", e)
|
||||
elif content_type in ("file", "input_file"):
|
||||
elif content_type in ("file", "input_file"): # pyright: ignore[reportUnnecessaryContains] # loose runtime dict
|
||||
# Extract file for inline_data (for tool results with PDF, audio, video, etc.)
|
||||
file_data = content.get("file_data", "")
|
||||
if not file_data:
|
||||
|
|
@ -1564,14 +1564,23 @@ def convert_to_anthropic_tool_result(
|
|||
}
|
||||
"""
|
||||
anthropic_content: (
|
||||
str | list[AnthropicMessagesToolResultContent | AnthropicMessagesImageParam | AnthropicMessagesDocumentParam]
|
||||
str
|
||||
| list[
|
||||
AnthropicMessagesToolResultContent
|
||||
| AnthropicMessagesImageParam
|
||||
| AnthropicMessagesDocumentParam
|
||||
| ToolReference
|
||||
]
|
||||
) = ""
|
||||
if isinstance(message["content"], str):
|
||||
anthropic_content = message["content"]
|
||||
elif isinstance(message["content"], list):
|
||||
content_list: Final = message["content"]
|
||||
anthropic_content_list: list[
|
||||
AnthropicMessagesToolResultContent | AnthropicMessagesImageParam | AnthropicMessagesDocumentParam
|
||||
AnthropicMessagesToolResultContent
|
||||
| AnthropicMessagesImageParam
|
||||
| AnthropicMessagesDocumentParam
|
||||
| ToolReference
|
||||
] = []
|
||||
for content in content_list:
|
||||
if content["type"] == "text":
|
||||
|
|
@ -1614,6 +1623,8 @@ def convert_to_anthropic_tool_result(
|
|||
original_content_element=content,
|
||||
)
|
||||
anthropic_content_list.append(cast(AnthropicMessagesImageParam, _anthropic_image_param))
|
||||
elif content["type"] == "tool_reference":
|
||||
anthropic_content_list.append(ToolReference(type="tool_reference", tool_name=content["tool_name"]))
|
||||
elif content["type"] == "file":
|
||||
file_content = cast(ChatCompletionFileObject, content)
|
||||
_file_block = anthropic_process_openai_file_message(file_content)
|
||||
|
|
|
|||
|
|
@ -8,11 +8,12 @@ import time
|
|||
import traceback
|
||||
from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, NoReturn, Protocol, TypeVar, cast
|
||||
|
||||
import anyio
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, ValidationError
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
|
||||
import litellm
|
||||
|
|
@ -182,6 +183,23 @@ class _VertexChunkLike(Protocol):
|
|||
candidates: Sequence[_VertexCandidateLike]
|
||||
|
||||
|
||||
class _ParsedChunkHiddenParams(BaseModel):
|
||||
provider_specific_fields: Mapping[str, object] | None = None
|
||||
|
||||
|
||||
def _provider_hidden_params(chunk: object) -> Mapping[str, object] | None:
|
||||
hidden: Final[object] = getattr(chunk, "_hidden_params", None)
|
||||
if not isinstance(hidden, dict):
|
||||
return None
|
||||
try:
|
||||
parsed: Final = _ParsedChunkHiddenParams.model_validate(hidden)
|
||||
except ValidationError:
|
||||
return None
|
||||
if not parsed.provider_specific_fields:
|
||||
return None
|
||||
return MappingProxyType({"provider_specific_fields": dict(parsed.provider_specific_fields)})
|
||||
|
||||
|
||||
class CustomStreamWrapper:
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -801,7 +819,7 @@ class CustomStreamWrapper:
|
|||
except Exception as e:
|
||||
raise e
|
||||
|
||||
def model_response_creator(self, chunk: dict | None = None, hidden_params: dict | None = None):
|
||||
def model_response_creator(self, chunk: dict | None = None, hidden_params: Mapping[str, object] | None = None):
|
||||
_model: Final = self._cached_model_name
|
||||
_logging_obj_llm_provider: Final = self._cached_logging_llm_provider
|
||||
|
||||
|
|
@ -1504,7 +1522,7 @@ class CustomStreamWrapper:
|
|||
def chunk_creator(self, chunk: Any):
|
||||
if hasattr(chunk, "id"):
|
||||
self.response_id = chunk.id
|
||||
model_response = self.model_response_creator()
|
||||
model_response = self.model_response_creator(hidden_params=_provider_hidden_params(chunk))
|
||||
response_obj: dict[str, Any] = {}
|
||||
try:
|
||||
# return this for all models
|
||||
|
|
|
|||
|
|
@ -24,10 +24,12 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
|
||||
LiteLLMAnthropicMessagesAdapter,
|
||||
is_provider_native_tool_dict,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
anthropic_tool_name,
|
||||
anthropic_tool_names,
|
||||
effective_scan_only_tool_results_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
effective_skip_tool_message_for_guardrail,
|
||||
|
|
@ -360,7 +362,13 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
structured_messages: Final = [full_structured_messages[index] for index in scoped_message_indices]
|
||||
|
||||
tools_to_check: Final[list[ChatCompletionToolParam]] = (
|
||||
[] if scan_only_tool_results else chat_completion_compatible_request.get("tools", [])
|
||||
[]
|
||||
if scan_only_tool_results
|
||||
else [
|
||||
tool
|
||||
for tool in chat_completion_compatible_request.get("tools", [])
|
||||
if not is_provider_native_tool_dict(tool)
|
||||
]
|
||||
)
|
||||
|
||||
# Step 1: Extract all text content and images
|
||||
|
|
@ -419,7 +427,10 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
tool_name=anthropic_tool_name,
|
||||
)
|
||||
if scan_only_tool_results
|
||||
else anthropic_tools
|
||||
else [
|
||||
*(tool for tool in data.get("tools") or [] if is_provider_native_tool_dict(tool)),
|
||||
*anthropic_tools,
|
||||
]
|
||||
)
|
||||
|
||||
guardrailed_structured_messages: Final = guardrailed_inputs.get("structured_messages")
|
||||
|
|
@ -677,12 +688,9 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
)
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> list[str]:
|
||||
"""Extract tool names from Anthropic messages request (tools[].name)."""
|
||||
names: Final[list[str]] = []
|
||||
for tool in data.get("tools") or []:
|
||||
if isinstance(tool, dict) and tool.get("name"):
|
||||
names.append(str(tool["name"]))
|
||||
return names
|
||||
"""Extract every tool name in an Anthropic messages request: tools[].name, plus
|
||||
tools[].function.name for OpenAI-format tools the bridge forwards verbatim."""
|
||||
return [name for tool in data.get("tools") or [] for name in anthropic_tool_names(tool)]
|
||||
|
||||
@classmethod
|
||||
def _extract_input_text_and_images(
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
import copy
|
||||
import hashlib
|
||||
import json
|
||||
from collections.abc import AsyncIterator, Iterator, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, cast
|
||||
from collections.abc import AsyncIterator, Iterator, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, TypeVar, cast
|
||||
|
||||
import litellm
|
||||
from litellm.llms.anthropic.experimental_pass_through.utils import (
|
||||
|
|
@ -18,6 +18,22 @@ TOOL_NAME_PREFIX_LENGTH: Final = OPENAI_MAX_TOOL_NAME_LENGTH - TOOL_NAME_HASH_LE
|
|||
PROVIDERS_PROXYING_AN_UNKNOWN_BACKEND: Final = frozenset({"litellm_proxy"})
|
||||
|
||||
|
||||
_ANTHROPIC_TOOL_SCHEMA_KEYS: Final = frozenset(
|
||||
{"name", "type", "input_schema", "description", "cache_control", "strict"}
|
||||
)
|
||||
|
||||
|
||||
def _is_openai_function_tool(tool: Mapping[str, object]) -> bool:
|
||||
return tool.get("type") == "function" and "function" in tool
|
||||
|
||||
|
||||
def is_provider_native_tool_dict(tool: Mapping[str, object]) -> bool:
|
||||
if len(tool) != 1:
|
||||
return False
|
||||
key, value = next(iter(tool.items()))
|
||||
return key not in _ANTHROPIC_TOOL_SCHEMA_KEYS and isinstance(value, dict)
|
||||
|
||||
|
||||
def truncate_tool_name(name: str) -> str:
|
||||
"""
|
||||
Truncate tool names that exceed OpenAI's 64-character limit.
|
||||
|
|
@ -126,7 +142,9 @@ from litellm.types.llms.openai import (
|
|||
ChatCompletionToolMessage,
|
||||
ChatCompletionToolParam,
|
||||
ChatCompletionToolParamFunctionChunk,
|
||||
ChatCompletionToolReferenceObject,
|
||||
ChatCompletionUserMessage,
|
||||
ToolMessageContentPart,
|
||||
)
|
||||
from litellm.types.utils import Choices, ModelResponse, StreamingChoices, Usage
|
||||
|
||||
|
|
@ -135,6 +153,8 @@ from .streaming_iterator import AnthropicStreamWrapper
|
|||
if TYPE_CHECKING:
|
||||
from litellm.types.llms.anthropic import ContentBlockContentBlockDict
|
||||
|
||||
ToolResultContent: TypeAlias = str | list[ToolMessageContentPart]
|
||||
|
||||
|
||||
class AnthropicAdapter:
|
||||
def __init__(self) -> None:
|
||||
|
|
@ -412,90 +432,13 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
self._add_cache_control_if_applicable(content, doc_obj, model)
|
||||
new_user_content_list.append(doc_obj)
|
||||
elif content.get("type") == "tool_result":
|
||||
if "content" not in content:
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content="",
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
elif isinstance(content.get("content"), str):
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content=str(content.get("content", "")),
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
elif isinstance(content.get("content"), list):
|
||||
# Combine all content items into a single tool message
|
||||
# to avoid creating multiple tool_result blocks with the same ID
|
||||
# (each tool_use must have exactly one tool_result)
|
||||
content_items = list(content.get("content", []))
|
||||
|
||||
# Single-item text keeps the backward-compatible string format; a single
|
||||
# image or document becomes a structured image_url part
|
||||
if len(content_items) == 1:
|
||||
c = content_items[0]
|
||||
if isinstance(c, str):
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content=c,
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
elif isinstance(c, dict):
|
||||
if c.get("type") == "text":
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content=c.get("text", ""),
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
elif c.get("type") in ("image", "document"):
|
||||
image_part = self._tool_result_image_part(c.get("source"))
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content=[image_part] # mutable-ok: content must be a json list
|
||||
if image_part
|
||||
else "",
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
else:
|
||||
# For multiple content items, combine into a single tool message
|
||||
# with list content to preserve all items while having one tool_use_id
|
||||
combined_content_parts: list[
|
||||
ChatCompletionTextObject | ChatCompletionImageObject
|
||||
] = []
|
||||
for c in content_items:
|
||||
if isinstance(c, str):
|
||||
combined_content_parts.append(ChatCompletionTextObject(type="text", text=c))
|
||||
elif isinstance(c, dict):
|
||||
if c.get("type") == "text":
|
||||
combined_content_parts.append(
|
||||
ChatCompletionTextObject(
|
||||
type="text",
|
||||
text=c.get("text", ""),
|
||||
)
|
||||
)
|
||||
elif c.get("type") in ("image", "document"):
|
||||
image_part = self._tool_result_image_part(c.get("source"))
|
||||
if image_part:
|
||||
combined_content_parts.append(image_part)
|
||||
# Create a single tool message with combined content
|
||||
if combined_content_parts:
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content=combined_content_parts,
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content=self._tool_result_content(content.get("content")),
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
|
||||
if len(tool_message_list) > 0:
|
||||
new_messages.extend(tool_message_list)
|
||||
|
|
@ -771,6 +714,10 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
new_tools.append(tool)
|
||||
continue
|
||||
|
||||
if _is_openai_function_tool(tool) or is_provider_native_tool_dict(tool):
|
||||
new_tools.append(cast(ChatCompletionToolParam, tool)) # cast-ok: passed through verbatim to provider
|
||||
continue
|
||||
|
||||
raw_name = tool.get("name")
|
||||
if raw_name is None or (isinstance(raw_name, str) and not str(raw_name).strip()):
|
||||
original_name = f"litellm_unnamed_tool_{idx}"
|
||||
|
|
@ -1210,6 +1157,39 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
|
||||
return None
|
||||
|
||||
def _tool_result_content(self, raw_content: object) -> ToolResultContent:
|
||||
if isinstance(raw_content, str):
|
||||
return raw_content
|
||||
if not isinstance(raw_content, list):
|
||||
return ""
|
||||
items: Final = cast(Sequence[object], raw_content) # cast-ok: untrusted client payload
|
||||
parts: Final = tuple(part for part in (self._tool_result_part(item) for item in items) if part is not None)
|
||||
match parts:
|
||||
case ():
|
||||
return ""
|
||||
case ({"type": "text", "text": str(text)},):
|
||||
return text
|
||||
case _:
|
||||
return list(parts) # mutable-ok: content must be a json list
|
||||
|
||||
def _tool_result_part(self, item: object) -> ToolMessageContentPart | None:
|
||||
if isinstance(item, str):
|
||||
return ChatCompletionTextObject(type="text", text=item)
|
||||
if not isinstance(item, dict):
|
||||
return None
|
||||
block: Final = cast(Mapping[str, object], item) # cast-ok: untrusted client payload
|
||||
match block.get("type"):
|
||||
case "text":
|
||||
return ChatCompletionTextObject(type="text", text=str(block.get("text") or ""))
|
||||
case "image" | "document":
|
||||
return self._tool_result_image_part(block.get("source"))
|
||||
case "tool_reference":
|
||||
return ChatCompletionToolReferenceObject(
|
||||
type="tool_reference", tool_name=str(block.get("tool_name") or "")
|
||||
)
|
||||
case _:
|
||||
return None
|
||||
|
||||
def _tool_result_image_part(self, image_source: object) -> ChatCompletionImageObject | None:
|
||||
if not isinstance(image_source, dict):
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
import os
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
|
|
@ -23,6 +25,29 @@ def is_reasoning_auto_summary_enabled() -> bool:
|
|||
return litellm.reasoning_auto_summary or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true"
|
||||
|
||||
|
||||
_DECLARED_DEGRADATION_CHAINS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType(
|
||||
{"max": ("max", "xhigh", "high"), "xhigh": ("xhigh", "high"), "minimal": ("minimal", "low")}
|
||||
)
|
||||
|
||||
|
||||
def _effort_from_declaration(model_info: ModelInfo, effort: str) -> str | None:
|
||||
"""A declared level set is the WHOLE answer for this gate, so a level it omits degrades even
|
||||
where a per-level flag would have allowed it. Honoring both would let /model_group/info and
|
||||
this path disagree about the same entry. None means the entry declares nothing, and the flag
|
||||
chain below decides as before.
|
||||
|
||||
A declaration that omits every level in a chain still lands on that chain's terminal, which can
|
||||
itself be undeclared. Picking a nearer declared level instead would need a strength ordering,
|
||||
and the advertisement order is presentation only by design, so the terminal stays the answer."""
|
||||
from litellm.router_utils.reasoning_effort_capability import declared_reasoning_efforts
|
||||
|
||||
declared: Final = declared_reasoning_efforts(model_info)
|
||||
if declared is None:
|
||||
return None
|
||||
chain: Final = _DECLARED_DEGRADATION_CHAINS[effort]
|
||||
return next((level for level in chain if level in declared), chain[-1])
|
||||
|
||||
|
||||
def normalize_reasoning_effort_value(
|
||||
effort: str,
|
||||
model: str,
|
||||
|
|
@ -48,6 +73,10 @@ def normalize_reasoning_effort_value(
|
|||
except Exception:
|
||||
model_info = None
|
||||
|
||||
declared_effort: Final = _effort_from_declaration(model_info, effort) if model_info is not None else None
|
||||
if declared_effort is not None:
|
||||
return declared_effort
|
||||
|
||||
if effort == "max":
|
||||
if model_info and model_info.get("supports_max_reasoning_effort"):
|
||||
return "max"
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ from httpx._models import Headers, Response
|
|||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
hoist_images_from_tool_messages,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
|
|
@ -252,7 +253,8 @@ class AzureOpenAIConfig(BaseConfig):
|
|||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
azure_messages: Final = convert_to_azure_openai_messages(hoist_images_from_tool_messages(messages))
|
||||
stripped_messages: Final = drop_tool_reference_parts_from_tool_messages(messages)
|
||||
azure_messages: Final = convert_to_azure_openai_messages(hoist_images_from_tool_messages(stripped_messages))
|
||||
return {
|
||||
"model": model,
|
||||
"messages": azure_messages,
|
||||
|
|
|
|||
|
|
@ -40,6 +40,16 @@ class BaseAudioTranscriptionConfig(BaseConfig, ABC):
|
|||
def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]:
|
||||
pass
|
||||
|
||||
@property
|
||||
def supports_subtitle_synthesis(self) -> bool:
|
||||
"""
|
||||
Opt-in for providers without a native srt/vtt response body: when True
|
||||
and the user asked for response_format srt/vtt, the http handler
|
||||
synthesizes the subtitle document from the word timestamps the
|
||||
provider's TranscriptionResponse carries in `words`.
|
||||
"""
|
||||
return False
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
|
|
|
|||
|
|
@ -209,9 +209,20 @@ def openai_tool_name(tool: object) -> str | None:
|
|||
return flat_name if isinstance(flat_name, str) else None
|
||||
|
||||
|
||||
def anthropic_tool_names(tool: object) -> tuple[str, ...]:
|
||||
"""Every name a /v1/messages tool dict can act under: the flat Anthropic ``name`` plus
|
||||
``function.name`` for OpenAI-format tools the bridge forwards verbatim. Allowlist checks
|
||||
must see both, or a decoy flat name could smuggle a disallowed ``function.name`` through."""
|
||||
if not isinstance(tool, dict):
|
||||
return ()
|
||||
function: Final = tool.get("function") if tool.get("type") == "function" else None
|
||||
function_name: Final = function.get("name") if isinstance(function, dict) else None
|
||||
return tuple(name for name in (tool.get("name"), function_name) if isinstance(name, str) and name)
|
||||
|
||||
|
||||
def anthropic_tool_name(tool: object) -> str | None:
|
||||
name: Final = tool.get("name") if isinstance(tool, dict) else None
|
||||
return name if isinstance(name, str) else None
|
||||
names: Final = anthropic_tool_names(tool)
|
||||
return names[0] if names else None
|
||||
|
||||
|
||||
def merge_returned_tools_into_request_tools(
|
||||
|
|
|
|||
|
|
@ -25,6 +25,10 @@ from litellm.litellm_core_utils.agentic_loop_settings import (
|
|||
validated_max_agentic_loops,
|
||||
)
|
||||
from litellm.litellm_core_utils.asyncify import run_async_function
|
||||
from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
|
||||
SUBTITLE_RESPONSE_FORMATS,
|
||||
synthesize_subtitle_document,
|
||||
)
|
||||
from litellm.litellm_core_utils.llm_request_utils import serialize_multipart_form_fields
|
||||
from litellm.litellm_core_utils.realtime_errors import realtime_error_event, websocket_close_reason
|
||||
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
|
||||
|
|
@ -1296,9 +1300,23 @@ class BaseLLMHTTPHandler:
|
|||
api_key: str | None,
|
||||
) -> TranscriptionResponse:
|
||||
"""Shared logic for transforming audio transcription responses."""
|
||||
return provider_config.transform_audio_transcription_response(
|
||||
transformed: Final = provider_config.transform_audio_transcription_response(
|
||||
raw_response=response,
|
||||
)
|
||||
if not provider_config.supports_subtitle_synthesis:
|
||||
return transformed
|
||||
requested_format: Final = optional_params.get("response_format")
|
||||
if not isinstance(requested_format, str) or requested_format not in SUBTITLE_RESPONSE_FORMATS:
|
||||
return transformed
|
||||
document: Final = synthesize_subtitle_document(
|
||||
words=transformed.get("words"),
|
||||
response_format=requested_format,
|
||||
)
|
||||
if document is not None:
|
||||
transformed.text = document
|
||||
if "words" in transformed:
|
||||
delattr(transformed, "words")
|
||||
return transformed
|
||||
|
||||
def audio_transcriptions(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ Request format:
|
|||
"input": {
|
||||
"messages": [{"role": "user", "content": [{"text": "<prompt>"}]}]
|
||||
},
|
||||
"parameters": {"size": "1024*1024", ...}
|
||||
"parameters": {"size": "1024*1024", "n": 1, ...}
|
||||
}
|
||||
|
||||
Response format:
|
||||
|
|
@ -19,7 +19,7 @@ Response format:
|
|||
"output": {
|
||||
"choices": [{"message": {"content": [{"image": "<url>"}]}}]
|
||||
},
|
||||
"usage": {"input_tokens": 0, "output_tokens": 0, "width": 1024, "height": 1024, "image_count": 1}
|
||||
"usage": {"output_width": 1024, "output_height": 1024, "output_image_count": 1}
|
||||
}
|
||||
"""
|
||||
|
||||
|
|
@ -46,6 +46,8 @@ else:
|
|||
|
||||
DEFAULT_API_BASE: Final = "https://dashscope-intl.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation"
|
||||
|
||||
CHAT_COMPATIBLE_MODE_PATH: Final = "/compatible-mode/v1"
|
||||
|
||||
# Maps OpenAI size strings (WxH) to DashScope size strings (W*H)
|
||||
OPENAI_TO_DASHSCOPE_SIZE: Final[dict] = {
|
||||
"256x256": "256*256",
|
||||
|
|
@ -59,7 +61,8 @@ OPENAI_TO_DASHSCOPE_SIZE: Final[dict] = {
|
|||
|
||||
class DashScopeImageGenerationConfig(BaseImageGenerationConfig):
|
||||
"""
|
||||
Configuration for DashScope image generation (qwen-image-2.0, qwen-image-2.0-pro).
|
||||
Configuration for DashScope image generation (qwen-image-2.0, qwen-image-2.0-pro,
|
||||
qwen-image-3.0, qwen-image-3.0-pro).
|
||||
"""
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]:
|
||||
|
|
@ -82,8 +85,8 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig):
|
|||
if k == "size":
|
||||
# Convert "WxH" → "W*H"
|
||||
mapped["size"] = OPENAI_TO_DASHSCOPE_SIZE.get(v, v.replace("x", "*"))
|
||||
elif k == "n":
|
||||
mapped["image_count"] = v
|
||||
else:
|
||||
mapped[k] = v
|
||||
return mapped
|
||||
|
||||
def get_complete_url(
|
||||
|
|
@ -95,7 +98,10 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig):
|
|||
litellm_params: dict,
|
||||
stream: bool | None = None,
|
||||
) -> str:
|
||||
return api_base or get_secret_str("DASHSCOPE_API_BASE_IMAGE") or DEFAULT_API_BASE
|
||||
image_api_base: Final = (
|
||||
api_base if api_base and not api_base.rstrip("/").endswith(CHAT_COMPATIBLE_MODE_PATH) else None
|
||||
)
|
||||
return image_api_base or get_secret_str("DASHSCOPE_API_BASE_IMAGE") or DEFAULT_API_BASE
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ from typing import Final
|
|||
|
||||
from httpx import Headers, Response
|
||||
|
||||
from litellm.litellm_core_utils.audio_utils.subtitle_utils import SUBTITLE_RESPONSE_FORMATS
|
||||
from litellm.litellm_core_utils.audio_utils.utils import (
|
||||
normalize_transcription_language_to_bcp47,
|
||||
process_audio_file,
|
||||
|
|
@ -48,6 +49,10 @@ class GeminiAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
|
|||
) -> list[OpenAIAudioTranscriptionOptionalParams]: # mutable-ok: BaseAudioTranscriptionConfig signature
|
||||
return ["language", "response_format", "timestamp_granularities"] # mutable-ok: base contract returns a list
|
||||
|
||||
@property
|
||||
def supports_subtitle_synthesis(self) -> bool:
|
||||
return True
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: Mapping[str, object],
|
||||
|
|
@ -215,16 +220,17 @@ def _language_config(language: object) -> GeminiTranscriptionConfig:
|
|||
return language_config
|
||||
|
||||
|
||||
def _timestamp_config(timestamp_granularities: object) -> GeminiTranscriptionConfig:
|
||||
if isinstance(timestamp_granularities, list) and "word" in timestamp_granularities:
|
||||
return _WORD_TIMESTAMP_CONFIG
|
||||
return _EMPTY_TRANSCRIPTION_CONFIG
|
||||
def _timestamp_config(timestamp_granularities: object, response_format: object) -> GeminiTranscriptionConfig:
|
||||
wants_word_timestamps: Final = (
|
||||
isinstance(timestamp_granularities, list) and "word" in timestamp_granularities
|
||||
) or (isinstance(response_format, str) and response_format in SUBTITLE_RESPONSE_FORMATS)
|
||||
return _WORD_TIMESTAMP_CONFIG if wants_word_timestamps else _EMPTY_TRANSCRIPTION_CONFIG
|
||||
|
||||
|
||||
def _build_transcription_config(optional_params: Mapping[str, object]) -> GeminiTranscriptionConfig:
|
||||
transcription_config: Final[GeminiTranscriptionConfig] = {
|
||||
**_language_config(optional_params.get("language")),
|
||||
**_timestamp_config(optional_params.get("timestamp_granularities")),
|
||||
**_timestamp_config(optional_params.get("timestamp_granularities"), optional_params.get("response_format")),
|
||||
}
|
||||
return transcription_config
|
||||
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
|
|||
from litellm.litellm_core_utils.prompt_templates.image_handling import (
|
||||
convert_url_to_base64,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionFileObject
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionFileObject, ChatCompletionImageObject
|
||||
from litellm.types.llms.vertex_ai import ContentType, PartType
|
||||
from litellm.utils import supports_reasoning
|
||||
|
||||
|
|
@ -16,6 +16,13 @@ from ...vertex_ai.gemini.transformation import _gemini_convert_messages_with_his
|
|||
from ...vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig
|
||||
|
||||
|
||||
def _image_url_fields(img_element: ChatCompletionImageObject) -> tuple[str | None, str | None, str | None]:
|
||||
image_value: Final = img_element.get("image_url")
|
||||
if isinstance(image_value, dict):
|
||||
return image_value.get("url"), image_value.get("format"), image_value.get("detail")
|
||||
return image_value, None, None
|
||||
|
||||
|
||||
class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
|
||||
"""
|
||||
Reference: https://ai.google.dev/api/rest/v1beta/GenerationConfig
|
||||
|
|
@ -118,16 +125,8 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
|
|||
_parts: list[PartType] = []
|
||||
for element in _message_content:
|
||||
if element.get("type") == "image_url":
|
||||
img_element = element
|
||||
_image_url: str | None = None
|
||||
format: str | None = None
|
||||
detail: str | None = None
|
||||
if isinstance(img_element.get("image_url"), dict):
|
||||
_image_url = img_element["image_url"].get("url")
|
||||
format = img_element["image_url"].get("format")
|
||||
detail = img_element["image_url"].get("detail")
|
||||
else:
|
||||
_image_url = img_element.get("image_url")
|
||||
img_element = cast(ChatCompletionImageObject, element) # cast-ok: runtime type tag checked
|
||||
_image_url, format, detail = _image_url_fields(img_element)
|
||||
if _image_url and "https://" in _image_url:
|
||||
image_obj = convert_to_anthropic_image_obj(_image_url, format=format)
|
||||
converted_image_url = convert_generic_image_chunk_to_openai_image_obj(image_obj)
|
||||
|
|
|
|||
|
|
@ -292,7 +292,7 @@ class MistralConfig(OpenAIGPTConfig):
|
|||
file_id = file_content.get("file", {}).get("file_id")
|
||||
if file_id:
|
||||
# Replace 'file' with 'file_id'
|
||||
file_content["file_id"] = file_id
|
||||
file_content["file_id"] = file_id # pyright: ignore[reportGeneralTypeIssues] # legacy in-place rewrite of the block shape
|
||||
file_content.pop("file", None)
|
||||
return messages
|
||||
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
|
|||
_should_convert_tool_call_to_json_mode,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
get_tool_call_names,
|
||||
hoist_images_from_tool_messages,
|
||||
)
|
||||
|
|
@ -336,7 +337,8 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
self, messages: list[AllMessageValues], model: str, is_async: bool = False
|
||||
) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]:
|
||||
"""OpenAI no longer supports image_url as a string, so we need to convert it to a dict"""
|
||||
hoisted_messages: Final = hoist_images_from_tool_messages(messages)
|
||||
stripped_messages: Final = drop_tool_reference_parts_from_tool_messages(messages)
|
||||
hoisted_messages: Final = hoist_images_from_tool_messages(stripped_messages)
|
||||
|
||||
async def _async_transform():
|
||||
for message in hoisted_messages:
|
||||
|
|
|
|||
|
|
@ -4,6 +4,11 @@ Shared utilities for the Soniox provider (https://soniox.com).
|
|||
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
|
||||
SubtitleToken,
|
||||
render_subtitle_tokens_as_srt,
|
||||
render_subtitle_tokens_as_vtt,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
# Soniox API base URL.
|
||||
|
|
@ -109,121 +114,13 @@ def render_soniox_tokens(tokens: list[dict[str, Any]]) -> str:
|
|||
return "".join(text_parts)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SRT / VTT subtitle rendering
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Maximum number of tokens to group into a single subtitle cue.
|
||||
_CUE_MAX_TOKENS: Final[int] = 15
|
||||
|
||||
# Maximum duration (in ms) for a single cue before forcing a break.
|
||||
_CUE_MAX_DURATION_MS: Final[int] = 5000
|
||||
|
||||
|
||||
def _format_timestamp_srt(ms: int) -> str:
|
||||
"""Format milliseconds as SRT timestamp: HH:MM:SS,mmm"""
|
||||
ms = max(ms, 0)
|
||||
hours: Final = ms // 3_600_000
|
||||
ms %= 3_600_000
|
||||
minutes: Final = ms // 60_000
|
||||
ms %= 60_000
|
||||
seconds: Final = ms // 1_000
|
||||
millis: Final = ms % 1_000
|
||||
return f"{hours:02d}:{minutes:02d}:{seconds:02d},{millis:03d}"
|
||||
|
||||
|
||||
def _format_timestamp_vtt(ms: int) -> str:
|
||||
"""Format milliseconds as VTT timestamp: HH:MM:SS.mmm"""
|
||||
ms = max(ms, 0)
|
||||
hours: Final = ms // 3_600_000
|
||||
ms %= 3_600_000
|
||||
minutes: Final = ms // 60_000
|
||||
ms %= 60_000
|
||||
seconds: Final = ms // 1_000
|
||||
millis: Final = ms % 1_000
|
||||
return f"{hours:02d}:{minutes:02d}:{seconds:02d}.{millis:03d}"
|
||||
|
||||
|
||||
def _group_tokens_into_cues(
|
||||
tokens: list[dict[str, Any]],
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Group Soniox tokens into subtitle cues.
|
||||
|
||||
Each cue has:
|
||||
- start_ms: int
|
||||
- end_ms: int
|
||||
- text: str
|
||||
|
||||
Grouping heuristics:
|
||||
- A new cue starts when token count exceeds _CUE_MAX_TOKENS.
|
||||
- A new cue starts when duration exceeds _CUE_MAX_DURATION_MS.
|
||||
- A new cue starts when the speaker changes (if diarization is on).
|
||||
- Tokens without timestamps are appended to the current cue.
|
||||
"""
|
||||
cues: Final[list[dict[str, Any]]] = []
|
||||
current_tokens: list[str] = []
|
||||
current_start: int | None = None
|
||||
current_end: int | None = None
|
||||
current_speaker: Any | None = None
|
||||
|
||||
def _flush() -> None:
|
||||
if current_tokens and current_start is not None:
|
||||
text: Final = "".join(current_tokens).strip()
|
||||
if text:
|
||||
cues.append(
|
||||
{
|
||||
"start_ms": current_start,
|
||||
"end_ms": (current_end if current_end is not None else current_start),
|
||||
"text": text,
|
||||
}
|
||||
)
|
||||
|
||||
for token in tokens:
|
||||
start_ms = token.get("start_ms")
|
||||
end_ms = token.get("end_ms")
|
||||
text = token.get("text", "")
|
||||
speaker = token.get("speaker")
|
||||
|
||||
# Skip tokens with no timestamp data entirely if we have no cue started
|
||||
if start_ms is None and current_start is None:
|
||||
continue
|
||||
|
||||
# Speaker change forces a new cue
|
||||
if speaker is not None and speaker != current_speaker:
|
||||
_flush()
|
||||
current_tokens = []
|
||||
current_start = start_ms
|
||||
current_end = end_ms
|
||||
current_speaker = speaker
|
||||
current_tokens.append(text)
|
||||
continue
|
||||
|
||||
# Duration or token count exceeded -> flush
|
||||
should_break = False
|
||||
if (
|
||||
len(current_tokens) >= _CUE_MAX_TOKENS
|
||||
or current_start is not None
|
||||
and start_ms is not None
|
||||
and (start_ms - current_start) >= _CUE_MAX_DURATION_MS
|
||||
):
|
||||
should_break = True
|
||||
|
||||
if should_break:
|
||||
_flush()
|
||||
current_tokens = []
|
||||
current_start = start_ms
|
||||
current_end = end_ms
|
||||
current_tokens.append(text)
|
||||
else:
|
||||
if current_start is None:
|
||||
current_start = start_ms
|
||||
if end_ms is not None:
|
||||
current_end = end_ms
|
||||
current_tokens.append(text)
|
||||
|
||||
_flush()
|
||||
return cues
|
||||
def _soniox_token_to_subtitle_token(token: dict[str, Any]) -> SubtitleToken:
|
||||
return SubtitleToken(
|
||||
text=token.get("text", ""),
|
||||
start_ms=token.get("start_ms"),
|
||||
end_ms=token.get("end_ms"),
|
||||
speaker=token.get("speaker"),
|
||||
)
|
||||
|
||||
|
||||
def render_soniox_tokens_as_srt(tokens: list[dict[str, Any]]) -> str:
|
||||
|
|
@ -232,20 +129,7 @@ def render_soniox_tokens_as_srt(tokens: list[dict[str, Any]]) -> str:
|
|||
|
||||
Returns an empty string if no tokens have timestamp data.
|
||||
"""
|
||||
cues: Final = _group_tokens_into_cues(tokens)
|
||||
if not cues:
|
||||
return ""
|
||||
|
||||
lines: Final[list[str]] = []
|
||||
for idx, cue in enumerate(cues, start=1):
|
||||
start = _format_timestamp_srt(cue["start_ms"])
|
||||
end = _format_timestamp_srt(cue["end_ms"])
|
||||
lines.append(str(idx))
|
||||
lines.append(f"{start} --> {end}")
|
||||
lines.append(cue["text"])
|
||||
lines.append("") # blank line between cues
|
||||
|
||||
return "\n".join(lines)
|
||||
return render_subtitle_tokens_as_srt(tuple(_soniox_token_to_subtitle_token(token) for token in tokens))
|
||||
|
||||
|
||||
def render_soniox_tokens_as_vtt(tokens: list[dict[str, Any]]) -> str:
|
||||
|
|
@ -254,14 +138,4 @@ def render_soniox_tokens_as_vtt(tokens: list[dict[str, Any]]) -> str:
|
|||
|
||||
Returns the VTT header even if no cues are present.
|
||||
"""
|
||||
cues: Final = _group_tokens_into_cues(tokens)
|
||||
|
||||
lines: Final[list[str]] = ["WEBVTT", ""]
|
||||
for cue in cues:
|
||||
start = _format_timestamp_vtt(cue["start_ms"])
|
||||
end = _format_timestamp_vtt(cue["end_ms"])
|
||||
lines.append(f"{start} --> {end}")
|
||||
lines.append(cue["text"])
|
||||
lines.append("") # blank line between cues
|
||||
|
||||
return "\n".join(lines)
|
||||
return render_subtitle_tokens_as_vtt(tuple(_soniox_token_to_subtitle_token(token) for token in tokens))
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ Handles calculating cost for together ai models
|
|||
"""
|
||||
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
from litellm.constants import (
|
||||
|
|
@ -18,6 +19,12 @@ from litellm.constants import (
|
|||
from litellm.types.utils import CallTypes
|
||||
|
||||
|
||||
def has_together_registry_pricing(model: str, cost_map: Mapping[str, object]) -> bool:
|
||||
stripped: Final = model.removeprefix("together_ai/")
|
||||
entry: Final = cost_map.get(f"together_ai/{stripped}")
|
||||
return isinstance(entry, Mapping) and "input_cost_per_token" in entry
|
||||
|
||||
|
||||
# Extract the number of billion parameters from the model name
|
||||
# only used for together_computer LLMs
|
||||
def get_model_params_and_category(model_name, call_type: CallTypes) -> str:
|
||||
|
|
|
|||
|
|
@ -9315,6 +9315,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-kimi-k3-through-fireworks-ai-on-microsoft-foundry/4540187",
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -14647,6 +14652,22 @@
|
|||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"dashscope/qwen-image-3.0": {
|
||||
"litellm_provider": "dashscope",
|
||||
"mode": "image_generation",
|
||||
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"dashscope/qwen-image-3.0-pro": {
|
||||
"litellm_provider": "dashscope",
|
||||
"mode": "image_generation",
|
||||
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"databricks/databricks-bge-large-en": {
|
||||
"cache_creation_input_token_cost": 1.0003e-07,
|
||||
"cache_read_input_token_cost": 1.0003e-07,
|
||||
|
|
@ -31881,6 +31902,11 @@
|
|||
"max_tokens": 1048576,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://platform.kimi.ai/docs/pricing/chat-k3",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -36489,6 +36515,14 @@
|
|||
"litellm_provider": "perplexity",
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"minimal",
|
||||
"low",
|
||||
"medium",
|
||||
"high",
|
||||
"xhigh",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.perplexity.ai/docs/agent-api/models",
|
||||
"supports_web_search": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -38791,14 +38825,14 @@
|
|||
"supports_reasoning": true
|
||||
},
|
||||
"together_ai/Qwen/Qwen3.7-Max": {
|
||||
"cache_read_input_token_cost": 1.3e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"input_cost_per_token": 2.5e-06,
|
||||
"litellm_provider": "together_ai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3.75e-06,
|
||||
"output_cost_per_token": 7.5e-06,
|
||||
"source": "https://docs.together.ai/docs/serverless-models",
|
||||
"supports_prompt_caching": true
|
||||
},
|
||||
|
|
@ -38970,6 +39004,11 @@
|
|||
"max_tokens": 1048576,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.together.ai/docs/serverless-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
|
|
@ -39053,6 +39092,24 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"together_ai/zai-org/GLM-5.3-Flash": {
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"litellm_provider": "together_ai",
|
||||
"max_input_tokens": 1048575,
|
||||
"max_output_tokens": 1048575,
|
||||
"max_tokens": 1048575,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-07,
|
||||
"source": "https://docs.together.ai/docs/serverless-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"tts-1": {
|
||||
"input_cost_per_character": 1.5e-05,
|
||||
"litellm_provider": "openai",
|
||||
|
|
@ -51740,6 +51797,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -51804,6 +51866,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -51820,6 +51887,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.25e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -51836,6 +51908,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -52008,6 +52085,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.25e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -52024,6 +52106,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
|
|||
|
|
@ -686,6 +686,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
aws_profile_name: Final = self.optional_params.get("aws_profile_name", None)
|
||||
aws_web_identity_token: Final = self.optional_params.get("aws_web_identity_token", None)
|
||||
aws_sts_endpoint: Final = self.optional_params.get("aws_sts_endpoint", None)
|
||||
aws_external_id: Final = self.optional_params.get("aws_external_id", None)
|
||||
|
||||
### SET REGION NAME ###
|
||||
aws_region_name = self.get_aws_region_name_for_non_llm_api_calls(
|
||||
|
|
@ -702,6 +703,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
aws_role_name=aws_role_name,
|
||||
aws_web_identity_token=aws_web_identity_token,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
aws_external_id=aws_external_id,
|
||||
)
|
||||
return credentials, aws_region_name
|
||||
|
||||
|
|
|
|||
|
|
@ -34,6 +34,7 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail):
|
|||
aws_role_name=litellm_params.aws_role_name,
|
||||
aws_web_identity_token=litellm_params.aws_web_identity_token,
|
||||
aws_sts_endpoint=litellm_params.aws_sts_endpoint,
|
||||
aws_external_id=litellm_params.aws_external_id,
|
||||
aws_bedrock_runtime_endpoint=litellm_params.aws_bedrock_runtime_endpoint,
|
||||
experimental_use_latest_role_message_only=litellm_params.experimental_use_latest_role_message_only,
|
||||
only_scan_new_messages=litellm_params.only_scan_new_messages or False,
|
||||
|
|
|
|||
|
|
@ -2661,6 +2661,7 @@ class LiteLLMCompletionResponsesConfig:
|
|||
optional_output_details: Final[dict[str, int]] = {
|
||||
field: value
|
||||
for field, value in (
|
||||
("audio_tokens", getattr(completion_details, "audio_tokens", None)),
|
||||
("text_tokens", getattr(completion_details, "text_tokens", None)),
|
||||
("image_tokens", getattr(completion_details, "image_tokens", None)),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,10 +1,13 @@
|
|||
"""Resolve which reasoning_effort values a deployment, and by intersection a model group, accepts.
|
||||
|
||||
The model map's supports_*_reasoning_effort flags are the only signal, and each level's polarity
|
||||
mirrors how a request path reads that same flag. medium and high are unconditional for a reasoning
|
||||
model. minimal and low are opt-out: openai/chat/gpt_5_transformation.py refuses them only when the
|
||||
map says false. xhigh and max are opt-in. none is opt-out everywhere except the azure gpt-5 family,
|
||||
whose config raises UnsupportedParamsError without an explicit true.
|
||||
An entry that states its levels outright in reasoning_effort_levels is read first and wins
|
||||
whole, for a model whose set the per-level flags cannot express: Kimi K3 takes low, high and max,
|
||||
and no flag can drop medium because medium has none. Every other entry answers through the
|
||||
supports_*_reasoning_effort flags below, whose polarity mirrors how a request path reads that same
|
||||
flag. medium and high are unconditional for a reasoning model. minimal and low are opt-out:
|
||||
openai/chat/gpt_5_transformation.py refuses them only when the map says false. xhigh and max are
|
||||
opt-in. none is opt-out everywhere except the azure gpt-5 family, whose config raises
|
||||
UnsupportedParamsError without an explicit true.
|
||||
|
||||
xhigh is gated on the request path by the openai and azure gpt-5 configs. max is not gated there at
|
||||
all: every entry carrying supports_max_reasoning_effort is Claude-family, and
|
||||
|
|
@ -41,6 +44,7 @@ _EFFORT_FLAGS: Final = (
|
|||
("xhigh", "supports_xhigh_reasoning_effort"),
|
||||
("max", "supports_max_reasoning_effort"),
|
||||
)
|
||||
_DECLARED_EFFORTS_KEY: Final = "reasoning_effort_levels"
|
||||
_OPT_OUT_EFFORTS: Final = ("minimal", "low")
|
||||
_OPT_IN_EFFORTS: Final = ("xhigh", "max")
|
||||
_UNCONDITIONAL_EFFORTS: Final = frozenset(("medium", "high"))
|
||||
|
|
@ -69,6 +73,20 @@ def _declared_effort_flags(model_info: Mapping[str, object]) -> Mapping[str, obj
|
|||
)
|
||||
|
||||
|
||||
def declared_reasoning_efforts(model_info: Mapping[str, object]) -> tuple[str, ...] | None:
|
||||
"""The entry's own answer, read through the same bare twin as the flags so both spellings of one
|
||||
model agree. Present-and-a-list IS the answer, so a declared [] correctly empties the group and
|
||||
an unknown level is dropped rather than raised: the bundled map is enum-validated by
|
||||
validate-model-prices-json, but an operator can put this key on a config.yaml model_info block
|
||||
where that schema never runs, and one mistyped level must not fail every sibling on the proxy."""
|
||||
own: Final = model_info.get(_DECLARED_EFFORTS_KEY)
|
||||
raw: Final = own if own is not None else _bare_model_entry(model_info).get(_DECLARED_EFFORTS_KEY)
|
||||
if not isinstance(raw, Sequence) or isinstance(raw, (str, bytes)):
|
||||
return None
|
||||
declared: Final = frozenset(effort for effort in raw if isinstance(effort, str))
|
||||
return tuple(effort for effort in REASONING_EFFORT_ADVERTISEMENT_ORDER if effort in declared)
|
||||
|
||||
|
||||
def _supports_none_reasoning_effort(model_info: Mapping[str, object], flag: object) -> bool:
|
||||
"""Opt-in only where a request path refuses the level. AzureOpenAIGPT5Config raises
|
||||
UnsupportedParamsError on reasoning_effort='none' without an explicit true, and it is selected
|
||||
|
|
@ -119,6 +137,10 @@ def resolve_supported_reasoning_efforts(
|
|||
if supports_reasoning is not True:
|
||||
return () if supports_reasoning is False or deployment_is_mapped else None
|
||||
|
||||
declared: Final = declared_reasoning_efforts(model_info)
|
||||
if declared is not None:
|
||||
return declared
|
||||
|
||||
flags: Final = _declared_effort_flags(model_info)
|
||||
if all(value is None for value in flags.values()):
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -496,6 +496,9 @@ class BedrockGuardrailConfigModel(BaseModel):
|
|||
aws_role_name: str | None = Field(default=None, description="AWS role name for assuming roles")
|
||||
aws_web_identity_token: str | None = Field(default=None, description="Web identity token for AWS role assumption")
|
||||
aws_sts_endpoint: str | None = Field(default=None, description="AWS STS endpoint URL")
|
||||
aws_external_id: str | None = Field(
|
||||
default=None, description="External ID required by the target role's trust policy on sts:AssumeRole"
|
||||
)
|
||||
aws_bedrock_runtime_endpoint: str | None = Field(default=None, description="AWS Bedrock runtime endpoint URL")
|
||||
checks: BedrockChecksConfigModel | None = Field(
|
||||
default=None,
|
||||
|
|
|
|||
|
|
@ -324,7 +324,12 @@ class AnthropicMessagesToolResultParam(TypedDict, total=False):
|
|||
is_error: bool
|
||||
content: (
|
||||
str
|
||||
| Iterable[AnthropicMessagesToolResultContent | AnthropicMessagesImageParam | AnthropicMessagesDocumentParam]
|
||||
| Iterable[
|
||||
AnthropicMessagesToolResultContent
|
||||
| AnthropicMessagesImageParam
|
||||
| AnthropicMessagesDocumentParam
|
||||
| ToolReference
|
||||
]
|
||||
)
|
||||
cache_control: dict | ChatCompletionCachedContent | None
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from collections.abc import Iterable, Mapping
|
||||
from enum import Enum
|
||||
from os import PathLike
|
||||
from typing import IO, Any, Final, Literal, Optional, Union
|
||||
from typing import IO, Any, Final, Literal, Optional, TypeAlias, Union
|
||||
|
||||
import httpx
|
||||
from openai import Omit
|
||||
|
|
@ -820,9 +820,21 @@ class ChatCompletionAssistantMessage(OpenAIChatCompletionAssistantMessage, total
|
|||
reasoning_items: list[ChatCompletionReasoningItem] | None
|
||||
|
||||
|
||||
class ChatCompletionToolReferenceObject(TypedDict):
|
||||
"""Anthropic tool-search result block, carried through untouched so it survives a round trip."""
|
||||
|
||||
type: Literal["tool_reference"] # writable-ok: Pydantic warns on ReadOnly TypedDict fields
|
||||
tool_name: str # writable-ok: Pydantic warns on ReadOnly TypedDict fields
|
||||
|
||||
|
||||
ToolMessageContentPart: TypeAlias = (
|
||||
ChatCompletionTextObject | ChatCompletionImageObject | ChatCompletionToolReferenceObject
|
||||
)
|
||||
|
||||
|
||||
class ChatCompletionToolMessage(TypedDict):
|
||||
role: Literal["tool"]
|
||||
content: str | Iterable[ChatCompletionTextObject | ChatCompletionImageObject]
|
||||
content: str | Iterable[ToolMessageContentPart] # writable-ok: Pydantic warns on ReadOnly TypedDict fields
|
||||
tool_call_id: str
|
||||
|
||||
|
||||
|
|
@ -1258,6 +1270,8 @@ class ResponsesAPIRequestParams(ResponsesAPIOptionalRequestParams, total=False):
|
|||
|
||||
|
||||
class OutputTokensDetails(BaseLiteLLMOpenAIResponseObject):
|
||||
audio_tokens: int | None = None
|
||||
|
||||
reasoning_tokens: int | None = None
|
||||
|
||||
text_tokens: int | None = None
|
||||
|
|
|
|||
|
|
@ -164,6 +164,7 @@ class ProviderSpecificModelInfo(TypedDict, total=False):
|
|||
supports_low_reasoning_effort: bool | None
|
||||
supports_xhigh_reasoning_effort: bool | None
|
||||
supports_max_reasoning_effort: bool | None
|
||||
reasoning_effort_levels: ReadOnly[Sequence[str] | None]
|
||||
supports_output_config: bool | None
|
||||
supports_image_size: bool | None
|
||||
bedrock_output_config_effort_ceiling: Literal["low", "medium", "high", "max", "xhigh"] | None
|
||||
|
|
|
|||
|
|
@ -5889,6 +5889,7 @@ def _get_model_info_helper(
|
|||
supports_low_reasoning_effort=_model_info.get("supports_low_reasoning_effort", None),
|
||||
supports_xhigh_reasoning_effort=_model_info.get("supports_xhigh_reasoning_effort", None),
|
||||
supports_max_reasoning_effort=_model_info.get("supports_max_reasoning_effort", None),
|
||||
reasoning_effort_levels=_model_info.get("reasoning_effort_levels", None),
|
||||
bedrock_output_config_effort_ceiling=_model_info.get("bedrock_output_config_effort_ceiling", None),
|
||||
bedrock_converse_supports_strict_tools=_model_info.get("bedrock_converse_supports_strict_tools", None),
|
||||
supports_computer_use=_model_info.get("supports_computer_use", None),
|
||||
|
|
|
|||
|
|
@ -9315,6 +9315,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-kimi-k3-through-fireworks-ai-on-microsoft-foundry/4540187",
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -14647,6 +14652,22 @@
|
|||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"dashscope/qwen-image-3.0": {
|
||||
"litellm_provider": "dashscope",
|
||||
"mode": "image_generation",
|
||||
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"dashscope/qwen-image-3.0-pro": {
|
||||
"litellm_provider": "dashscope",
|
||||
"mode": "image_generation",
|
||||
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"databricks/databricks-bge-large-en": {
|
||||
"cache_creation_input_token_cost": 1.0003e-07,
|
||||
"cache_read_input_token_cost": 1.0003e-07,
|
||||
|
|
@ -31881,6 +31902,11 @@
|
|||
"max_tokens": 1048576,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://platform.kimi.ai/docs/pricing/chat-k3",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -36489,6 +36515,14 @@
|
|||
"litellm_provider": "perplexity",
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"minimal",
|
||||
"low",
|
||||
"medium",
|
||||
"high",
|
||||
"xhigh",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.perplexity.ai/docs/agent-api/models",
|
||||
"supports_web_search": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -38791,14 +38825,14 @@
|
|||
"supports_reasoning": true
|
||||
},
|
||||
"together_ai/Qwen/Qwen3.7-Max": {
|
||||
"cache_read_input_token_cost": 1.3e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"input_cost_per_token": 2.5e-06,
|
||||
"litellm_provider": "together_ai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3.75e-06,
|
||||
"output_cost_per_token": 7.5e-06,
|
||||
"source": "https://docs.together.ai/docs/serverless-models",
|
||||
"supports_prompt_caching": true
|
||||
},
|
||||
|
|
@ -38970,6 +39004,11 @@
|
|||
"max_tokens": 1048576,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.together.ai/docs/serverless-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
|
|
@ -39053,6 +39092,24 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"together_ai/zai-org/GLM-5.3-Flash": {
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"litellm_provider": "together_ai",
|
||||
"max_input_tokens": 1048575,
|
||||
"max_output_tokens": 1048575,
|
||||
"max_tokens": 1048575,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-07,
|
||||
"source": "https://docs.together.ai/docs/serverless-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"tts-1": {
|
||||
"input_cost_per_character": 1.5e-05,
|
||||
"litellm_provider": "openai",
|
||||
|
|
@ -51740,6 +51797,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -51804,6 +51866,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -51820,6 +51887,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.25e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -51836,6 +51908,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -52008,6 +52085,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.25e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -52024,6 +52106,11 @@
|
|||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"low",
|
||||
"high",
|
||||
"max"
|
||||
],
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
|
|||
|
|
@ -532,6 +532,22 @@
|
|||
"type": "object",
|
||||
"description": "Provider-internal routing hints (e.g. bedrock_invocation_schema)."
|
||||
},
|
||||
"reasoning_effort_levels": {
|
||||
"type": "array",
|
||||
"description": "Exact reasoning_effort levels this deployment accepts; wins over supports_* flags.",
|
||||
"items": {
|
||||
"type": "string",
|
||||
"enum": [
|
||||
"none",
|
||||
"minimal",
|
||||
"low",
|
||||
"medium",
|
||||
"high",
|
||||
"xhigh",
|
||||
"max"
|
||||
]
|
||||
}
|
||||
},
|
||||
"regional_endpoint_uplift_multiplier": {
|
||||
"type": "number",
|
||||
"minimum": 1,
|
||||
|
|
|
|||
BIN
tests/e2e/llm_translation/fixtures/cat.jpg
Normal file
BIN
tests/e2e/llm_translation/fixtures/cat.jpg
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 30 KiB |
|
|
@ -16,7 +16,10 @@ via /model/new (Cohere, Gemini, hosted_vllm), each deleted on teardown.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -79,18 +82,22 @@ def _streamed_tool_call(events: list[str]) -> tuple[str, str]:
|
|||
return name, arguments
|
||||
|
||||
|
||||
CAT_IMAGE_URL = "https://upload.wikimedia.org/wikipedia/commons/3/3a/Cat03.jpg"
|
||||
_FIXTURES_DIR: Final = Path(__file__).parent / "fixtures"
|
||||
CAT_IMAGE: Final = _FIXTURES_DIR / "cat.jpg"
|
||||
OPENAI_VISION_BACKEND = "openai/gpt-4o"
|
||||
|
||||
# OpenAI caches a shared prompt prefix once it exceeds ~1024 tokens; this is well
|
||||
# past that, so a repeat call reports cached prompt tokens.
|
||||
|
||||
def _cat_image_data_url() -> str:
|
||||
return "data:image/jpeg;base64," + base64.b64encode(CAT_IMAGE.read_bytes()).decode()
|
||||
|
||||
|
||||
def _vision_messages() -> list[ChatMessage]:
|
||||
return [
|
||||
ChatMessage(
|
||||
role="user",
|
||||
content=[
|
||||
TextContentPart(text="What animal is in this image? Answer in one word."),
|
||||
ImageContentPart(image_url=ImageUrl(url=CAT_IMAGE_URL)),
|
||||
ImageContentPart(image_url=ImageUrl(url=_cat_image_data_url())),
|
||||
],
|
||||
)
|
||||
]
|
||||
|
|
|
|||
|
|
@ -3903,3 +3903,39 @@ def test_stored_reasoning_items_win_over_thinking_blocks():
|
|||
reasoning_items = [item for item in input_items if item.get("type") == "reasoning"]
|
||||
assert len(reasoning_items) == 1
|
||||
assert reasoning_items[0]["id"] == "rs_real"
|
||||
|
||||
|
||||
def test_convert_chat_completion_messages_to_responses_api_tool_result_with_tool_reference():
|
||||
"""Tool-search tool_reference blocks have no Responses API equivalent: skip them, never stringify them."""
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
)
|
||||
|
||||
handler = LiteLLMResponsesTransformationHandler()
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_abc123",
|
||||
"type": "function",
|
||||
"function": {"name": "ToolSearch", "arguments": '{"query": "web"}'},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_abc123",
|
||||
"content": [
|
||||
{"type": "tool_reference", "tool_name": "WebFetch"},
|
||||
{"type": "text", "text": "1 tool found"},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
response, _ = handler.convert_chat_completion_messages_to_responses_api(messages)
|
||||
|
||||
function_call_output = next(item for item in response if item.get("type") == "function_call_output")
|
||||
assert function_call_output["output"] == [{"type": "input_text", "text": "1 tool found"}]
|
||||
|
|
|
|||
|
|
@ -0,0 +1,134 @@
|
|||
from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
|
||||
SubtitleToken,
|
||||
render_subtitle_tokens_as_srt,
|
||||
render_subtitle_tokens_as_vtt,
|
||||
synthesize_subtitle_document,
|
||||
)
|
||||
|
||||
|
||||
class TestRenderSubtitleTokensAsSrt:
|
||||
def test_single_cue_full_document(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hello ", start_ms=0, end_ms=500),
|
||||
SubtitleToken(text="world.", start_ms=500, end_ms=1000),
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:00,000 --> 00:00:01,000\nHello world.\n"
|
||||
|
||||
def test_speaker_change_starts_a_new_cue(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hi.", start_ms=0, end_ms=1000, speaker="spk:0"),
|
||||
SubtitleToken(text="Hey.", start_ms=1500, end_ms=2500, speaker="spk:1"),
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == (
|
||||
"1\n00:00:00,000 --> 00:00:01,000\nHi.\n\n2\n00:00:01,500 --> 00:00:02,500\nHey.\n"
|
||||
)
|
||||
|
||||
def test_token_cap_starts_a_new_cue_after_15_tokens(self):
|
||||
tokens = tuple(
|
||||
SubtitleToken(text=f"{index} ", start_ms=index * 100, end_ms=index * 100 + 100) for index in range(16)
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == (
|
||||
"1\n00:00:00,000 --> 00:00:01,500\n0 1 2 3 4 5 6 7 8 9 10 11 12 13 14\n"
|
||||
"\n2\n00:00:01,500 --> 00:00:01,600\n15\n"
|
||||
)
|
||||
|
||||
def test_duration_cap_starts_a_new_cue_at_5000ms(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Alpha ", start_ms=0, end_ms=400),
|
||||
SubtitleToken(text="beta ", start_ms=2000, end_ms=2400),
|
||||
SubtitleToken(text="gamma.", start_ms=5000, end_ms=5400),
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == (
|
||||
"1\n00:00:00,000 --> 00:00:02,400\nAlpha beta\n\n2\n00:00:05,000 --> 00:00:05,400\ngamma.\n"
|
||||
)
|
||||
|
||||
def test_timestampless_token_joins_the_current_cue(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hello ", start_ms=0, end_ms=500),
|
||||
SubtitleToken(text="there "),
|
||||
SubtitleToken(text="world.", start_ms=900, end_ms=1300),
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:00,000 --> 00:00:01,300\nHello there world.\n"
|
||||
|
||||
def test_only_timestampless_tokens_renders_empty(self):
|
||||
assert render_subtitle_tokens_as_srt((SubtitleToken(text="no timestamps"),)) == ""
|
||||
|
||||
def test_empty_tokens_render_empty(self):
|
||||
assert render_subtitle_tokens_as_srt(()) == ""
|
||||
|
||||
def test_timestamps_past_one_hour(self):
|
||||
tokens = (SubtitleToken(text="Late.", start_ms=3_661_001, end_ms=3_662_002),)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n01:01:01,001 --> 01:01:02,002\nLate.\n"
|
||||
|
||||
def test_negative_timestamps_clamp_to_zero(self):
|
||||
tokens = (SubtitleToken(text="Early.", start_ms=-100, end_ms=-50),)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:00,000 --> 00:00:00,000\nEarly.\n"
|
||||
|
||||
def test_missing_end_falls_back_to_cue_start(self):
|
||||
tokens = (SubtitleToken(text="Open.", start_ms=1200),)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:01,200 --> 00:00:01,200\nOpen.\n"
|
||||
|
||||
|
||||
class TestRenderSubtitleTokensAsVtt:
|
||||
def test_single_cue_full_document(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hello ", start_ms=0, end_ms=500),
|
||||
SubtitleToken(text="world.", start_ms=500, end_ms=1000),
|
||||
)
|
||||
assert render_subtitle_tokens_as_vtt(tokens) == "WEBVTT\n\n00:00:00.000 --> 00:00:01.000\nHello world.\n"
|
||||
|
||||
def test_empty_tokens_render_header_only(self):
|
||||
assert render_subtitle_tokens_as_vtt(()) == "WEBVTT\n"
|
||||
|
||||
def test_timestamps_past_one_hour_use_dot_separator(self):
|
||||
tokens = (SubtitleToken(text="Late.", start_ms=3_661_001, end_ms=3_662_002),)
|
||||
assert render_subtitle_tokens_as_vtt(tokens) == "WEBVTT\n\n01:01:01.001 --> 01:01:02.002\nLate.\n"
|
||||
|
||||
def test_speaker_change_starts_a_new_cue(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hi.", start_ms=0, end_ms=1000, speaker=1),
|
||||
SubtitleToken(text="Hey.", start_ms=1500, end_ms=2500, speaker=2),
|
||||
)
|
||||
assert render_subtitle_tokens_as_vtt(tokens) == (
|
||||
"WEBVTT\n\n00:00:00.000 --> 00:00:01.000\nHi.\n\n00:00:01.500 --> 00:00:02.500\nHey.\n"
|
||||
)
|
||||
|
||||
|
||||
class TestSynthesizeSubtitleDocument:
|
||||
WORDS = [
|
||||
{"word": "Four", "start": 0.4, "end": 0.7, "speaker": "spk:0"},
|
||||
{"word": "score", "start": 0.7, "end": 1.1, "speaker": "spk:0"},
|
||||
]
|
||||
|
||||
def test_srt_from_words_converts_seconds_to_milliseconds(self):
|
||||
assert synthesize_subtitle_document(self.WORDS, "srt") == "1\n00:00:00,400 --> 00:00:01,100\nFour score\n"
|
||||
|
||||
def test_vtt_from_words_converts_seconds_to_milliseconds(self):
|
||||
assert synthesize_subtitle_document(self.WORDS, "vtt") == (
|
||||
"WEBVTT\n\n00:00:00.400 --> 00:00:01.100\nFour score\n"
|
||||
)
|
||||
|
||||
def test_speaker_change_splits_cues(self):
|
||||
words = [
|
||||
{"word": "Hi", "start": 0.0, "end": 0.5, "speaker": "spk:0"},
|
||||
{"word": "Hey", "start": 0.6, "end": 1.0, "speaker": "spk:1"},
|
||||
]
|
||||
assert synthesize_subtitle_document(words, "srt") == (
|
||||
"1\n00:00:00,000 --> 00:00:00,500\nHi\n\n2\n00:00:00,600 --> 00:00:01,000\nHey\n"
|
||||
)
|
||||
|
||||
def test_non_subtitle_format_returns_none(self):
|
||||
assert synthesize_subtitle_document(self.WORDS, "verbose_json") is None
|
||||
assert synthesize_subtitle_document(self.WORDS, "json") is None
|
||||
|
||||
def test_missing_words_returns_none(self):
|
||||
assert synthesize_subtitle_document(None, "srt") is None
|
||||
assert synthesize_subtitle_document([], "srt") is None
|
||||
|
||||
def test_words_without_timestamps_return_none(self):
|
||||
assert synthesize_subtitle_document([{"word": "Hello"}], "srt") is None
|
||||
assert synthesize_subtitle_document([{"word": "Hello"}], "vtt") is None
|
||||
|
||||
def test_malformed_words_return_none(self):
|
||||
assert synthesize_subtitle_document("not words", "srt") is None
|
||||
assert synthesize_subtitle_document([{"word": "ok", "start": "not-a-number"}], "srt") is None
|
||||
|
|
@ -1027,3 +1027,70 @@ def test_update_messages_xlitellm_decode_does_not_override_mapping():
|
|||
updated = update_messages_with_model_file_ids(messages, "model-A", mapping)
|
||||
|
||||
assert updated[0]["content"][0]["file"]["file_id"] == "provider-explicit-id"
|
||||
|
||||
|
||||
def test_drop_tool_reference_parts_keeps_text_parts():
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
)
|
||||
|
||||
messages = [
|
||||
_assistant_tool_call_msg("call_1"),
|
||||
_tool_msg(
|
||||
[
|
||||
{"type": "text", "text": "WebFetch tool loaded successfully."},
|
||||
{"type": "tool_reference", "tool_name": "WebFetch"},
|
||||
]
|
||||
),
|
||||
]
|
||||
|
||||
result = drop_tool_reference_parts_from_tool_messages(messages)
|
||||
|
||||
assert result[1]["content"] == [{"type": "text", "text": "WebFetch tool loaded successfully."}]
|
||||
assert result[1]["tool_call_id"] == "call_1"
|
||||
|
||||
|
||||
def test_drop_tool_reference_parts_reference_only_becomes_empty_text():
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
)
|
||||
|
||||
messages = [
|
||||
_assistant_tool_call_msg("call_1"),
|
||||
_tool_msg([{"type": "tool_reference", "tool_name": "WebFetch"}]),
|
||||
]
|
||||
|
||||
result = drop_tool_reference_parts_from_tool_messages(messages)
|
||||
|
||||
assert result[1] == {"role": "tool", "tool_call_id": "call_1", "content": ""}
|
||||
|
||||
|
||||
def test_drop_tool_reference_parts_without_references_passes_through():
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
)
|
||||
|
||||
messages = [
|
||||
_assistant_tool_call_msg("call_1"),
|
||||
_tool_msg([{"type": "text", "text": "plain result"}]),
|
||||
]
|
||||
|
||||
assert drop_tool_reference_parts_from_tool_messages(messages) is messages
|
||||
|
||||
|
||||
def test_drop_tool_reference_parts_leaves_non_tool_messages_alone():
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
)
|
||||
|
||||
user_message = {"role": "user", "content": [{"type": "tool_reference", "tool_name": "WebFetch"}]}
|
||||
messages = [
|
||||
user_message,
|
||||
_assistant_tool_call_msg("call_1"),
|
||||
_tool_msg([{"type": "tool_reference", "tool_name": "WebFetch"}]),
|
||||
]
|
||||
|
||||
result = drop_tool_reference_parts_from_tool_messages(messages)
|
||||
|
||||
assert result[0] == user_message
|
||||
assert result[2]["content"] == ""
|
||||
|
|
|
|||
|
|
@ -3578,3 +3578,52 @@ async def test_bedrock_converse_pdf_only_user_message_gets_text_block_async():
|
|||
assert len(result) == 1
|
||||
assert any("document" in block for block in result[0]["content"])
|
||||
assert _text_blocks(result[0]) == [BEDROCK_DOCUMENT_PLACEHOLDER_TEXT]
|
||||
|
||||
|
||||
def test_convert_to_anthropic_tool_result_keeps_tool_reference_blocks():
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import convert_to_anthropic_tool_result
|
||||
|
||||
result = convert_to_anthropic_tool_result(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "toolu_01",
|
||||
"content": [
|
||||
{"type": "text", "text": "loaded"},
|
||||
{"type": "tool_reference", "tool_name": "WebFetch"},
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
assert result == {
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "toolu_01",
|
||||
"content": [
|
||||
{"type": "text", "text": "loaded"},
|
||||
{"type": "tool_reference", "tool_name": "WebFetch"},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def test_convert_gemini_tool_call_result_answers_tool_reference_only_result():
|
||||
"""Every Gemini function call needs a function response, even when the tool result carries no text.
|
||||
Fixes: https://github.com/BerriAI/litellm/issues/37462
|
||||
"""
|
||||
result = convert_to_gemini_tool_call_result(
|
||||
message=ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id="toolu_01",
|
||||
content=[{"type": "tool_reference", "tool_name": "WebFetch"}],
|
||||
),
|
||||
last_message_with_tool_calls={
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "toolu_01",
|
||||
"type": "function",
|
||||
"function": {"name": "ToolSearch", "arguments": '{"query": "select:WebFetch"}'},
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
assert result == {"function_response": {"name": "ToolSearch", "response": {"content": ""}}}
|
||||
|
|
|
|||
|
|
@ -4460,3 +4460,51 @@ def test_handle_stream_fallback_error_restores_context_only_after_exception_mapp
|
|||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_chunk_creator_preserves_hidden_provider_specific_fields_from_parsed_chunk():
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=None,
|
||||
model="gemini-3.5-flash",
|
||||
logging_obj=MagicMock(),
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
parsed_chunk = ModelResponseStream(
|
||||
choices=[StreamingChoices(index=0, delta=Delta(content="hello", role="assistant"), finish_reason=None)],
|
||||
)
|
||||
parsed_chunk._hidden_params["provider_specific_fields"] = {"traffic_type": "ON_DEMAND_FLEX"}
|
||||
|
||||
result = wrapper.chunk_creator(chunk=parsed_chunk)
|
||||
|
||||
assert result is not None
|
||||
assert result._hidden_params["provider_specific_fields"] == {"traffic_type": "ON_DEMAND_FLEX"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_stream_assembled_response_keeps_vertex_traffic_type(logging_obj: Logging):
|
||||
content_chunk = ModelResponseStream(
|
||||
choices=[StreamingChoices(index=0, delta=Delta(content="hello", role="assistant"), finish_reason=None)],
|
||||
)
|
||||
final_chunk = ModelResponseStream(
|
||||
choices=[StreamingChoices(index=0, delta=Delta(content=""), finish_reason="stop")],
|
||||
)
|
||||
setattr(final_chunk, "usage", Usage(prompt_tokens=7, completion_tokens=5, total_tokens=12))
|
||||
final_chunk._hidden_params["provider_specific_fields"] = {"traffic_type": "ON_DEMAND_FLEX"}
|
||||
|
||||
async def _stream():
|
||||
yield content_chunk
|
||||
yield final_chunk
|
||||
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=_stream(),
|
||||
model="gemini-3.5-flash",
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider="vertex_ai",
|
||||
stream_options={"include_usage": True},
|
||||
)
|
||||
|
||||
received = [chunk async for chunk in wrapper]
|
||||
|
||||
assembled = litellm.stream_chunk_builder(chunks=received, messages=[{"role": "user", "content": "hi"}])
|
||||
assert assembled is not None
|
||||
assert assembled._hidden_params["provider_specific_fields"]["traffic_type"] == "ON_DEMAND_FLEX"
|
||||
|
|
|
|||
|
|
@ -290,6 +290,24 @@ class TestAnthropicMessagesHandlerInputProcessing:
|
|||
assert data.get("litellm_metadata", {}).get("guardrails")
|
||||
assert guardrail.dynamic_params == {"policy_id": "policy-123"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_native_tools_survive_guardrail_round_trip(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockPassThroughGuardrail(guardrail_name="test")
|
||||
data = {
|
||||
"model": "gemini-2.5-flash",
|
||||
"messages": [{"role": "user", "content": "coffee shops near Union Square?"}],
|
||||
"tools": [
|
||||
{"googleMaps": {"enable_widget": True}},
|
||||
{"name": "get_weather", "input_schema": {"type": "object", "properties": {}}},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert {"googleMaps": {"enable_widget": True}} in data["tools"]
|
||||
assert [tool["name"] for tool in data["tools"] if "name" in tool] == ["get_weather"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_midturn_system_correction_is_guardrailed_when_top_level_system_is_skipped(
|
||||
self,
|
||||
|
|
@ -1818,3 +1836,72 @@ class TestAnthropicMessagesScanOnlyToolResults:
|
|||
|
||||
assert guardrail.captured_inputs is not None
|
||||
assert guardrail.captured_inputs.get("images") == ["TOOL_IMG"]
|
||||
|
||||
|
||||
class TestStructuredWriteBackKeepsToolResults:
|
||||
"""A guardrail rewrite must never leave a tool_use without its tool_result (Claude Code ToolSearch, LIT-6103)."""
|
||||
|
||||
@staticmethod
|
||||
def _claude_code_tool_search_turns(tool_result_content):
|
||||
return [
|
||||
{"role": "user", "content": "load WebFetch for bob@example.com"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "toolu_01",
|
||||
"name": "ToolSearch",
|
||||
"input": {"query": "select:WebFetch"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "tool_result", "tool_use_id": "toolu_01", "content": tool_result_content},
|
||||
{"type": "text", "text": "Now fetch the page."},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _blocks(message):
|
||||
return message["content"] if isinstance(message["content"], list) else []
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("tool_result_content", "expected_written_back_content"),
|
||||
[
|
||||
(
|
||||
[{"type": "tool_reference", "tool_name": "WebFetch"}],
|
||||
[{"type": "tool_reference", "tool_name": "WebFetch"}],
|
||||
),
|
||||
([], ""),
|
||||
],
|
||||
ids=["tool_reference", "empty"],
|
||||
)
|
||||
async def test_tool_result_stays_right_after_its_tool_use(
|
||||
self, tool_result_content, expected_written_back_content
|
||||
):
|
||||
handler = AnthropicMessagesHandler()
|
||||
data = {"model": "claude-fable-5", "messages": self._claude_code_tool_search_turns(tool_result_content)}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=MockStructuredMaskingGuardrail())
|
||||
|
||||
serialized = json.dumps(data["messages"])
|
||||
assert "bob@example.com" not in serialized
|
||||
assert "<EMAIL>" in serialized
|
||||
|
||||
messages = data["messages"]
|
||||
tool_use_index = next(
|
||||
i for i, m in enumerate(messages) if any(b.get("type") == "tool_use" for b in self._blocks(m))
|
||||
)
|
||||
answer = messages[tool_use_index + 1]
|
||||
assert answer["role"] == "user"
|
||||
assert answer["content"][0] == {
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "toolu_01",
|
||||
"content": expected_written_back_content,
|
||||
}
|
||||
later_blocks = [b for m in messages[tool_use_index + 1 :] for b in self._blocks(m)]
|
||||
assert {"type": "text", "text": "Now fetch the page."} in later_blocks
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
|
|||
)
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
|
||||
OPENAI_MAX_TOOL_NAME_LENGTH,
|
||||
AnthropicAdapter,
|
||||
LiteLLMAnthropicMessagesAdapter,
|
||||
create_tool_name_mapping,
|
||||
truncate_tool_name,
|
||||
|
|
@ -2307,6 +2308,53 @@ def test_translate_anthropic_tools_to_openai_fills_missing_tool_name():
|
|||
assert result[1]["function"]["name"] == "litellm_unnamed_tool_1"
|
||||
|
||||
|
||||
def test_translate_anthropic_tools_to_openai_passes_provider_native_tool_dicts_through():
|
||||
"""Deployment-level provider-native tools (e.g. Gemini googleMaps) must reach the provider transformation verbatim (LIT-6286)."""
|
||||
tools = [
|
||||
{"googleMaps": {}},
|
||||
{"googleSearch": {}},
|
||||
{
|
||||
"name": "get_weather",
|
||||
"input_schema": {"type": "object", "properties": {"location": {"type": "string"}}},
|
||||
},
|
||||
]
|
||||
adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
result, tool_name_mapping = adapter.translate_anthropic_tools_to_openai(tools=tools, model=None)
|
||||
assert result[0] == {"googleMaps": {}}
|
||||
assert result[1] == {"googleSearch": {}}
|
||||
assert result[2]["function"]["name"] == "get_weather"
|
||||
assert tool_name_mapping == {}
|
||||
|
||||
|
||||
def test_translate_anthropic_tools_to_openai_passes_openai_function_tools_through():
|
||||
"""A tool already in OpenAI function format must pass through unchanged instead of becoming litellm_unnamed_tool_N."""
|
||||
openai_tool = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"parameters": {"type": "object", "properties": {"location": {"type": "string"}}},
|
||||
},
|
||||
}
|
||||
adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
result, _ = adapter.translate_anthropic_tools_to_openai(tools=[openai_tool], model=None)
|
||||
assert result == [openai_tool]
|
||||
|
||||
|
||||
def test_translate_completion_input_params_keeps_provider_native_tools():
|
||||
"""/v1/messages request translation must keep router-merged provider-native tools in kwargs['tools'] (LIT-6286)."""
|
||||
adapter = AnthropicAdapter()
|
||||
translated = adapter.translate_completion_input_params(
|
||||
{
|
||||
"model": "gemini/gemini-2.5-flash",
|
||||
"max_tokens": 1024,
|
||||
"messages": [{"role": "user", "content": "coffee shops near Union Square"}],
|
||||
"tools": [{"googleMaps": {}}],
|
||||
}
|
||||
)
|
||||
assert translated is not None
|
||||
assert translated["tools"] == [{"googleMaps": {}}]
|
||||
|
||||
|
||||
def test_translate_openai_content_to_anthropic_reasoning_content_without_thinking_blocks():
|
||||
"""
|
||||
Test that reasoning_content is converted to thinking block when thinking_blocks is not present.
|
||||
|
|
@ -3999,6 +4047,75 @@ def test_translate_anthropic_messages_to_openai_carries_midturn_system_prompt_ca
|
|||
]
|
||||
|
||||
|
||||
def _tool_reference_block(tool_name="WebFetch"):
|
||||
return {"type": "tool_reference", "tool_name": tool_name}
|
||||
|
||||
|
||||
def test_tool_result_tool_reference_is_carried_through_untouched():
|
||||
adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
|
||||
result = adapter.translate_anthropic_messages_to_openai(
|
||||
messages=[
|
||||
_anthropic_tool_use_turn("toolu_01"),
|
||||
_anthropic_tool_result_turn({"toolu_01": [_tool_reference_block()]}),
|
||||
]
|
||||
)
|
||||
|
||||
assert [m["role"] for m in result] == ["assistant", "tool"]
|
||||
assert result[1]["tool_call_id"] == "toolu_01"
|
||||
assert result[1]["content"] == [{"type": "tool_reference", "tool_name": "WebFetch"}]
|
||||
|
||||
|
||||
def test_tool_result_text_beside_tool_reference_keeps_both_parts_in_order():
|
||||
adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
|
||||
result = adapter.translate_anthropic_messages_to_openai(
|
||||
messages=[
|
||||
_anthropic_tool_use_turn("toolu_01"),
|
||||
_anthropic_tool_result_turn(
|
||||
{"toolu_01": [{"type": "text", "text": "loaded"}, _tool_reference_block("Grep")]}
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
assert result[1]["content"] == [
|
||||
{"type": "text", "text": "loaded"},
|
||||
{"type": "tool_reference", "tool_name": "Grep"},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"tool_result_content",
|
||||
[
|
||||
[],
|
||||
None,
|
||||
"",
|
||||
{"not": "a list"},
|
||||
[{"type": "future_block", "payload": 1}],
|
||||
[{"type": "search_result", "source": "https://example.com", "title": "t", "content": []}],
|
||||
],
|
||||
ids=["empty_list", "null", "empty_string", "non_list", "unknown_block", "search_result_only"],
|
||||
)
|
||||
def test_tool_result_without_translatable_content_still_answers_its_tool_use(tool_result_content):
|
||||
adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
|
||||
result = adapter.translate_anthropic_messages_to_openai(
|
||||
messages=[
|
||||
_anthropic_tool_use_turn("toolu_01"),
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": tool_result_content}],
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
assert result == [
|
||||
result[0],
|
||||
{"role": "tool", "tool_call_id": "toolu_01", "content": ""},
|
||||
]
|
||||
assert result[0]["role"] == "assistant"
|
||||
|
||||
|
||||
def _openai_response_with_usage(usage: Usage) -> ModelResponse:
|
||||
return ModelResponse(
|
||||
id="resp_web_search",
|
||||
|
|
|
|||
|
|
@ -14,6 +14,8 @@ from unittest.mock import patch
|
|||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.utils import (
|
||||
normalize_reasoning_effort_value,
|
||||
)
|
||||
|
|
@ -291,3 +293,91 @@ class TestAdapterAdaptiveThinking:
|
|||
)
|
||||
assert result is not None
|
||||
assert result["effort"] == "medium"
|
||||
|
||||
|
||||
class TestDeclaredEffortsAnswerTheDegradationGate:
|
||||
"""Without this the chain reads only the per-level booleans, so a kimi-k3 request asking for
|
||||
max silently arrives as high."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model, provider",
|
||||
[("kimi-k3", "moonshot"), ("kimi-k3", "fireworks_ai"), ("kimi-k3-us", "fireworks_ai")],
|
||||
)
|
||||
def test_a_declared_level_survives_instead_of_degrading(self, local_model_cost_map, model, provider):
|
||||
assert normalize_reasoning_effort_value("max", model, provider) == "max"
|
||||
|
||||
def test_a_level_the_entry_does_not_declare_still_degrades(self, local_model_cost_map):
|
||||
"""xhigh is not on kimi-k3's declaration, so it must keep degrading rather than be waved
|
||||
past by the mere presence of one."""
|
||||
assert normalize_reasoning_effort_value("xhigh", "kimi-k3", "moonshot") == "high"
|
||||
assert normalize_reasoning_effort_value("minimal", "kimi-k3", "moonshot") == "low"
|
||||
|
||||
def test_the_wider_perplexity_entry_keeps_the_levels_it_declares(self, local_model_cost_map):
|
||||
assert normalize_reasoning_effort_value("xhigh", "perplexity/kimi-k3", "perplexity") == "xhigh"
|
||||
assert normalize_reasoning_effort_value("minimal", "perplexity/kimi-k3", "perplexity") == "minimal"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model, provider, effort, expected",
|
||||
[
|
||||
("claude-opus-4-7", "anthropic", "max", "max"),
|
||||
("claude-sonnet-4-6", "anthropic", "minimal", "low"),
|
||||
("gpt-5-mini", "azure", "max", "high"),
|
||||
],
|
||||
)
|
||||
def test_an_entry_on_the_per_level_flags_is_untouched(
|
||||
self, local_model_cost_map, model, provider, effort, expected
|
||||
):
|
||||
"""The negative class that bounds this change to entries carrying the key."""
|
||||
assert normalize_reasoning_effort_value(effort, model, provider) == expected
|
||||
|
||||
|
||||
class TestDeclarationBeatsThePerLevelFlags:
|
||||
"""An entry can carry both shapes. The declaration wins whole, or /model_group/info and this
|
||||
path would disagree about the same deployment. Driven through the public entry point over a
|
||||
seeded map entry rather than a patched get_model_info, so it pins behaviour and not wiring."""
|
||||
|
||||
MODEL = "declared-and-flagged"
|
||||
|
||||
@pytest.fixture
|
||||
def seeded(self, local_model_cost_map, monkeypatch):
|
||||
def _seed(**entry):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
self.MODEL,
|
||||
{"litellm_provider": "openai", "mode": "chat", "supports_reasoning": True, **entry},
|
||||
)
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
return _seed
|
||||
|
||||
@pytest.mark.parametrize("effort, expected", [("max", "max"), ("xhigh", "high"), ("minimal", "low")])
|
||||
def test_a_flag_cannot_re_add_a_level_the_declaration_omits(self, seeded, effort, expected):
|
||||
seeded(
|
||||
reasoning_effort_levels=["low", "high", "max"],
|
||||
supports_xhigh_reasoning_effort=True,
|
||||
supports_minimal_reasoning_effort=True,
|
||||
supports_max_reasoning_effort=False,
|
||||
)
|
||||
|
||||
assert normalize_reasoning_effort_value(effort, self.MODEL, "openai") == expected
|
||||
|
||||
def test_a_flag_cannot_keep_max_when_the_declaration_drops_it(self, seeded):
|
||||
seeded(
|
||||
reasoning_effort_levels=["low", "high"],
|
||||
supports_max_reasoning_effort=True,
|
||||
supports_xhigh_reasoning_effort=True,
|
||||
)
|
||||
|
||||
assert normalize_reasoning_effort_value("max", self.MODEL, "openai") == "high"
|
||||
|
||||
def test_a_false_flag_cannot_remove_a_level_the_declaration_names(self, seeded):
|
||||
seeded(reasoning_effort_levels=["high", "xhigh"], supports_xhigh_reasoning_effort=False)
|
||||
|
||||
assert normalize_reasoning_effort_value("xhigh", self.MODEL, "openai") == "xhigh"
|
||||
assert normalize_reasoning_effort_value("max", self.MODEL, "openai") == "xhigh"
|
||||
|
||||
def test_a_chain_the_declaration_omits_entirely_lands_on_its_terminal(self, seeded):
|
||||
"""Documented residual: no strength ordering exists to pick a nearer declared level."""
|
||||
seeded(reasoning_effort_levels=["high", "xhigh"])
|
||||
|
||||
assert normalize_reasoning_effort_value("minimal", self.MODEL, "openai") == "low"
|
||||
|
|
|
|||
|
|
@ -102,6 +102,35 @@ def test_transform_request_hoists_tool_message_image():
|
|||
]
|
||||
|
||||
|
||||
def test_transform_request_drops_tool_reference_parts():
|
||||
"""Azure's transform_request shares the tool-message sanitizing with OpenAI:
|
||||
tool_reference parts are dropped, a reference-only result keeps its tool
|
||||
message with empty text (#37462 round trip)."""
|
||||
messages = [
|
||||
{"role": "user", "content": "load the WebFetch tool"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "ToolSearch", "arguments": "{}"}}],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_1",
|
||||
"content": [{"type": "tool_reference", "tool_name": "WebFetch"}],
|
||||
},
|
||||
]
|
||||
|
||||
request = AzureOpenAIConfig().transform_request(
|
||||
model="gpt-4o",
|
||||
messages=messages,
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert request["messages"][2]["content"] == ""
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model, emitted_key, absent_key",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -1901,6 +1901,76 @@ async def test_async_audio_transcriptions_sends_dict_data_as_json_body():
|
|||
assert response.text == "transcribed"
|
||||
|
||||
|
||||
class _WordTimestampAudioTranscriptionConfig(_JSONBodyAudioTranscriptionConfig):
|
||||
def transform_audio_transcription_response(self, raw_response):
|
||||
payload = raw_response.json()
|
||||
response = TranscriptionResponse(text=payload["text"])
|
||||
response["words"] = payload["words"]
|
||||
return response
|
||||
|
||||
|
||||
def test_transform_audio_transcription_response_without_subtitle_opt_in_keeps_text_and_words():
|
||||
words = [
|
||||
{"word": "hello", "start": 0.0, "end": 0.5},
|
||||
{"word": "world", "start": 0.5, "end": 1.0},
|
||||
]
|
||||
raw_response = httpx.Response(200, json={"text": "hello world", "words": words})
|
||||
|
||||
response = BaseLLMHTTPHandler()._transform_audio_transcription_response(
|
||||
provider_config=_WordTimestampAudioTranscriptionConfig(),
|
||||
model="test-model",
|
||||
response=raw_response,
|
||||
model_response=TranscriptionResponse(),
|
||||
logging_obj=Mock(),
|
||||
optional_params={"response_format": "srt"},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert response.text == "hello world"
|
||||
assert response["words"] == words
|
||||
|
||||
|
||||
class _SubtitleSynthesisAudioTranscriptionConfig(_JSONBodyAudioTranscriptionConfig):
|
||||
@property
|
||||
def supports_subtitle_synthesis(self) -> bool:
|
||||
return True
|
||||
|
||||
def transform_audio_transcription_response(self, raw_response):
|
||||
payload = raw_response.json()
|
||||
response = TranscriptionResponse(text=payload["text"])
|
||||
if "words" in payload:
|
||||
response["words"] = payload["words"]
|
||||
return response
|
||||
|
||||
|
||||
def _transform_subtitle_response(payload):
|
||||
return BaseLLMHTTPHandler()._transform_audio_transcription_response(
|
||||
provider_config=_SubtitleSynthesisAudioTranscriptionConfig(),
|
||||
model="test-model",
|
||||
response=httpx.Response(200, json=payload),
|
||||
model_response=TranscriptionResponse(),
|
||||
logging_obj=Mock(),
|
||||
optional_params={"response_format": "srt"},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
|
||||
def test_subtitle_synthesis_fallback_without_timings_drops_words():
|
||||
response = _transform_subtitle_response(
|
||||
{"text": "hello world", "words": [{"word": "hello"}, {"word": "world"}]}
|
||||
)
|
||||
|
||||
assert response.text == "hello world"
|
||||
assert "words" not in response
|
||||
|
||||
|
||||
def test_subtitle_synthesis_without_words_keeps_plain_text():
|
||||
response = _transform_subtitle_response({"text": "hello world"})
|
||||
|
||||
assert response.text == "hello world"
|
||||
assert "words" not in response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_retrieve_file_content_raises_on_http_error():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -169,6 +169,42 @@ class TestTransformRequest:
|
|||
}
|
||||
}
|
||||
|
||||
@pytest.mark.parametrize("response_format", ["srt", "vtt"])
|
||||
def test_subtitle_response_format_requests_word_timestamps(self, config, response_format):
|
||||
request_data = config.transform_audio_transcription_request(
|
||||
model="gemini-3.5-transcribe",
|
||||
audio_file=("sample.wav", AUDIO_BYTES, "audio/wav"),
|
||||
optional_params={"response_format": response_format},
|
||||
litellm_params={},
|
||||
)
|
||||
transcription_config = request_data.data["generation_config"]["transcription_config"]
|
||||
assert json.loads(json.dumps(transcription_config)) == {
|
||||
"mode": {
|
||||
"type": "verbatim",
|
||||
"timestamp_granularities": ["word"],
|
||||
"diarization_mode": "speaker",
|
||||
}
|
||||
}
|
||||
|
||||
@pytest.mark.parametrize("response_format", ["json", "text", "verbose_json"])
|
||||
def test_non_subtitle_response_format_sends_no_mode(self, config, response_format):
|
||||
request_data = config.transform_audio_transcription_request(
|
||||
model="gemini-3.5-transcribe",
|
||||
audio_file=("sample.wav", AUDIO_BYTES, "audio/wav"),
|
||||
optional_params={"response_format": response_format},
|
||||
litellm_params={},
|
||||
)
|
||||
assert "generation_config" not in request_data.data
|
||||
|
||||
def test_non_string_response_format_sends_no_mode(self, config):
|
||||
request_data = config.transform_audio_transcription_request(
|
||||
model="gemini-3.5-transcribe",
|
||||
audio_file=("sample.wav", AUDIO_BYTES, "audio/wav"),
|
||||
optional_params={"response_format": {"type": "json_object"}},
|
||||
litellm_params={},
|
||||
)
|
||||
assert "generation_config" not in request_data.data
|
||||
|
||||
def test_segment_granularity_sends_no_mode(self, config):
|
||||
request_data = config.transform_audio_transcription_request(
|
||||
model="gemini-3.5-transcribe",
|
||||
|
|
@ -214,6 +250,54 @@ class TestTransformResponse:
|
|||
assert response.get("duration") is None
|
||||
|
||||
|
||||
class TestSubtitleSynthesisThroughHandler:
|
||||
def _transform(self, config, response_format):
|
||||
from unittest.mock import Mock
|
||||
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.utils import TranscriptionResponse
|
||||
|
||||
return BaseLLMHTTPHandler()._transform_audio_transcription_response(
|
||||
provider_config=config,
|
||||
model="gemini-3.5-transcribe",
|
||||
response=make_response(COMPLETED_RESPONSE),
|
||||
model_response=TranscriptionResponse(),
|
||||
logging_obj=Mock(),
|
||||
optional_params={"response_format": response_format},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
def test_supports_subtitle_synthesis(self, config):
|
||||
assert config.supports_subtitle_synthesis is True
|
||||
|
||||
def test_srt_synthesizes_subtitle_document_and_drops_words(self, config):
|
||||
response = self._transform(config, "srt")
|
||||
assert response.text == (
|
||||
"1\n00:00:00,100 --> 00:00:00,400\nHello\n\n2\n00:00:00,500 --> 00:00:00,900\nworld.\n"
|
||||
)
|
||||
assert "words" not in response
|
||||
assert response["task"] == "transcribe"
|
||||
assert response["duration"] == 0.9
|
||||
assert response.usage.total_tokens == 200
|
||||
|
||||
def test_vtt_synthesizes_subtitle_document_and_drops_words(self, config):
|
||||
response = self._transform(config, "vtt")
|
||||
assert response.text == (
|
||||
"WEBVTT\n\n00:00:00.100 --> 00:00:00.400\nHello\n\n00:00:00.500 --> 00:00:00.900\nworld.\n"
|
||||
)
|
||||
assert "words" not in response
|
||||
assert response.usage.total_tokens == 200
|
||||
|
||||
@pytest.mark.parametrize("response_format", ["json", "verbose_json"])
|
||||
def test_non_subtitle_formats_keep_plain_text_and_words(self, config, response_format):
|
||||
response = self._transform(config, response_format)
|
||||
assert response.text == "Hello world."
|
||||
assert response["words"] == [
|
||||
{"word": "Hello", "start": 0.1, "end": 0.4, "speaker": "spk:0"},
|
||||
{"word": "world.", "start": 0.5, "end": 0.9, "speaker": "spk:1"},
|
||||
]
|
||||
|
||||
|
||||
class TestCostRegression:
|
||||
@pytest.fixture
|
||||
def local_cost_map(self, monkeypatch):
|
||||
|
|
|
|||
|
|
@ -1866,6 +1866,54 @@ def test_map_openai_params_drops_stock_voice_case_insensitively():
|
|||
assert passthrough["generationConfig"]["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore"
|
||||
|
||||
|
||||
def test_gemini_response_done_bills_audio_output_tokens_at_audio_rate(monkeypatch):
|
||||
"""Regression for the Gemini Live AUDIO output breakdown: responseTokensDetails
|
||||
must survive into response.done usage and bill at output_cost_per_audio_token,
|
||||
not the text rate."""
|
||||
from litellm.cost_calculator import (
|
||||
RealtimeAPITokenUsageProcessor,
|
||||
handle_realtime_stream_cost_calculation,
|
||||
)
|
||||
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
|
||||
|
||||
config = GeminiRealtimeConfig()
|
||||
done_event = config.transform_response_done_event(
|
||||
message={
|
||||
"serverContent": {"turnComplete": True},
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 377,
|
||||
"responseTokenCount": 51,
|
||||
"totalTokenCount": 428,
|
||||
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 377}],
|
||||
"responseTokensDetails": [{"modality": "AUDIO", "tokenCount": 51}],
|
||||
"thoughtsTokenCount": 37,
|
||||
},
|
||||
},
|
||||
current_response_id="resp_lit6277",
|
||||
current_conversation_id="conv_lit6277",
|
||||
output_items=None,
|
||||
)
|
||||
|
||||
usage = done_event["response"]["usage"]
|
||||
assert usage["output_tokens_details"]["audio_tokens"] == 51
|
||||
assert usage["output_token_details"]["audio_tokens"] == 51
|
||||
|
||||
results = [done_event]
|
||||
combined_usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(
|
||||
results=results,
|
||||
)
|
||||
assert combined_usage.completion_tokens_details is not None
|
||||
assert combined_usage.completion_tokens_details.audio_tokens == 51
|
||||
|
||||
cost = handle_realtime_stream_cost_calculation(
|
||||
results=results,
|
||||
combined_usage_object=combined_usage,
|
||||
custom_llm_provider="gemini",
|
||||
litellm_model_name="gemini-2.5-flash-native-audio-preview-12-2025",
|
||||
)
|
||||
assert cost == pytest.approx(377 * 5e-07 + 51 * 1.2e-05 + 37 * 2e-06)
|
||||
@pytest.fixture(autouse=False)
|
||||
def patch_gemini_transcribe_live_cost_map_entry(monkeypatch):
|
||||
"""Inject the gemini-3.5-transcribe-live registry entry locally.
|
||||
|
|
|
|||
|
|
@ -869,6 +869,69 @@ class TestToolMessageImageHoisting:
|
|||
assert result[3]["content"] == self.HOISTED_USER_CONTENT
|
||||
|
||||
|
||||
class TestToolReferenceStripping:
|
||||
"""transform_request drops tool_reference parts from tool messages: OpenAI's
|
||||
chat API rejects them, and the reference names an already-declared tool
|
||||
rather than carrying content (#37462 round trip)."""
|
||||
|
||||
def setup_method(self):
|
||||
self.config = OpenAIGPTConfig()
|
||||
|
||||
def _messages_with_tool_reference(self, extra_parts=()):
|
||||
return [
|
||||
{"role": "user", "content": "load the WebFetch tool"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{"id": "call_1", "type": "function", "function": {"name": "ToolSearch", "arguments": "{}"}}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_1",
|
||||
"content": [*extra_parts, {"type": "tool_reference", "tool_name": "WebFetch"}],
|
||||
},
|
||||
]
|
||||
|
||||
def test_transform_request_keeps_text_and_drops_reference(self):
|
||||
request = self.config.transform_request(
|
||||
model="gpt-4.1",
|
||||
messages=self._messages_with_tool_reference(extra_parts=({"type": "text", "text": "loaded"},)),
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
tool_message = request["messages"][2]
|
||||
assert tool_message["content"] == [{"type": "text", "text": "loaded"}]
|
||||
assert tool_message["tool_call_id"] == "call_1"
|
||||
|
||||
def test_transform_request_reference_only_keeps_tool_message_with_empty_text(self):
|
||||
request = self.config.transform_request(
|
||||
model="gpt-4.1",
|
||||
messages=self._messages_with_tool_reference(),
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert [m.get("role") for m in request["messages"]] == ["user", "assistant", "tool"]
|
||||
assert request["messages"][2]["content"] == ""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_transform_request_drops_reference(self):
|
||||
request = await self.config.async_transform_request(
|
||||
model="gpt-4.1",
|
||||
messages=self._messages_with_tool_reference(),
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert request["messages"][2]["content"] == ""
|
||||
|
||||
|
||||
class TestOpenAIPromptCacheBreakpointChatPath:
|
||||
"""Chat-path shape for OpenAI explicit prompt caching (#37509)."""
|
||||
|
||||
|
|
|
|||
|
|
@ -477,12 +477,12 @@ class TestBuildResponseWithResponseFormat:
|
|||
}
|
||||
}
|
||||
# SRT requested but tokens have no start_ms/end_ms -> empty SRT
|
||||
# falls back gracefully since _group_tokens_into_cues skips them
|
||||
# falls back gracefully since group_subtitle_tokens_into_cues skips them
|
||||
resp = cfg._build_response_from_payload(payload, response_format="srt")
|
||||
# With no timestamp data, SRT rendering produces empty string,
|
||||
# but we still get output because the code checks `tokens` truthiness
|
||||
# before choosing SRT path. Actually the tokens list is truthy but
|
||||
# _group_tokens_into_cues will produce no cues -> empty SRT string.
|
||||
# group_subtitle_tokens_into_cues will produce no cues -> empty SRT string.
|
||||
# Let's verify it doesn't crash.
|
||||
assert isinstance(resp.text, str)
|
||||
|
||||
|
|
|
|||
|
|
@ -5274,3 +5274,74 @@ async def test_terminal_failure_logs_usage_and_cost_of_prior_passed_chunks(monke
|
|||
assert logged["guardrail_cost"] == pytest.approx(0.0003)
|
||||
assert logged["guardrail_response"]["usage"] == {"contentPolicyUnits": 2, "wordPolicyUnits": 1}
|
||||
assert "error" in logged["guardrail_response"]
|
||||
|
||||
|
||||
def test_load_credentials_assumes_role_with_external_id():
|
||||
"""A trust policy requiring sts:ExternalId must be satisfied by the guardrail's aws_external_id."""
|
||||
import datetime
|
||||
|
||||
import boto3
|
||||
from botocore.exceptions import ClientError
|
||||
|
||||
class FakeSTSClient:
|
||||
"""STS that mirrors a cross-account role whose trust policy requires an ExternalId."""
|
||||
|
||||
def get_caller_identity(self):
|
||||
return {"Arn": "arn:aws:iam::111111111111:user/litellm-proxy-pod"}
|
||||
|
||||
def assume_role(self, **params):
|
||||
if params.get("ExternalId") != "external-id-123":
|
||||
raise ClientError(
|
||||
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:AssumeRole"}},
|
||||
"AssumeRole",
|
||||
)
|
||||
return {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "ASIAASSUMEDROLEKEY",
|
||||
"SecretAccessKey": "assumed-secret",
|
||||
"SessionToken": "assumed-session-token",
|
||||
"Expiration": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30),
|
||||
}
|
||||
}
|
||||
|
||||
guardrail = BedrockGuardrail(
|
||||
guardrail_name="bedrock-external-id",
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
guardrailIdentifier="gr-1",
|
||||
guardrailVersion="DRAFT",
|
||||
aws_region_name="us-east-1",
|
||||
aws_access_key_id="AKIAPODCALLERKEY",
|
||||
aws_secret_access_key="pod-caller-secret",
|
||||
aws_role_name="arn:aws:iam::999999999999:role/litellm-guardrail-role",
|
||||
aws_session_name="litellm-session",
|
||||
aws_external_id="external-id-123",
|
||||
)
|
||||
|
||||
with patch.object(boto3, "client", return_value=FakeSTSClient()):
|
||||
credentials, aws_region_name = guardrail._load_credentials()
|
||||
|
||||
assert credentials.access_key == "ASIAASSUMEDROLEKEY"
|
||||
assert credentials.token == "assumed-session-token"
|
||||
assert aws_region_name == "us-east-1"
|
||||
|
||||
|
||||
def test_initialize_bedrock_forwards_aws_external_id():
|
||||
"""aws_external_id configured on the guardrail must survive LitellmParams and the initializer."""
|
||||
from litellm.proxy.guardrails.guardrail_initializers import initialize_bedrock
|
||||
from litellm.types.guardrails import LitellmParams
|
||||
|
||||
litellm_params = LitellmParams(
|
||||
guardrail="bedrock",
|
||||
mode="pre_call",
|
||||
guardrailIdentifier="gr-1",
|
||||
guardrailVersion="DRAFT",
|
||||
aws_region_name="us-east-1",
|
||||
aws_role_name="arn:aws:iam::999999999999:role/litellm-guardrail-role",
|
||||
aws_external_id="external-id-123",
|
||||
)
|
||||
|
||||
guardrail = initialize_bedrock(litellm_params, {"guardrail_name": "bedrock-external-id"})
|
||||
try:
|
||||
assert guardrail.optional_params["aws_external_id"] == "external-id-123"
|
||||
finally:
|
||||
litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm.callbacks, guardrail)
|
||||
|
|
|
|||
|
|
@ -93,6 +93,32 @@ class TestExtractRequestToolNames:
|
|||
"run_sql",
|
||||
]
|
||||
|
||||
def test_anthropic_openai_format_tools_forwarded_by_bridge(self):
|
||||
data = {
|
||||
"tools": [
|
||||
{"type": "function", "function": {"name": "get_weather"}},
|
||||
{"name": "run_sql"},
|
||||
{"googleSearch": {}},
|
||||
]
|
||||
}
|
||||
assert extract_request_tool_names("/v1/messages", data) == [
|
||||
"get_weather",
|
||||
"run_sql",
|
||||
]
|
||||
|
||||
def test_anthropic_hybrid_tool_yields_every_name(self):
|
||||
data = {
|
||||
"tools": [
|
||||
{"type": "function", "name": "decoy", "function": {"name": "blocked_fn"}},
|
||||
{"type": "function", "name": "", "function": {"name": "hidden_fn"}},
|
||||
]
|
||||
}
|
||||
assert extract_request_tool_names("/v1/messages", data) == [
|
||||
"decoy",
|
||||
"blocked_fn",
|
||||
"hidden_fn",
|
||||
]
|
||||
|
||||
def test_generate_content_tools(self):
|
||||
data = {
|
||||
"tools": [
|
||||
|
|
@ -159,6 +185,34 @@ class TestCheckToolsAllowlist:
|
|||
assert exc_info.value.type == ProxyErrorTypes.tool_access_denied
|
||||
assert "get_weather" in str(exc_info.value.message)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disallowed_openai_format_tool_raises_on_messages_route(self):
|
||||
token = _token(metadata={"allowed_tools": ["other_tool"]})
|
||||
body = {"tools": [{"type": "function", "function": {"name": "get_weather"}}]}
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await check_tools_allowlist(
|
||||
request_body=body,
|
||||
valid_token=token,
|
||||
team_object=None,
|
||||
route="/v1/messages",
|
||||
)
|
||||
assert exc_info.value.type == ProxyErrorTypes.tool_access_denied
|
||||
assert "get_weather" in str(exc_info.value.message)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hybrid_tool_with_decoy_name_raises_on_messages_route(self):
|
||||
token = _token(metadata={"allowed_tools": ["decoy"]})
|
||||
body = {"tools": [{"type": "function", "name": "decoy", "function": {"name": "run_sql"}}]}
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await check_tools_allowlist(
|
||||
request_body=body,
|
||||
valid_token=token,
|
||||
team_object=None,
|
||||
route="/v1/messages",
|
||||
)
|
||||
assert exc_info.value.type == ProxyErrorTypes.tool_access_denied
|
||||
assert "run_sql" in str(exc_info.value.message)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disallowed_custom_tool_raises_on_responses_route(self):
|
||||
token = _token(metadata={"allowed_tools": ["other_tool"]})
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.router_utils.reasoning_effort_capability import (
|
||||
deployment_is_catalog_mapped,
|
||||
intersect_supported_reasoning_efforts,
|
||||
|
|
@ -196,3 +197,158 @@ class TestIntersectSupportedReasoningEfforts:
|
|||
|
||||
def test_disjoint_sets_intersect_to_empty(self):
|
||||
assert intersect_supported_reasoning_efforts(["max"], ["minimal"]) == ()
|
||||
|
||||
|
||||
class TestDeclaredEffortList:
|
||||
"""reasoning_effort_levels is what the catalog DECLARES per deployment;
|
||||
ModelGroupInfo.supported_reasoning_efforts is what a group COMPUTED. test_router.py pins that
|
||||
the computed one is never seeded from model_info, so the two names must stay apart."""
|
||||
|
||||
def test_a_declared_list_answers_where_no_flag_could(self):
|
||||
"""No flag can drop medium, so before this key the entry could only stay silent or
|
||||
over-advertise a level the model does not document."""
|
||||
resolved = resolve_supported_reasoning_efforts(
|
||||
{"supports_reasoning": True, "reasoning_effort_levels": ["low", "high", "max"]},
|
||||
deployment_is_mapped=True,
|
||||
)
|
||||
assert resolved == ("low", "high", "max")
|
||||
|
||||
def test_a_declared_list_wins_whole_over_the_flags(self):
|
||||
resolved = resolve_supported_reasoning_efforts(
|
||||
{
|
||||
"supports_reasoning": True,
|
||||
"reasoning_effort_levels": ["low", "high", "max"],
|
||||
"supports_none_reasoning_effort": True,
|
||||
"supports_minimal_reasoning_effort": True,
|
||||
"supports_xhigh_reasoning_effort": True,
|
||||
"supports_max_reasoning_effort": False,
|
||||
},
|
||||
deployment_is_mapped=True,
|
||||
)
|
||||
assert resolved == ("low", "high", "max")
|
||||
|
||||
def test_a_declaration_is_reordered_into_the_advertisement_order(self):
|
||||
resolved = resolve_supported_reasoning_efforts(
|
||||
{"supports_reasoning": True, "reasoning_effort_levels": ["max", "low", "high"]},
|
||||
deployment_is_mapped=True,
|
||||
)
|
||||
assert resolved == ("low", "high", "max")
|
||||
|
||||
def test_a_declared_empty_list_empties_the_group(self):
|
||||
assert (
|
||||
resolve_supported_reasoning_efforts(
|
||||
{"supports_reasoning": True, "reasoning_effort_levels": []},
|
||||
deployment_is_mapped=True,
|
||||
)
|
||||
== ()
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("declared", [["low", "bogus"], ["bogus"], ["low", 7, None]])
|
||||
def test_an_unknown_level_is_dropped_rather_than_raised(self, declared):
|
||||
"""A config.yaml model_info block bypasses the map's enum schema, and one mistyped level
|
||||
must not fail every sibling on the proxy."""
|
||||
resolved = resolve_supported_reasoning_efforts(
|
||||
{"supports_reasoning": True, "reasoning_effort_levels": declared},
|
||||
deployment_is_mapped=True,
|
||||
)
|
||||
assert resolved == tuple(effort for effort in ("low",) if effort in declared)
|
||||
|
||||
@pytest.mark.parametrize("malformed", ["low,high,max", {"low": True}, 3, True])
|
||||
def test_a_malformed_declaration_falls_through_to_the_flags(self, malformed):
|
||||
resolved = resolve_supported_reasoning_efforts(
|
||||
{
|
||||
"supports_reasoning": True,
|
||||
"reasoning_effort_levels": malformed,
|
||||
"supports_max_reasoning_effort": True,
|
||||
},
|
||||
deployment_is_mapped=True,
|
||||
)
|
||||
assert resolved == ("none", "minimal", "low", "medium", "high", "max")
|
||||
|
||||
def test_a_model_the_map_calls_non_reasoning_ignores_its_declaration(self):
|
||||
assert (
|
||||
resolve_supported_reasoning_efforts(
|
||||
{"supports_reasoning": False, "reasoning_effort_levels": ["low", "high", "max"]},
|
||||
deployment_is_mapped=True,
|
||||
)
|
||||
== ()
|
||||
)
|
||||
|
||||
def test_a_declaration_is_read_through_the_bare_twin(self, monkeypatch):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"some-declared-reasoner",
|
||||
{"supports_reasoning": True, "reasoning_effort_levels": ["low", "max"]},
|
||||
)
|
||||
resolved = resolve_supported_reasoning_efforts(
|
||||
{
|
||||
"supports_reasoning": True,
|
||||
"litellm_provider": "openai",
|
||||
"key": "openai/some-declared-reasoner",
|
||||
},
|
||||
deployment_is_mapped=True,
|
||||
)
|
||||
assert resolved == ("low", "max")
|
||||
|
||||
|
||||
KIMI_K3_PASSTHROUGH_KEYS = (
|
||||
"azure_ai/FW-Kimi-K3",
|
||||
"moonshot/kimi-k3",
|
||||
"together_ai/moonshotai/Kimi-K3",
|
||||
"fireworks_ai/kimi-k3",
|
||||
"fireworks_ai/kimi-k3-fast",
|
||||
"fireworks_ai/kimi-k3-us",
|
||||
"fireworks_ai/accounts/fireworks/models/kimi-k3",
|
||||
"fireworks_ai/accounts/fireworks/routers/kimi-k3-fast",
|
||||
"fireworks_ai/accounts/fireworks/routers/kimi-k3-us",
|
||||
)
|
||||
KIMI_K3_PERPLEXITY_KEY = "perplexity/perplexity/kimi-k3"
|
||||
|
||||
|
||||
class TestKimiK3AdvertisesItsDocumentedLevels:
|
||||
@pytest.mark.parametrize("model_key", KIMI_K3_PASSTHROUGH_KEYS)
|
||||
def test_a_passthrough_entry_advertises_the_models_own_levels(self, local_model_cost_map, model_key):
|
||||
"""platform.kimi.ai documents exactly low, high and max, and these providers forward the
|
||||
level unchanged. Undeclared, each entry resolves to unknown and the dashboard falls back to
|
||||
a capability-blind list that omits max."""
|
||||
entry = dict(litellm.model_cost[model_key], key=model_key)
|
||||
|
||||
assert resolve_supported_reasoning_efforts(entry, deployment_is_mapped=True) == ("low", "high", "max")
|
||||
|
||||
def test_the_perplexity_entry_advertises_the_wider_set_it_maps_down(self, local_model_cost_map):
|
||||
"""Perplexity's Agent API takes a six-value enum and maps it down internally, so this
|
||||
deployment is legitimately wider than a passthrough. One blanket list could not say both."""
|
||||
entry = dict(litellm.model_cost[KIMI_K3_PERPLEXITY_KEY], key=KIMI_K3_PERPLEXITY_KEY)
|
||||
|
||||
assert resolve_supported_reasoning_efforts(entry, deployment_is_mapped=True) == (
|
||||
"minimal",
|
||||
"low",
|
||||
"medium",
|
||||
"high",
|
||||
"xhigh",
|
||||
"max",
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("model, provider", [("kimi-k3", "moonshot"), ("kimi-k3", "fireworks_ai")])
|
||||
def test_the_declaration_survives_model_info_hydration(self, local_model_cost_map, model, provider):
|
||||
"""The hydration line is the load-bearing seam: without it the key the map carries never
|
||||
reaches the resolver and reads as absent everywhere downstream."""
|
||||
from litellm.utils import _get_model_info_helper
|
||||
|
||||
model_info = dict(_get_model_info_helper(model=model, custom_llm_provider=provider))
|
||||
|
||||
assert model_info["reasoning_effort_levels"] == ["low", "high", "max"]
|
||||
assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == ("low", "high", "max")
|
||||
|
||||
def test_a_kimi_k3_deployment_now_narrows_a_mixed_group(self, local_model_cost_map):
|
||||
"""kimi used to contribute unknown, which never narrows, so the group advertised whatever
|
||||
its other deployments agreed on."""
|
||||
kimi = resolve_supported_reasoning_efforts(
|
||||
dict(litellm.model_cost["fireworks_ai/kimi-k3"], key="fireworks_ai/kimi-k3"),
|
||||
deployment_is_mapped=True,
|
||||
)
|
||||
|
||||
assert intersect_supported_reasoning_efforts(("none", "minimal", "low", "medium", "high", "xhigh"), kimi) == (
|
||||
"low",
|
||||
"high",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3977,6 +3977,74 @@ def test_completion_cost_prices_anthropic_shaped_cache_read_tokens(_local_model_
|
|||
assert cost == pytest.approx(3 * 4e-6 + 4014 * 4e-7 + 5 * 2e-5, rel=1e-9)
|
||||
|
||||
|
||||
def _together_chat_response(model: str, prompt_tokens: int, completion_tokens: int, cached_tokens: int) -> ModelResponse:
|
||||
return ModelResponse(
|
||||
id="chatcmpl-together-cache",
|
||||
choices=[{"finish_reason": "stop", "index": 0, "message": {"content": "acknowledged", "role": "assistant"}}],
|
||||
created=1756164000,
|
||||
model=model,
|
||||
object="chat.completion",
|
||||
usage=Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_tokens + completion_tokens,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cached_tokens),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_completion_cost_prices_together_cached_tokens_at_cache_read_rate(_local_model_cost_map):
|
||||
"""Regression: Together reports prompt_tokens_details.cached_tokens but no together_ai
|
||||
registry entry carried cache_read_input_token_cost, so cache-hit tokens were priced at
|
||||
0.0 and spend on cache-heavy workloads was understated."""
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=_together_chat_response(
|
||||
model="deepseek-ai/DeepSeek-V4-Flash-0731", prompt_tokens=7864, completion_tokens=16, cached_tokens=7863
|
||||
),
|
||||
custom_llm_provider="together_ai",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx(1 * 1.4e-07 + 7863 * 3e-08 + 16 * 2.8e-07, rel=1e-9)
|
||||
|
||||
|
||||
def test_completion_cost_together_mapped_model_skips_size_bucket(_local_model_cost_map):
|
||||
"""Regression: any together model whose name matches (\\d+b) was rewritten to a
|
||||
together-ai-* size bucket before the registry lookup, so mapped models like
|
||||
Muse-Glimmer-30B never used their per-model rates, cache fields included."""
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=_together_chat_response(
|
||||
model="meta-models/Muse-Glimmer-30B", prompt_tokens=63, completion_tokens=16, cached_tokens=0
|
||||
),
|
||||
custom_llm_provider="together_ai",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx(63 * 3.5e-07 + 16 * 1.5e-06, rel=1e-9)
|
||||
|
||||
|
||||
def test_completion_cost_together_unmapped_model_still_uses_size_bucket(_local_model_cost_map):
|
||||
cost = completion_cost(
|
||||
completion_response=_together_chat_response(
|
||||
model="qwen/Qwen2-72B-Instruct", prompt_tokens=23, completion_tokens=15, cached_tokens=0
|
||||
),
|
||||
custom_llm_provider="together_ai",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx((23 + 15) * 9e-07, rel=1e-9)
|
||||
|
||||
|
||||
def test_completion_cost_together_metadata_only_model_still_uses_size_bucket(_local_model_cost_map):
|
||||
assert "input_cost_per_token" not in litellm.model_cost["together_ai/togethercomputer/CodeLlama-34b-Instruct"]
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=_together_chat_response(
|
||||
model="togethercomputer/CodeLlama-34b-Instruct", prompt_tokens=23, completion_tokens=15, cached_tokens=0
|
||||
),
|
||||
custom_llm_provider="together_ai",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx((23 + 15) * 8e-07, rel=1e-9)
|
||||
def test_select_model_name_strips_unregistered_alias_prefix(_local_model_cost_map):
|
||||
"""A router-facing model_name alias containing "/" whose leading segment is NOT a
|
||||
registered provider must not be double-prefixed into a non-existent cost key.
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""
|
||||
Unit tests for DashScope image generation support (qwen-image-2.0, qwen-image-2.0-pro).
|
||||
Unit tests for DashScope image generation support (qwen-image-2.0, qwen-image-2.0-pro,
|
||||
qwen-image-3.0, qwen-image-3.0-pro).
|
||||
|
||||
Run in docker: pytest tests/test_litellm/test_dashscope_image_generation.py -v
|
||||
"""
|
||||
|
|
@ -30,6 +31,8 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
|||
[
|
||||
"dashscope/qwen-image-2.0",
|
||||
"dashscope/qwen-image-2.0-pro",
|
||||
"dashscope/qwen-image-3.0",
|
||||
"dashscope/qwen-image-3.0-pro",
|
||||
],
|
||||
)
|
||||
def test_get_llm_provider_returns_dashscope(model_string: str):
|
||||
|
|
@ -48,6 +51,8 @@ def test_get_llm_provider_returns_dashscope(model_string: str):
|
|||
[
|
||||
("dashscope/qwen-image-2.0", "dashscope"),
|
||||
("dashscope/qwen-image-2.0-pro", "dashscope"),
|
||||
("dashscope/qwen-image-3.0", "dashscope"),
|
||||
("dashscope/qwen-image-3.0-pro", "dashscope"),
|
||||
],
|
||||
)
|
||||
def test_get_model_info_mode_is_image_generation(
|
||||
|
|
@ -93,6 +98,19 @@ class TestDashScopeImageGenerationConfig:
|
|||
url = self.cfg.get_complete_url(custom, None, "qwen-image-2.0", {}, {})
|
||||
assert url == custom
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"chat_api_base",
|
||||
[
|
||||
"https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||||
"https://dashscope-intl.aliyuncs.com/compatible-mode/v1/",
|
||||
],
|
||||
)
|
||||
def test_get_complete_url_ignores_chat_compatible_mode_base(
|
||||
self, chat_api_base: str
|
||||
):
|
||||
url = self.cfg.get_complete_url(chat_api_base, None, "qwen-image-3.0", {}, {})
|
||||
assert url == DEFAULT_API_BASE
|
||||
|
||||
def test_validate_environment_sets_auth_header(self):
|
||||
headers = self.cfg.validate_environment(
|
||||
headers={},
|
||||
|
|
@ -135,6 +153,27 @@ class TestDashScopeImageGenerationConfig:
|
|||
assert messages[0]["content"][0]["text"] == "a puppy on green grass"
|
||||
assert req["parameters"]["size"] == "1024*1024"
|
||||
|
||||
@pytest.mark.parametrize("model", ["qwen-image-3.0", "qwen-image-3.0-pro"])
|
||||
def test_transform_request_qwen_image_3(self, model: str):
|
||||
req = self.cfg.transform_image_generation_request(
|
||||
model=model,
|
||||
prompt="a poster with small multilingual text",
|
||||
optional_params=self.cfg.map_openai_params(
|
||||
non_default_params={"size": "2048x2048", "n": 6},
|
||||
optional_params={},
|
||||
model=model,
|
||||
drop_params=False,
|
||||
),
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert req["model"] == model
|
||||
assert req["input"]["messages"][0]["content"][0]["text"] == (
|
||||
"a poster with small multilingual text"
|
||||
)
|
||||
assert req["parameters"]["size"] == "2048*2048"
|
||||
assert req["parameters"]["n"] == 6
|
||||
|
||||
def test_transform_request_empty_params(self):
|
||||
req = self.cfg.transform_image_generation_request(
|
||||
model="qwen-image-2.0-pro",
|
||||
|
|
@ -238,6 +277,48 @@ class TestDashScopeImageGenerationConfig:
|
|||
assert result.data[0].url == "https://example.com/img1.png"
|
||||
assert result.data[1].url == "https://example.com/img2.png"
|
||||
|
||||
def test_transform_response_multiple_images_in_one_choice(self):
|
||||
body = {
|
||||
"output": {
|
||||
"choices": [
|
||||
{
|
||||
"finish_reason": "stop",
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"image": "https://example.com/img1.png", "type": "image"},
|
||||
{"image": "https://example.com/img2.png", "type": "image"},
|
||||
],
|
||||
},
|
||||
}
|
||||
]
|
||||
},
|
||||
"usage": {
|
||||
"output_width": 1024,
|
||||
"output_height": 1024,
|
||||
"output_image_count": 2,
|
||||
},
|
||||
}
|
||||
mock_resp = MagicMock(spec=httpx.Response)
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.headers = {}
|
||||
mock_resp.json.return_value = body
|
||||
|
||||
result = self.cfg.transform_image_generation_response(
|
||||
model="qwen-image-3.0",
|
||||
raw_response=mock_resp,
|
||||
model_response=ImageResponse(),
|
||||
logging_obj=MagicMock(),
|
||||
request_data={},
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
assert [image.url for image in result.data] == [
|
||||
"https://example.com/img1.png",
|
||||
"https://example.com/img2.png",
|
||||
]
|
||||
|
||||
def test_transform_response_raises_on_non_200_status(self):
|
||||
mock_resp = MagicMock(spec=httpx.Response)
|
||||
mock_resp.status_code = 400
|
||||
|
|
@ -294,14 +375,14 @@ class TestDashScopeImageGenerationConfig:
|
|||
)
|
||||
assert mapped["size"] == "1024*1024"
|
||||
|
||||
def test_map_openai_params_n_to_image_count(self):
|
||||
def test_map_openai_params_n_passthrough(self):
|
||||
mapped = self.cfg.map_openai_params(
|
||||
non_default_params={"n": 2},
|
||||
optional_params={},
|
||||
model="qwen-image-2.0",
|
||||
drop_params=False,
|
||||
)
|
||||
assert mapped["image_count"] == 2
|
||||
assert mapped == {"n": 2}
|
||||
|
||||
def test_map_openai_params_unknown_size_uses_asterisk(self):
|
||||
mapped = self.cfg.map_openai_params(
|
||||
|
|
@ -338,7 +419,15 @@ class TestDashScopeImageGenerationConfig:
|
|||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_litellm_image_generation_dashscope_end_to_end():
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"dashscope/qwen-image-2.0",
|
||||
"dashscope/qwen-image-3.0",
|
||||
"dashscope/qwen-image-3.0-pro",
|
||||
],
|
||||
)
|
||||
def test_litellm_image_generation_dashscope_end_to_end(model: str):
|
||||
mock_response_body = {
|
||||
"output": {
|
||||
"choices": [
|
||||
|
|
@ -374,7 +463,7 @@ def test_litellm_image_generation_dashscope_end_to_end():
|
|||
mock_post.return_value = mock_http_response
|
||||
|
||||
response = litellm.image_generation(
|
||||
model="dashscope/qwen-image-2.0",
|
||||
model=model,
|
||||
prompt="a puppy playing on green grass",
|
||||
api_key="sk-test-key",
|
||||
size="1024x1024",
|
||||
|
|
@ -392,7 +481,7 @@ def test_litellm_image_generation_dashscope_end_to_end():
|
|||
called_url = (
|
||||
call_args[0][0] if call_args[0] else call_args.kwargs.get("url", "")
|
||||
)
|
||||
assert "dashscope" in called_url or "aliyuncs" in called_url
|
||||
assert called_url == DEFAULT_API_BASE
|
||||
|
||||
# Verify request body contains DashScope format
|
||||
call_kwargs = call_args[1] if call_args[1] else {}
|
||||
|
|
@ -400,3 +489,4 @@ def test_litellm_image_generation_dashscope_end_to_end():
|
|||
body = call_kwargs["json"]
|
||||
assert "input" in body
|
||||
assert "messages" in body["input"]
|
||||
assert body["parameters"]["size"] == "1024*1024"
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ COST_MAP_ADAPTER: Final = TypeAdapter(CostMap)
|
|||
SERVERLESS_CHAT_MODELS: Final = (
|
||||
"together_ai/moonshotai/Kimi-K3",
|
||||
"together_ai/zai-org/GLM-5.2",
|
||||
"together_ai/zai-org/GLM-5.3-Flash",
|
||||
"together_ai/deepseek-ai/DeepSeek-V4-Pro-0813",
|
||||
"together_ai/deepseek-ai/DeepSeek-V4-Flash-0731",
|
||||
"together_ai/MiniMaxAI/MiniMax-M3",
|
||||
|
|
@ -110,6 +111,22 @@ def test_together_glm_52_pricing(cost_map: CostMap):
|
|||
assert info["supports_reasoning"] is True
|
||||
|
||||
|
||||
def test_together_glm_53_flash_pricing_and_capabilities(cost_map: CostMap):
|
||||
info = cost_map["together_ai/zai-org/GLM-5.3-Flash"]
|
||||
assert info["input_cost_per_token"] == 1.5e-07
|
||||
assert info["output_cost_per_token"] == 5e-07
|
||||
assert info["cache_read_input_token_cost"] == 3e-08
|
||||
assert info["max_input_tokens"] == 1048575
|
||||
assert info["max_output_tokens"] == 1048575
|
||||
assert info["supports_function_calling"] is True
|
||||
assert info["supports_parallel_function_calling"] is True
|
||||
assert info["supports_prompt_caching"] is True
|
||||
assert info["supports_tool_choice"] is True
|
||||
assert info["supports_response_schema"] is True
|
||||
assert info["supports_vision"] is True
|
||||
assert info["supports_reasoning"] is True
|
||||
|
||||
|
||||
def test_together_multilingual_e5_embedding_entry(cost_map: CostMap):
|
||||
info = cost_map["together_ai/intfloat/multilingual-e5-large-instruct"]
|
||||
assert info["mode"] == "embedding"
|
||||
|
|
@ -159,3 +176,51 @@ def test_together_backup_cost_map_in_sync(cost_map: CostMap):
|
|||
together_main = {k: v for k, v in cost_map.items() if k.startswith("together_ai/")}
|
||||
together_backup = {k: v for k, v in backup.items() if k.startswith("together_ai/")}
|
||||
assert together_backup == together_main
|
||||
|
||||
|
||||
CACHED_INPUT_MODELS: Final = (
|
||||
"together_ai/moonshotai/Kimi-K3",
|
||||
"together_ai/zai-org/GLM-5.2",
|
||||
"together_ai/meta-models/Muse-Glimmer-30B",
|
||||
"together_ai/Qwen/Qwen3.8-2.4T-A95B",
|
||||
"together_ai/deepseek-ai/DeepSeek-V4-Pro-0813",
|
||||
"together_ai/deepseek-ai/DeepSeek-V4-Flash-0731",
|
||||
"together_ai/thinkingmachines/Inkling",
|
||||
"together_ai/MiniMaxAI/MiniMax-M3",
|
||||
"together_ai/thinkingmachines/Inkling-Small",
|
||||
"together_ai/moonshotai/Kimi-K2.7-Code",
|
||||
"together_ai/deepseek-ai/DeepSeek-V4-Pro",
|
||||
"together_ai/nvidia/nemotron-3-ultra-550b-a55b",
|
||||
"together_ai/Qwen/Qwen3.7-Max",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", CACHED_INPUT_MODELS)
|
||||
def test_together_cached_input_model_carries_cache_read_pricing(cost_map: CostMap, model: str):
|
||||
info = cost_map.get(model)
|
||||
assert info is not None, f"{model} missing from model_prices_and_context_window.json"
|
||||
assert info.get("supports_prompt_caching") is True
|
||||
cache_read = info.get("cache_read_input_token_cost")
|
||||
assert isinstance(cache_read, float)
|
||||
assert 0 < cache_read < info["input_cost_per_token"]
|
||||
assert "cache_creation_input_token_cost" not in info
|
||||
|
||||
|
||||
def test_together_prompt_caching_flag_implies_cache_read_rate(cost_map: CostMap):
|
||||
for model, info in cost_map.items():
|
||||
if model.startswith("together_ai/") and info.get("supports_prompt_caching"):
|
||||
assert "cache_read_input_token_cost" in info, f"{model} flags caching without a cache read rate"
|
||||
|
||||
|
||||
def test_together_deepseek_v4_flash_cache_read_rate(cost_map: CostMap):
|
||||
info = cost_map["together_ai/deepseek-ai/DeepSeek-V4-Flash-0731"]
|
||||
assert info["input_cost_per_token"] == 1.4e-07
|
||||
assert info["cache_read_input_token_cost"] == 3e-08
|
||||
assert info["output_cost_per_token"] == 2.8e-07
|
||||
|
||||
|
||||
def test_together_qwen_37_max_repriced_to_current_together_rate(cost_map: CostMap):
|
||||
info = cost_map["together_ai/Qwen/Qwen3.7-Max"]
|
||||
assert info["input_cost_per_token"] == 2.5e-06
|
||||
assert info["output_cost_per_token"] == 7.5e-06
|
||||
assert info["cache_read_input_token_cost"] == 5e-07
|
||||
|
|
|
|||
|
|
@ -1014,6 +1014,10 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
|||
"supports_none_reasoning_effort": {"type": "boolean"},
|
||||
"supports_xhigh_reasoning_effort": {"type": "boolean"},
|
||||
"supports_max_reasoning_effort": {"type": "boolean"},
|
||||
"reasoning_effort_levels": {
|
||||
"type": "array",
|
||||
"items": {"type": "string", "enum": ["none", "minimal", "low", "medium", "high", "xhigh", "max"]},
|
||||
},
|
||||
"supports_adaptive_thinking": {"type": "boolean"},
|
||||
"supports_legacy_thinking": {"type": "boolean"},
|
||||
"thinking_always_on": {"type": "boolean"},
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
"limit": 22733
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26863
|
||||
"limit": 26860
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 269
|
||||
|
|
@ -33,6 +33,6 @@
|
|||
"limit": 5583
|
||||
},
|
||||
"LIT012": {
|
||||
"limit": 4510
|
||||
"limit": 4509
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import { render, screen } from "@testing-library/react";
|
||||
import { render, screen, within } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { useState } from "react";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { useInfiniteTeams } from "@/app/(dashboard)/hooks/teams/useTeams";
|
||||
import TeamMultiSelect from "./team_multi_select";
|
||||
|
|
@ -58,9 +59,9 @@ describe("TeamMultiSelect", () => {
|
|||
await user.click(combobox());
|
||||
|
||||
expect(screen.getByText("Alpha Team")).toBeInTheDocument();
|
||||
expect(screen.getByText("(team-1)")).toBeInTheDocument();
|
||||
expect(screen.getByText("team-1")).toBeInTheDocument();
|
||||
expect(screen.getByText("Beta Team")).toBeInTheDocument();
|
||||
expect(screen.getByText("(team-2)")).toBeInTheDocument();
|
||||
expect(screen.getByText("team-2")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("deduplicates a team that appears on more than one page", async () => {
|
||||
|
|
@ -116,6 +117,29 @@ describe("TeamMultiSelect", () => {
|
|||
expect(screen.getByText(/No teams found/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("keeps a picked team's alias on its chip once a later search drops it from the loaded page", async () => {
|
||||
const user = userEvent.setup();
|
||||
|
||||
function Controlled() {
|
||||
const [value, setValue] = useState<string[]>([]);
|
||||
return <TeamMultiSelect value={value} onChange={setValue} />;
|
||||
}
|
||||
const { rerender } = render(<Controlled />);
|
||||
|
||||
await user.click(combobox());
|
||||
const matches = screen.getAllByText("Beta Team");
|
||||
await user.click(matches[matches.length - 1]);
|
||||
|
||||
mockUseInfiniteTeams.mockReturnValue(
|
||||
mockTeamsResult({ pages: [{ teams: [team("team-3", "Gamma Team")] }] }) as never,
|
||||
);
|
||||
rerender(<Controlled />);
|
||||
|
||||
const chips = document.querySelector('[data-slot="combobox-chips"]') as HTMLElement;
|
||||
expect(within(chips).getByText("Beta Team")).toBeInTheDocument();
|
||||
expect(within(chips).queryByText("team-2")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("passes the page size and organization filter through to the teams query", () => {
|
||||
render(<TeamMultiSelect pageSize={25} organizationId="org-7" />);
|
||||
|
||||
|
|
|
|||
|
|
@ -1,22 +1,7 @@
|
|||
import React, { useMemo, useState, type UIEvent } from "react";
|
||||
import { Loader2 } from "lucide-react";
|
||||
import { useDebouncedCallback } from "@tanstack/react-pacer/debouncer";
|
||||
import {
|
||||
Combobox,
|
||||
ComboboxChip,
|
||||
ComboboxChips,
|
||||
ComboboxChipsInput,
|
||||
ComboboxClear,
|
||||
ComboboxContent,
|
||||
ComboboxEmpty,
|
||||
ComboboxItem,
|
||||
ComboboxList,
|
||||
ComboboxValue,
|
||||
useComboboxAnchor,
|
||||
} from "@/components/ui/combobox";
|
||||
import React, { useMemo, useState } from "react";
|
||||
import { PaginatedMultiSelect } from "@/components/shared/PaginatedMultiSelect";
|
||||
import type { SearchSelectOption } from "@/components/shared/SearchSelect";
|
||||
import { useInfiniteTeams } from "@/app/(dashboard)/hooks/teams/useTeams";
|
||||
import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants";
|
||||
import { Team } from "../key_team_helpers/key_list";
|
||||
|
||||
interface TeamMultiSelectProps {
|
||||
value?: string[];
|
||||
|
|
@ -27,8 +12,6 @@ interface TeamMultiSelectProps {
|
|||
placeholder?: string;
|
||||
}
|
||||
|
||||
const SCROLL_THRESHOLD = 0.8;
|
||||
|
||||
const TeamMultiSelect: React.FC<TeamMultiSelectProps> = ({
|
||||
value = [],
|
||||
onChange,
|
||||
|
|
@ -37,9 +20,7 @@ const TeamMultiSelect: React.FC<TeamMultiSelectProps> = ({
|
|||
pageSize = 20,
|
||||
placeholder = "Search teams by alias...",
|
||||
}) => {
|
||||
const anchor = useComboboxAnchor();
|
||||
const [search, setSearch] = useState("");
|
||||
const debouncedSetSearch = useDebouncedCallback(setSearch, { wait: DEBOUNCE_WAIT_MS });
|
||||
|
||||
const { data, fetchNextPage, hasNextPage, isFetchingNextPage, isLoading } = useInfiniteTeams(
|
||||
pageSize,
|
||||
|
|
@ -47,68 +28,40 @@ const TeamMultiSelect: React.FC<TeamMultiSelectProps> = ({
|
|||
organizationId,
|
||||
);
|
||||
|
||||
const teamById = useMemo(
|
||||
const options = useMemo<SearchSelectOption[]>(
|
||||
() =>
|
||||
new Map<string, Team>(
|
||||
(data?.pages ?? []).flatMap((page) => page.teams).map((team) => [team.team_id, team] as const),
|
||||
Array.from(
|
||||
new Map(
|
||||
(data?.pages ?? [])
|
||||
.flatMap((page) => page.teams)
|
||||
.map(
|
||||
(team) =>
|
||||
[
|
||||
team.team_id,
|
||||
{ label: team.team_alias || team.team_id, value: team.team_id, sublabel: team.team_id },
|
||||
] as const,
|
||||
),
|
||||
).values(),
|
||||
),
|
||||
[data],
|
||||
);
|
||||
const teamIds = useMemo(() => Array.from(teamById.keys()), [teamById]);
|
||||
|
||||
const aliasOf = (teamId: string) => teamById.get(teamId)?.team_alias ?? teamId;
|
||||
|
||||
const handleScroll = (event: UIEvent<HTMLDivElement>) => {
|
||||
const target = event.currentTarget;
|
||||
if (target.scrollHeight === 0) return;
|
||||
const scrollRatio = (target.scrollTop + target.clientHeight) / target.scrollHeight;
|
||||
if (scrollRatio >= SCROLL_THRESHOLD && hasNextPage && !isFetchingNextPage) {
|
||||
fetchNextPage();
|
||||
}
|
||||
};
|
||||
|
||||
return (
|
||||
<Combobox
|
||||
multiple
|
||||
items={teamIds}
|
||||
<PaginatedMultiSelect
|
||||
options={options}
|
||||
value={value}
|
||||
onValueChange={(next: string[]) => onChange?.(next)}
|
||||
filter={null}
|
||||
onInputValueChange={debouncedSetSearch}
|
||||
onSearchChange={setSearch}
|
||||
onLoadMore={fetchNextPage}
|
||||
hasNextPage={hasNextPage}
|
||||
isLoading={isLoading}
|
||||
isFetchingNextPage={isFetchingNextPage}
|
||||
placeholder={placeholder}
|
||||
emptyText="No teams found"
|
||||
loadingText="Loading teams..."
|
||||
clearAllLabel="Clear all teams"
|
||||
disabled={disabled}
|
||||
>
|
||||
<ComboboxChips render={<div ref={anchor} />} className="w-full" aria-busy={isLoading}>
|
||||
<ComboboxValue>
|
||||
{(selected: string[]) =>
|
||||
selected.map((teamId) => (
|
||||
<ComboboxChip key={teamId} aria-label={aliasOf(teamId)}>
|
||||
{aliasOf(teamId)}
|
||||
</ComboboxChip>
|
||||
))
|
||||
}
|
||||
</ComboboxValue>
|
||||
<ComboboxChipsInput placeholder={placeholder} aria-label={placeholder} disabled={disabled} />
|
||||
{value.length > 0 && <ComboboxClear aria-label="Clear all teams" disabled={disabled} />}
|
||||
</ComboboxChips>
|
||||
<ComboboxContent anchor={anchor}>
|
||||
<ComboboxEmpty>
|
||||
{isLoading ? <Loader2 className="size-4 animate-spin text-muted-foreground" /> : "No teams found"}
|
||||
</ComboboxEmpty>
|
||||
<ComboboxList onScroll={handleScroll}>
|
||||
{(teamId: string) => (
|
||||
<ComboboxItem key={teamId} value={teamId}>
|
||||
<span className="font-medium">{aliasOf(teamId)}</span>{" "}
|
||||
<span className="text-muted-foreground">({teamId})</span>
|
||||
</ComboboxItem>
|
||||
)}
|
||||
</ComboboxList>
|
||||
{isFetchingNextPage && (
|
||||
<div className="flex justify-center py-2">
|
||||
<Loader2 className="size-4 animate-spin text-muted-foreground" />
|
||||
</div>
|
||||
)}
|
||||
</ComboboxContent>
|
||||
</Combobox>
|
||||
/>
|
||||
);
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -199,6 +199,52 @@ describe("UserSearchModal submit payload", () => {
|
|||
});
|
||||
});
|
||||
|
||||
describe("UserSearchModal search lifecycle", () => {
|
||||
const directory = [
|
||||
{ user_id: "u-jones", user_email: "alice.jones@example.com" },
|
||||
{ user_id: "u-smith", user_email: "alice.smith@example.com" },
|
||||
{ user_id: "u-bob", user_email: "bob@example.com" },
|
||||
];
|
||||
|
||||
beforeEach(() => {
|
||||
vi.mocked(userFilterUICall).mockReset();
|
||||
vi.mocked(userFilterUICall).mockImplementation((_accessToken, params) => {
|
||||
const query = params.get("user_email") ?? "";
|
||||
return Promise.resolve(directory.filter((user) => user.user_email.includes(query))) as never;
|
||||
});
|
||||
});
|
||||
|
||||
const searchedFor = (): string[] =>
|
||||
vi.mocked(userFilterUICall).mock.calls.map((call) => {
|
||||
const email = call[1].get("user_email");
|
||||
return email === null ? `user_id=${call[1].get("user_id")}` : `user_email=${email}`;
|
||||
});
|
||||
|
||||
const settleDebounce = () =>
|
||||
act(async () => {
|
||||
await new Promise((resolve) => setTimeout(resolve, DEBOUNCE_WAIT_MS + 100));
|
||||
});
|
||||
|
||||
it("leaves the search unfiltered after a pick, so reopening searches the newly typed text", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<UserSearchModal isVisible onCancel={vi.fn()} onSubmit={vi.fn()} accessToken="sk-test" />);
|
||||
|
||||
const input = getEmailSearchInput();
|
||||
await user.click(input);
|
||||
await user.type(input, "ali");
|
||||
await user.click(await screen.findByRole("option", { name: "alice.jones@example.com" }));
|
||||
|
||||
await settleDebounce();
|
||||
expect(searchedFor()).toEqual(["user_email=ali"]);
|
||||
|
||||
await user.click(input);
|
||||
await user.type(input, "bob");
|
||||
|
||||
expect(await screen.findByRole("option", { name: "bob@example.com" })).toBeInTheDocument();
|
||||
expect(searchedFor()).toEqual(["user_email=ali", "user_email=bob"]);
|
||||
});
|
||||
});
|
||||
|
||||
describe("UserSearchModal out-of-order search results", () => {
|
||||
const answers = new Map<string, (users: { user_id: string; user_email: string }[]) => void>();
|
||||
|
||||
|
|
|
|||
|
|
@ -1,21 +1,12 @@
|
|||
import { useRef, useState } from "react";
|
||||
import { Info, UserPlus } from "lucide-react";
|
||||
import { Alert, AlertTitle } from "@/components/shared/Alert";
|
||||
import { useDebouncedCallback } from "@tanstack/react-pacer/debouncer";
|
||||
import { useForm } from "react-hook-form";
|
||||
import { userFilterUICall } from "@/components/networking";
|
||||
import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants";
|
||||
import { FieldGroup } from "@/components/ui/field";
|
||||
import { FormField } from "@/components/shared/form/FormField";
|
||||
import { PaginatedSearchSelect } from "@/components/shared/PaginatedSearchSelect";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import {
|
||||
Combobox,
|
||||
ComboboxContent,
|
||||
ComboboxEmpty,
|
||||
ComboboxInput,
|
||||
ComboboxItem,
|
||||
ComboboxList,
|
||||
} from "@/components/ui/combobox";
|
||||
import { Dialog, DialogContent, DialogHeader, DialogTitle } from "@/components/ui/dialog";
|
||||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip";
|
||||
|
|
@ -119,14 +110,9 @@ const UserSearchModal: React.FC<UserSearchModalProps> = ({
|
|||
}
|
||||
};
|
||||
|
||||
const debouncedSearch = useDebouncedCallback(
|
||||
(text: string, fieldName: "user_email" | "user_id") => fetchUsers(text, fieldName),
|
||||
{ wait: DEBOUNCE_WAIT_MS },
|
||||
);
|
||||
|
||||
const handleSearch = (value: string, fieldName: "user_email" | "user_id"): void => {
|
||||
setSelectedField(fieldName);
|
||||
debouncedSearch(value, fieldName);
|
||||
void fetchUsers(value, fieldName);
|
||||
};
|
||||
|
||||
const handleSelect = (option: UserOption | null): void => {
|
||||
|
|
@ -154,54 +140,30 @@ const UserSearchModal: React.FC<UserSearchModalProps> = ({
|
|||
if (event.key === "Enter") event.preventDefault();
|
||||
};
|
||||
|
||||
const optionsFor = (fieldName: "user_email" | "user_id", value: string | undefined): UserOption[] => {
|
||||
const visible = selectedField === fieldName ? userOptions : [];
|
||||
if (value == null || value === "" || visible.some((option) => option.value === value)) return visible;
|
||||
return [{ label: value, value, user: null }, ...visible];
|
||||
};
|
||||
|
||||
const renderUserSearch = (
|
||||
fieldName: "user_email" | "user_id",
|
||||
placeholder: string,
|
||||
controlProps: { id: string; value: string | undefined; onChange: (value: string | undefined) => void },
|
||||
testId?: string,
|
||||
) => {
|
||||
const items = optionsFor(fieldName, controlProps.value);
|
||||
const selected = items.find((option) => option.value === controlProps.value) ?? null;
|
||||
const items = selectedField === fieldName ? userOptions : [];
|
||||
return (
|
||||
<div data-testid={testId}>
|
||||
<Combobox
|
||||
items={items}
|
||||
value={selected}
|
||||
// @ts-expect-error TS2322 -- Combobox.Root narrows autoHighlight to boolean; the AriaCombobox it wraps
|
||||
// accepts "always", the only value that highlights a list this component filters server-side
|
||||
autoHighlight="always"
|
||||
filter={null}
|
||||
onValueChange={(option: UserOption | null) => {
|
||||
controlProps.onChange(option?.value);
|
||||
handleSelect(option);
|
||||
<div data-testid={testId} onKeyDown={swallowEnter}>
|
||||
<PaginatedSearchSelect
|
||||
options={items}
|
||||
value={controlProps.value}
|
||||
onValueChange={(value: string) => {
|
||||
controlProps.onChange(value === "" ? undefined : value);
|
||||
handleSelect(items.find((option) => option.value === value) ?? null);
|
||||
}}
|
||||
onInputValueChange={(text: string) => handleSearch(text, fieldName)}
|
||||
isItemEqualToValue={(a: UserOption, b: UserOption) => a.value === b.value}
|
||||
itemToStringLabel={(option: UserOption) => option.label}
|
||||
>
|
||||
<ComboboxInput
|
||||
id={controlProps.id}
|
||||
placeholder={placeholder}
|
||||
showClear={selected !== null}
|
||||
onKeyDown={swallowEnter}
|
||||
/>
|
||||
<ComboboxContent>
|
||||
<ComboboxEmpty>{loading ? "Loading..." : "No results"}</ComboboxEmpty>
|
||||
<ComboboxList>
|
||||
{(option: UserOption) => (
|
||||
<ComboboxItem key={option.value} value={option}>
|
||||
{option.label}
|
||||
</ComboboxItem>
|
||||
)}
|
||||
</ComboboxList>
|
||||
</ComboboxContent>
|
||||
</Combobox>
|
||||
onSearchChange={(query: string) => handleSearch(query, fieldName)}
|
||||
autoHighlight="always"
|
||||
isLoading={loading}
|
||||
placeholder={placeholder}
|
||||
emptyText="No results"
|
||||
loadingText="Loading..."
|
||||
inputId={controlProps.id}
|
||||
/>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -829,14 +829,14 @@ describe("CreateKey", () => {
|
|||
await act(async () => {
|
||||
answers.get("alice.smith@example.com")?.([{ user_id: "u-smith", user_email: "alice.smith@example.com" }]);
|
||||
});
|
||||
await screen.findByTitle("alice.smith@example.com (u-smith)");
|
||||
await screen.findByRole("option", { name: "alice.smith@example.com (u-smith)" });
|
||||
|
||||
await act(async () => {
|
||||
answers.get("ali")?.([{ user_id: "u-jones", user_email: "alice.jones@example.com" }]);
|
||||
});
|
||||
|
||||
expect(screen.queryByTitle("alice.jones@example.com (u-jones)")).not.toBeInTheDocument();
|
||||
expect(screen.getByTitle("alice.smith@example.com (u-smith)")).toBeInTheDocument();
|
||||
expect(screen.queryByRole("option", { name: "alice.jones@example.com (u-jones)" })).not.toBeInTheDocument();
|
||||
expect(screen.getByRole("option", { name: "alice.smith@example.com (u-smith)" })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("stops searching once the box is cleared and the abandoned search answers", async () => {
|
||||
|
|
@ -863,7 +863,7 @@ describe("CreateKey", () => {
|
|||
answers.get("ali")?.([{ user_id: "u-jones", user_email: "alice.jones@example.com" }]);
|
||||
});
|
||||
|
||||
expect(screen.queryByTitle("alice.jones@example.com (u-jones)")).not.toBeInTheDocument();
|
||||
expect(screen.queryByRole("option", { name: "alice.jones@example.com (u-jones)" })).not.toBeInTheDocument();
|
||||
expect(screen.getByText("No users found")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
|
|
@ -896,7 +896,7 @@ describe("CreateKey", () => {
|
|||
await act(async () => {
|
||||
answers.get("alice.smith@example.com")?.([{ user_id: "u-smith", user_email: "alice.smith@example.com" }]);
|
||||
});
|
||||
await screen.findByTitle("alice.smith@example.com (u-smith)");
|
||||
await screen.findByRole("option", { name: "alice.smith@example.com (u-smith)" });
|
||||
});
|
||||
|
||||
it("only warns about a failed search when it is the one the box is waiting on", async () => {
|
||||
|
|
@ -926,14 +926,14 @@ describe("CreateKey", () => {
|
|||
.get("alice.smith@example.com")
|
||||
?.resolve([{ user_id: "u-smith", user_email: "alice.smith@example.com" }]);
|
||||
});
|
||||
await screen.findByTitle("alice.smith@example.com (u-smith)");
|
||||
await screen.findByRole("option", { name: "alice.smith@example.com (u-smith)" });
|
||||
|
||||
await act(async () => {
|
||||
answers.get("ali")?.reject(new Error("search failed"));
|
||||
});
|
||||
|
||||
expect(toast.fromError).not.toHaveBeenCalled();
|
||||
expect(screen.getByTitle("alice.smith@example.com (u-smith)")).toBeInTheDocument();
|
||||
expect(screen.getByRole("option", { name: "alice.smith@example.com (u-smith)" })).toBeInTheDocument();
|
||||
|
||||
await user.type(search, "x");
|
||||
await waitFor(() => expect(answers.has("alice.smith@example.comx")).toBe(true), { timeout: 3000 });
|
||||
|
|
@ -946,6 +946,45 @@ describe("CreateKey", () => {
|
|||
});
|
||||
});
|
||||
|
||||
describe("user picker selection", () => {
|
||||
it("keeps the picked user in the box instead of searching for its own label", async () => {
|
||||
const directory = [
|
||||
{ user_id: "u-77", user_email: "alice@example.com" },
|
||||
{ user_id: "u-88", user_email: "bob@example.com" },
|
||||
];
|
||||
vi.mocked(userFilterUICall).mockImplementation(
|
||||
(_accessToken, params) =>
|
||||
Promise.resolve(
|
||||
directory.filter((entry) => entry.user_email.includes(params.get("user_email") ?? "")),
|
||||
) as never,
|
||||
);
|
||||
|
||||
const user = userEvent.setup({ advanceTimers: vi.advanceTimersByTime });
|
||||
vi.useFakeTimers({ shouldAdvanceTime: true });
|
||||
try {
|
||||
renderCreateKey({
|
||||
autoOpenCreate: true,
|
||||
prefillData: { owned_by: "another_user", key_alias: "contract-key" },
|
||||
});
|
||||
const search = await userSearchInput();
|
||||
|
||||
await user.type(search, "alice");
|
||||
await user.click(await screen.findByRole("option", { name: "alice@example.com (u-77)" }));
|
||||
await act(async () => {
|
||||
await vi.advanceTimersByTimeAsync(1000);
|
||||
});
|
||||
|
||||
expect(search).toHaveValue("alice@example.com (u-77)");
|
||||
expect(vi.mocked(userFilterUICall)).toHaveBeenCalledTimes(1);
|
||||
} finally {
|
||||
vi.useRealTimers();
|
||||
}
|
||||
|
||||
await submit();
|
||||
expect((await createdPayload()).user_id).toBe("u-77");
|
||||
});
|
||||
});
|
||||
|
||||
describe("created key display", () => {
|
||||
it("surfaces the generated key after a successful create", async () => {
|
||||
await openModal();
|
||||
|
|
|
|||
|
|
@ -13,25 +13,16 @@ import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/component
|
|||
import { Input } from "@/components/ui/input";
|
||||
import { Field, FieldLabel } from "@/components/ui/field";
|
||||
import { Badge } from "@/components/ui/badge";
|
||||
import {
|
||||
Combobox,
|
||||
ComboboxContent,
|
||||
ComboboxEmpty,
|
||||
ComboboxInput,
|
||||
ComboboxItem,
|
||||
ComboboxList,
|
||||
} from "@/components/ui/combobox";
|
||||
import { RadioGroup, RadioGroupItem } from "@/components/ui/radio-group";
|
||||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import { Switch } from "@/components/ui/switch";
|
||||
import { Textarea } from "@/components/ui/textarea";
|
||||
import { SimpleTooltip } from "@/components/ui/tooltip";
|
||||
import { MultiSelect, type MultiSelectOption } from "@/components/shared/MultiSelect";
|
||||
import { SearchSelect } from "@/components/shared/SearchSelect";
|
||||
import { PaginatedSearchSelect } from "@/components/shared/PaginatedSearchSelect";
|
||||
import { SearchSelect, type SearchSelectOption } from "@/components/shared/SearchSelect";
|
||||
import { TagsInput } from "@/app/(dashboard)/guardrails/_components/content_filter/TagsInput";
|
||||
import { ChevronDown, Info } from "lucide-react";
|
||||
import { useDebouncedCallback } from "@tanstack/react-pacer/debouncer";
|
||||
import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants";
|
||||
import React, { useEffect, useMemo, useRef, useState } from "react";
|
||||
import { type Control, useForm, useWatch, type UseFormSetValue } from "react-hook-form";
|
||||
import { rolesWithWriteAccess } from "../../utils/roles";
|
||||
|
|
@ -169,12 +160,6 @@ interface User {
|
|||
role?: string;
|
||||
}
|
||||
|
||||
interface UserOption {
|
||||
label: string;
|
||||
value: string;
|
||||
user: User;
|
||||
}
|
||||
|
||||
export const fetchTeamModels = async (
|
||||
userID: string,
|
||||
userRole: string,
|
||||
|
|
@ -270,7 +255,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
|
|||
const [selectedProjectId, setSelectedProjectId] = useState<string | null>(null);
|
||||
const [isCreateUserModalVisible, setIsCreateUserModalVisible] = useState(false);
|
||||
const [possibleUIRoles, setPossibleUIRoles] = useState<Record<string, Record<string, string>>>({});
|
||||
const [userOptions, setUserOptions] = useState<UserOption[]>([]);
|
||||
const [userOptions, setUserOptions] = useState<SearchSelectOption[]>([]);
|
||||
const [userSearchLoading, setUserSearchLoading] = useState<boolean>(false);
|
||||
const latestUserSearchRef = useRef(0);
|
||||
const [disabledCallbacks, setDisabledCallbacks] = useState<string[]>([]);
|
||||
|
|
@ -588,10 +573,9 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
|
|||
if (!isLatestSearch()) return;
|
||||
|
||||
const data: User[] = response;
|
||||
const options: UserOption[] = data.map((user) => ({
|
||||
const options: SearchSelectOption[] = data.map((user) => ({
|
||||
label: `${user.user_email} (${user.user_id})`,
|
||||
value: user.user_id,
|
||||
user,
|
||||
}));
|
||||
|
||||
setUserOptions(options);
|
||||
|
|
@ -603,8 +587,6 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
|
|||
}
|
||||
};
|
||||
|
||||
const handleUserSearch = useDebouncedCallback((text: string) => fetchUsers(text), { wait: DEBOUNCE_WAIT_MS });
|
||||
|
||||
const changeOrganization = (write: FieldWrite) => (orgId: string) => {
|
||||
write(orgId);
|
||||
setSelectedOrganizationId(orgId || null);
|
||||
|
|
@ -736,36 +718,20 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
|
|||
{(control) => (
|
||||
<div>
|
||||
<div className="mb-2 flex">
|
||||
<Combobox
|
||||
items={userOptions}
|
||||
value={userOptions.find((option) => option.value === control.value) ?? null}
|
||||
filter={null}
|
||||
onValueChange={(option: UserOption | null) => control.onChange(option?.value)}
|
||||
onInputValueChange={handleUserSearch}
|
||||
isItemEqualToValue={(a: UserOption, b: UserOption) => a.value === b.value}
|
||||
itemToStringLabel={(option: UserOption) => option.label}
|
||||
>
|
||||
<ComboboxInput
|
||||
id={control.id}
|
||||
className="w-full"
|
||||
placeholder="Type email to search for users"
|
||||
aria-required={control["aria-required"]}
|
||||
aria-invalid={control["aria-invalid"]}
|
||||
aria-describedby={control["aria-describedby"]}
|
||||
showClear={control.value != null && control.value !== ""}
|
||||
onBlur={control.onBlur}
|
||||
/>
|
||||
<ComboboxContent>
|
||||
<ComboboxEmpty>{userSearchLoading ? "Searching..." : "No users found"}</ComboboxEmpty>
|
||||
<ComboboxList>
|
||||
{(option: UserOption) => (
|
||||
<ComboboxItem key={option.value} value={option} title={option.label}>
|
||||
{option.label}
|
||||
</ComboboxItem>
|
||||
)}
|
||||
</ComboboxList>
|
||||
</ComboboxContent>
|
||||
</Combobox>
|
||||
<PaginatedSearchSelect
|
||||
options={userOptions}
|
||||
value={typeof control.value === "string" ? control.value : undefined}
|
||||
onValueChange={control.onChange}
|
||||
onSearchChange={fetchUsers}
|
||||
isLoading={userSearchLoading}
|
||||
placeholder="Type email to search for users"
|
||||
emptyText="No users found"
|
||||
loadingText="Searching..."
|
||||
inputId={control.id}
|
||||
aria-required={control["aria-required"] === "true" ? true : undefined}
|
||||
aria-invalid={control["aria-invalid"] === "true" ? true : undefined}
|
||||
aria-describedby={control["aria-describedby"]}
|
||||
/>
|
||||
<Button variant="outline" className="ml-2" onClick={() => setIsCreateUserModalVisible(true)}>
|
||||
Create User
|
||||
</Button>
|
||||
|
|
|
|||
|
|
@ -69,6 +69,38 @@ describe("PaginatedMultiSelect", () => {
|
|||
await waitFor(() => expect(onSearchChange).toHaveBeenLastCalledWith(""), { timeout: 2000 });
|
||||
});
|
||||
|
||||
it("puts the unfiltered page back when a typed query is abandoned by closing", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSearchChange = vi.fn();
|
||||
renderSelect({ onSearchChange });
|
||||
|
||||
const input = screen.getByRole("combobox");
|
||||
await user.click(input);
|
||||
await user.type(input, "gamma");
|
||||
await waitFor(() => expect(onSearchChange).toHaveBeenCalledWith("gamma"), { timeout: 2000 });
|
||||
|
||||
await user.keyboard("{Escape}");
|
||||
|
||||
await waitFor(() => expect(onSearchChange).toHaveBeenLastCalledWith(""), { timeout: 2000 });
|
||||
expect(input).toHaveValue("");
|
||||
});
|
||||
|
||||
it("puts the unfiltered page back when the popup is dismissed by clicking away", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSearchChange = vi.fn();
|
||||
renderSelect({ onSearchChange });
|
||||
|
||||
const input = screen.getByRole("combobox");
|
||||
await user.click(input);
|
||||
await user.type(input, "gamma");
|
||||
await waitFor(() => expect(onSearchChange).toHaveBeenCalledWith("gamma"), { timeout: 2000 });
|
||||
|
||||
await user.click(document.body);
|
||||
|
||||
await waitFor(() => expect(onSearchChange).toHaveBeenLastCalledWith(""), { timeout: 2000 });
|
||||
expect(input).toHaveValue("");
|
||||
});
|
||||
|
||||
it("selects multiple values and reports them cumulatively", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onValueChange = vi.fn();
|
||||
|
|
@ -192,6 +224,33 @@ describe("PaginatedMultiSelect", () => {
|
|||
expect(within(chips).queryByText("hash-alpha")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("clears every selection through the clear-all control when a label is provided", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onValueChange = vi.fn();
|
||||
renderSelect({ value: ["alias-alpha", "alias-beta"], onValueChange, clearAllLabel: "Clear all" });
|
||||
|
||||
await user.click(screen.getByLabelText("Clear all"));
|
||||
|
||||
expect(onValueChange).toHaveBeenCalledWith([]);
|
||||
});
|
||||
|
||||
it("shows no clear-all control without a label or without selections", () => {
|
||||
const { unmount } = render(
|
||||
<PaginatedMultiSelect
|
||||
options={OPTIONS}
|
||||
value={["alias-alpha"]}
|
||||
onValueChange={vi.fn()}
|
||||
onSearchChange={vi.fn()}
|
||||
onLoadMore={vi.fn()}
|
||||
/>,
|
||||
);
|
||||
expect(document.querySelector('[data-slot="combobox-clear"]')).not.toBeInTheDocument();
|
||||
unmount();
|
||||
|
||||
renderSelect({ value: [], clearAllLabel: "Clear all" });
|
||||
expect(document.querySelector('[data-slot="combobox-clear"]')).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("anchors the dropdown to the chips container so it tracks the growing chip box", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderSelect({});
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ import {
|
|||
ComboboxChip,
|
||||
ComboboxChips,
|
||||
ComboboxChipsInput,
|
||||
ComboboxClear,
|
||||
ComboboxContent,
|
||||
ComboboxEmpty,
|
||||
ComboboxItem,
|
||||
|
|
@ -32,6 +33,7 @@ interface PaginatedMultiSelectProps {
|
|||
emptyText?: string;
|
||||
errorText?: string;
|
||||
loadingText?: string;
|
||||
clearAllLabel?: string;
|
||||
disabled?: boolean;
|
||||
className?: string;
|
||||
inputId?: string;
|
||||
|
|
@ -52,6 +54,7 @@ export function PaginatedMultiSelect({
|
|||
emptyText = "No results",
|
||||
errorText,
|
||||
loadingText = "Loading…",
|
||||
clearAllLabel,
|
||||
disabled = false,
|
||||
className,
|
||||
inputId,
|
||||
|
|
@ -119,6 +122,7 @@ export function PaginatedMultiSelect({
|
|||
className="h-5 min-w-24 flex-1 border-0 bg-transparent py-0 text-sm"
|
||||
aria-label={placeholder}
|
||||
/>
|
||||
{clearAllLabel != null && value.length > 0 && <ComboboxClear aria-label={clearAllLabel} disabled={disabled} />}
|
||||
</ComboboxChips>
|
||||
<ComboboxContent anchor={anchor}>
|
||||
<ComboboxEmpty className={errorText == null ? undefined : "text-destructive"}>
|
||||
|
|
|
|||
|
|
@ -276,6 +276,57 @@ describe("PaginatedSearchSelect", () => {
|
|||
await waitFor(() => expect(onSearchChange).toHaveBeenLastCalledWith("gamma"));
|
||||
});
|
||||
|
||||
it("commits the first server-filtered match on Enter when autoHighlight is always", async () => {
|
||||
const user = userEvent.setup();
|
||||
|
||||
function ServerBacked() {
|
||||
const [search, setSearch] = useState("");
|
||||
const [value, setValue] = useState("");
|
||||
return (
|
||||
<PaginatedSearchSelect
|
||||
options={OPTIONS.filter((option) => option.label.includes(search))}
|
||||
value={value}
|
||||
onValueChange={setValue}
|
||||
onSearchChange={setSearch}
|
||||
onLoadMore={vi.fn()}
|
||||
autoHighlight="always"
|
||||
/>
|
||||
);
|
||||
}
|
||||
render(<ServerBacked />);
|
||||
|
||||
const input = screen.getByRole("combobox");
|
||||
await user.click(input);
|
||||
await user.type(input, "gamma");
|
||||
await waitFor(() => expect(screen.queryByText("alias-alpha")).not.toBeInTheDocument());
|
||||
await user.keyboard("{Enter}");
|
||||
|
||||
await waitFor(() => expect(input).toHaveValue("gamma-key"));
|
||||
});
|
||||
|
||||
it("highlights the picked label on focus so typing starts over", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderSelect({ value: "alias-alpha" });
|
||||
|
||||
await user.tab();
|
||||
|
||||
const input = screen.getByRole("combobox") as HTMLInputElement;
|
||||
expect(input.selectionStart).toBe(0);
|
||||
expect(input.selectionEnd).toBe("alias-alpha".length);
|
||||
});
|
||||
|
||||
it("takes a paste over the highlighted label wholesale even when it shares a prefix", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSearchChange = vi.fn();
|
||||
renderSelect({ onSearchChange, value: "alias-alpha" });
|
||||
|
||||
await user.tab();
|
||||
await user.paste("alias-alphabet");
|
||||
|
||||
expect(screen.getByRole("combobox")).toHaveValue("alias-alphabet");
|
||||
await waitFor(() => expect(onSearchChange).toHaveBeenLastCalledWith("alias-alphabet"));
|
||||
});
|
||||
|
||||
it("starts a fresh query when typing lands inside the selected label", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSearchChange = vi.fn();
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
"use client";
|
||||
|
||||
import { Loader2 } from "lucide-react";
|
||||
import { useMemo, useState } from "react";
|
||||
import { useMemo, useRef, useState, type SyntheticEvent } from "react";
|
||||
|
||||
import {
|
||||
Combobox,
|
||||
|
|
@ -28,9 +28,11 @@ interface PaginatedSearchSelectProps {
|
|||
emptyText?: string;
|
||||
errorText?: string;
|
||||
loadingText?: string;
|
||||
autoHighlight?: boolean | "always";
|
||||
disabled?: boolean;
|
||||
className?: string;
|
||||
inputId?: string;
|
||||
"aria-required"?: true | undefined;
|
||||
"aria-invalid"?: true | undefined;
|
||||
"aria-describedby"?: string;
|
||||
}
|
||||
|
|
@ -61,13 +63,22 @@ export function PaginatedSearchSelect({
|
|||
emptyText = "No results",
|
||||
errorText,
|
||||
loadingText = "Loading…",
|
||||
autoHighlight = false,
|
||||
disabled = false,
|
||||
className,
|
||||
inputId,
|
||||
"aria-required": ariaRequired,
|
||||
"aria-invalid": ariaInvalid,
|
||||
"aria-describedby": ariaDescribedBy,
|
||||
}: PaginatedSearchSelectProps) {
|
||||
const [pickedOption, setPickedOption] = useState<SearchSelectOption | null>(null);
|
||||
const wholeSelectionRef = useRef(false);
|
||||
|
||||
const snapshotWholeSelection = (event: SyntheticEvent<HTMLInputElement>) => {
|
||||
const input = event.currentTarget;
|
||||
wholeSelectionRef.current =
|
||||
input.value.length > 0 && input.selectionStart === 0 && input.selectionEnd === input.value.length;
|
||||
};
|
||||
|
||||
const selected = useMemo<SearchSelectOption | null>(() => {
|
||||
if (value === undefined || value === "") return null;
|
||||
|
|
@ -86,6 +97,15 @@ export function PaginatedSearchSelect({
|
|||
const pagination = { onSearchChange, onLoadMore, hasNextPage, isFetchingNextPage };
|
||||
const { typedQuery, handleInputValueChange, handleOpenChange, handleScroll } = usePaginatedCombobox(pagination);
|
||||
|
||||
const handleTypedInput = (next: string, reason: string) => {
|
||||
const replacedWholeInput = wholeSelectionRef.current;
|
||||
wholeSelectionRef.current = false;
|
||||
handleInputValueChange(
|
||||
typedQuery === null && !replacedWholeInput ? typedInsertion(selected?.label ?? "", next) : next,
|
||||
reason,
|
||||
);
|
||||
};
|
||||
|
||||
return (
|
||||
<Combobox
|
||||
items={items}
|
||||
|
|
@ -95,22 +115,24 @@ export function PaginatedSearchSelect({
|
|||
setPickedOption(item);
|
||||
onValueChange(item?.value ?? "");
|
||||
}}
|
||||
onInputValueChange={(next, eventDetails) =>
|
||||
handleInputValueChange(
|
||||
typedQuery === null ? typedInsertion(selected?.label ?? "", next) : next,
|
||||
eventDetails.reason,
|
||||
)
|
||||
}
|
||||
onInputValueChange={(next, eventDetails) => handleTypedInput(next, eventDetails.reason)}
|
||||
onOpenChange={(nextOpen, eventDetails) => handleOpenChange(nextOpen, eventDetails.reason)}
|
||||
isItemEqualToValue={(a: SearchSelectOption, b: SearchSelectOption) => a.value === b.value}
|
||||
itemToStringLabel={(item: SearchSelectOption) => item.label}
|
||||
// @ts-expect-error TS2322 -- Combobox.Root narrows autoHighlight to boolean; the AriaCombobox it wraps
|
||||
// accepts "always", the only value that highlights a list filtered server-side
|
||||
autoHighlight={autoHighlight}
|
||||
filter={null}
|
||||
disabled={disabled}
|
||||
>
|
||||
<ComboboxInput
|
||||
id={inputId}
|
||||
aria-required={ariaRequired}
|
||||
aria-invalid={ariaInvalid}
|
||||
aria-describedby={ariaDescribedBy}
|
||||
onFocus={(event) => event.currentTarget.select()}
|
||||
onKeyDown={snapshotWholeSelection}
|
||||
onPaste={snapshotWholeSelection}
|
||||
placeholder={placeholder}
|
||||
showClear={value !== undefined && value !== ""}
|
||||
className={`w-full ${className ?? ""}`}
|
||||
|
|
|
|||
|
|
@ -1,8 +1,10 @@
|
|||
import { fireEvent, screen, waitFor } from "@testing-library/react";
|
||||
import { fireEvent, screen, waitFor, within } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { useState } from "react";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
|
||||
import { renderWithProviders, testQueryClient } from "../../../tests/test-utils";
|
||||
import { ERROR_CODE_OPTIONS } from "./constants";
|
||||
import { LOG_FILTER_IDS } from "./log_filter_logic";
|
||||
import { RequestLogsFilters } from "./RequestLogsFilters";
|
||||
|
||||
|
|
@ -45,6 +47,20 @@ function renderFilters(filters: Record<string, string> = {}) {
|
|||
return { set };
|
||||
}
|
||||
|
||||
function StatefulFilters() {
|
||||
const [filters, setFilters] = useState<Record<string, string | undefined>>({});
|
||||
return (
|
||||
<RequestLogsFilters
|
||||
get={(id: string) => filters[id]}
|
||||
set={(id: string, value: unknown) =>
|
||||
setFilters((previous) => ({ ...previous, [id]: typeof value === "string" ? value : undefined }))
|
||||
}
|
||||
teams={[]}
|
||||
logsWindow={LOGS_WINDOW}
|
||||
/>
|
||||
);
|
||||
}
|
||||
|
||||
describe("RequestLogsFilters", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
|
|
@ -284,6 +300,42 @@ describe("RequestLogsFilters", () => {
|
|||
expect(set).toHaveBeenCalledWith(LOG_FILTER_IDS.CACHE_STATUS, expected);
|
||||
});
|
||||
|
||||
it("stores the raw status code when a labeled error code is picked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const { set } = renderFilters();
|
||||
|
||||
await user.click(await screen.findByPlaceholderText("Select or type an error code"));
|
||||
await user.click(await screen.findByRole("option", { name: "429 - Rate Limited" }));
|
||||
|
||||
expect(set).toHaveBeenCalledWith(LOG_FILTER_IDS.ERROR_CODE, "429");
|
||||
});
|
||||
|
||||
it("offers every error code again after one was picked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<StatefulFilters />);
|
||||
|
||||
const input = await screen.findByPlaceholderText("Select or type an error code");
|
||||
await user.click(input);
|
||||
await user.click(await screen.findByRole("option", { name: "429 - Rate Limited" }));
|
||||
await user.click(input);
|
||||
|
||||
const list = await screen.findByTestId("error-code-filter-list");
|
||||
expect(within(list).getAllByRole("option")).toHaveLength(ERROR_CODE_OPTIONS.length);
|
||||
expect(within(list).queryByText(/^Use custom code:/)).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("filters by an error code the list does not offer", async () => {
|
||||
const user = userEvent.setup();
|
||||
const { set } = renderFilters();
|
||||
|
||||
const input = await screen.findByPlaceholderText("Select or type an error code");
|
||||
await user.click(input);
|
||||
await user.type(input, "418");
|
||||
await user.click(await screen.findByRole("option", { name: "Use custom code: 418" }));
|
||||
|
||||
expect(set).toHaveBeenCalledWith(LOG_FILTER_IDS.ERROR_CODE, "418");
|
||||
});
|
||||
|
||||
it("selecting All Requests clears the cache filter", async () => {
|
||||
const user = userEvent.setup();
|
||||
const { set } = renderFilters({ [LOG_FILTER_IDS.CACHE_STATUS]: "hit" });
|
||||
|
|
|
|||
|
|
@ -39,6 +39,8 @@ const CACHE_FILTER_ITEMS = [
|
|||
] as const;
|
||||
const PAGE_SIZE = 50;
|
||||
|
||||
const SEARCH_INPUT_REASONS: ReadonlySet<string> = new Set(["input-change", "input-clear", "clear-press"]);
|
||||
|
||||
const asString = (value: unknown): string => (typeof value === "string" ? value : "");
|
||||
const emptyToUndefined = (value: string): string | undefined => (value === "" ? undefined : value);
|
||||
|
||||
|
|
@ -254,7 +256,10 @@ function ErrorCodeFilterField({ value, onChange }: { value: string; onChange: (v
|
|||
const trimmed = query.trim();
|
||||
const lowered = trimmed.toLowerCase();
|
||||
const matches = ERROR_CODE_OPTIONS.filter((option) => option.label.toLowerCase().includes(lowered));
|
||||
if (trimmed === "" || ERROR_CODE_OPTIONS.some((option) => option.value === trimmed)) return matches;
|
||||
const isKnownCode = ERROR_CODE_OPTIONS.some(
|
||||
(option) => option.value === trimmed || option.label.toLowerCase() === lowered,
|
||||
);
|
||||
if (trimmed === "" || isKnownCode) return matches;
|
||||
return [...matches, { label: `Use custom code: ${trimmed}`, value: trimmed }];
|
||||
}, [query]);
|
||||
|
||||
|
|
@ -275,12 +280,20 @@ function ErrorCodeFilterField({ value, onChange }: { value: string; onChange: (v
|
|||
items={items}
|
||||
value={selected}
|
||||
onValueChange={(item: SearchSelectOption | null) => onChange(emptyToUndefined(item?.value ?? ""))}
|
||||
onInputValueChange={setQuery}
|
||||
onInputValueChange={(next, eventDetails) => setQuery(SEARCH_INPUT_REASONS.has(eventDetails.reason) ? next : "")}
|
||||
onOpenChange={(nextOpen) => {
|
||||
if (!nextOpen) setQuery("");
|
||||
}}
|
||||
isItemEqualToValue={(a: SearchSelectOption, b: SearchSelectOption) => a.value === b.value}
|
||||
itemToStringLabel={(item: SearchSelectOption) => item.label}
|
||||
filter={null}
|
||||
>
|
||||
<ComboboxInput placeholder="Select or type an error code" showClear={value !== ""} className="w-full" />
|
||||
<ComboboxInput
|
||||
onFocus={(event) => event.currentTarget.select()}
|
||||
placeholder="Select or type an error code"
|
||||
showClear={value !== ""}
|
||||
className="w-full"
|
||||
/>
|
||||
<ComboboxContent>
|
||||
<ComboboxEmpty>No error codes found</ComboboxEmpty>
|
||||
<ComboboxList data-testid="error-code-filter-list">
|
||||
|
|
|
|||
15
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
15
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -24353,7 +24353,7 @@ export interface components {
|
|||
/** ChatCompletionToolMessage */
|
||||
ChatCompletionToolMessage: {
|
||||
/** Content */
|
||||
content: string | (components["schemas"]["ChatCompletionTextObject"] | components["schemas"]["ChatCompletionImageObject"])[];
|
||||
content: string | (components["schemas"]["ChatCompletionTextObject"] | components["schemas"]["ChatCompletionImageObject"] | components["schemas"]["ChatCompletionToolReferenceObject"])[];
|
||||
/**
|
||||
* Role
|
||||
* @constant
|
||||
|
|
@ -24384,6 +24384,19 @@ export interface components {
|
|||
/** Strict */
|
||||
strict?: boolean;
|
||||
};
|
||||
/**
|
||||
* ChatCompletionToolReferenceObject
|
||||
* @description Anthropic tool-search result block, carried through untouched so it survives a round trip.
|
||||
*/
|
||||
ChatCompletionToolReferenceObject: {
|
||||
/** Tool Name */
|
||||
tool_name: string;
|
||||
/**
|
||||
* Type
|
||||
* @constant
|
||||
*/
|
||||
type: "tool_reference";
|
||||
};
|
||||
/** ChatCompletionUserMessage */
|
||||
ChatCompletionUserMessage: {
|
||||
cache_control?: components["schemas"]["ChatCompletionCachedContent"];
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
Loading…
Add table
Reference in a new issue