mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge 5e6ce0f6db into 461a58c40a
This commit is contained in:
commit
765b62351c
8 changed files with 450 additions and 34 deletions
|
|
@ -801,7 +801,15 @@ def _get_hidden_str_for_cost_calc(hidden_params: object, key: str) -> str | None
|
|||
|
||||
|
||||
_NON_TOKEN_RATE_FIELDS: Final = frozenset(
|
||||
{"cost_per_second", "input_cost_per_second", "output_cost_per_second", "input_cost_per_query", "tiered_pricing"}
|
||||
{
|
||||
"cost_per_second",
|
||||
"input_cost_per_second",
|
||||
"output_cost_per_second",
|
||||
"input_cost_per_query",
|
||||
"input_cost_per_character",
|
||||
"output_cost_per_character",
|
||||
"tiered_pricing",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -90,11 +90,51 @@ WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY: Final = "_websearch_interception_emit_native_b
|
|||
# ``web_search_tool_result`` blocks to inject into the final response.
|
||||
WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY: Final = "websearch_native_blocks"
|
||||
|
||||
# Key used to flag, on per-request kwargs, that the originating client sent
|
||||
# domain filters (``allowed_domains`` / ``blocked_domains``) on an
|
||||
# Anthropic-native ``web_search_*`` tool. The standard LiteLLM tool drops
|
||||
# them (the model must not see client policy), so they are stashed here and
|
||||
# applied to the downstream ``litellm.asearch()`` call as
|
||||
# ``search_domain_filter``.
|
||||
WEBSEARCH_DOMAIN_FILTER_KEY: Final = "_websearch_interception_domain_filter"
|
||||
|
||||
_RESPONSE_CONTENT_FIELD: Final = "content"
|
||||
|
||||
_ResponseT: Final = TypeVar("_ResponseT")
|
||||
|
||||
|
||||
def _extract_web_search_domain_filters(
|
||||
tools: Sequence[dict[str, object]],
|
||||
) -> dict[str, list[str]] | None:
|
||||
"""Collect ``allowed_domains`` / ``blocked_domains`` from web search tools.
|
||||
|
||||
Anthropic-native ``web_search_*`` tools carry optional domain limits. The
|
||||
standard LiteLLM tool deliberately drops them (client policy, not model
|
||||
input), so they are collected here, stashed on the request kwargs, and
|
||||
applied to the downstream ``litellm.asearch()`` call as
|
||||
``search_domain_filter``.
|
||||
|
||||
Returns None when no web search tool carries a domain limit.
|
||||
"""
|
||||
allowed: list[str] = []
|
||||
blocked: list[str] = []
|
||||
for tool in tools:
|
||||
if not is_web_search_tool(tool):
|
||||
continue
|
||||
for key, bucket in (("allowed_domains", allowed), ("blocked_domains", blocked)):
|
||||
value = tool.get(key)
|
||||
if isinstance(value, list):
|
||||
bucket.extend(item for item in value if isinstance(item, str) and item)
|
||||
if not allowed and not blocked:
|
||||
return None
|
||||
domain_filters: dict[str, list[str]] = {}
|
||||
if allowed:
|
||||
domain_filters["allowed_domains"] = allowed
|
||||
if blocked:
|
||||
domain_filters["blocked_domains"] = blocked
|
||||
return domain_filters
|
||||
|
||||
|
||||
class _PlanMetadataView(TypedDict):
|
||||
websearch_native_blocks: Sequence[Mapping[str, object]] | None
|
||||
|
||||
|
|
@ -142,7 +182,7 @@ class _AcreateNamedParams(TypedDict, total=False):
|
|||
|
||||
class _AsearchNamedParams(TypedDict, total=False):
|
||||
max_results: ReadOnly[int | None]
|
||||
search_domain_filter: ReadOnly[Never]
|
||||
search_domain_filter: ReadOnly[list[str] | None]
|
||||
max_tokens_per_page: ReadOnly[int | None]
|
||||
country: ReadOnly[str | None]
|
||||
api_key: ReadOnly[str | None]
|
||||
|
|
@ -377,6 +417,12 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
None,
|
||||
)
|
||||
|
||||
# Apply any domain limits the client set on the native web search
|
||||
# tool before the search executes.
|
||||
domain_filters: Final = _extract_web_search_domain_filters(tools)
|
||||
if domain_filters is not None and isinstance(kwargs, dict):
|
||||
kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] = domain_filters
|
||||
|
||||
outcome: Final = await self._short_circuit_search_outcome(query, kwargs=kwargs)
|
||||
search_result_text: Final = WebSearchTransformation.search_outcome_text(outcome)
|
||||
|
||||
|
|
@ -468,6 +514,12 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
if any(is_anthropic_native_web_search_tool(t) for t in tools):
|
||||
kwargs[WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY] = True
|
||||
|
||||
# Same for domain limits: stash them before the native tool is
|
||||
# replaced, so the downstream search can apply them.
|
||||
domain_filters: Final = _extract_web_search_domain_filters(tools)
|
||||
if domain_filters is not None:
|
||||
kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] = domain_filters
|
||||
|
||||
# Convert native/custom web_search tools to LiteLLM standard
|
||||
converted_tools: Final = []
|
||||
for tool in tools:
|
||||
|
|
@ -643,6 +695,12 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
if any(is_anthropic_native_web_search_tool(t) for t in tools):
|
||||
kwargs[WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY] = True
|
||||
|
||||
# Same for domain limits: stash them before the native tool is
|
||||
# replaced, so the downstream search can apply them.
|
||||
domain_filters: Final = _extract_web_search_domain_filters(tools)
|
||||
if domain_filters is not None:
|
||||
kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] = domain_filters
|
||||
|
||||
# Convert native web search tools to LiteLLM standard
|
||||
converted_tools: Final[list[dict[str, object]]] = []
|
||||
for tool in tools:
|
||||
|
|
@ -1595,17 +1653,39 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
rich_objective = rich.get("objective")
|
||||
if rich_objective and "objective" not in configured_search_kwargs:
|
||||
configured_search_kwargs["objective"] = rich_objective
|
||||
# Domain limits stashed by the pre-call hooks from the client's
|
||||
# native web_search tool. ``allowed_domains`` pass through as an
|
||||
# allowlist and ``blocked_domains`` as '-'-prefixed exclusions —
|
||||
# the convention litellm.asearch()'s providers (e.g. Perplexity)
|
||||
# use for search_domain_filter.
|
||||
search_domain_filter: list[str] | None = None
|
||||
request_domain_filters: Final = kwargs.get(WEBSEARCH_DOMAIN_FILTER_KEY) if kwargs is not None else None
|
||||
if isinstance(request_domain_filters, dict):
|
||||
allowed: Final[list[str]] = [
|
||||
item for item in request_domain_filters.get("allowed_domains", []) if isinstance(item, str) and item
|
||||
]
|
||||
blocked: Final[list[str]] = [
|
||||
item for item in request_domain_filters.get("blocked_domains", []) if isinstance(item, str) and item
|
||||
]
|
||||
if allowed or blocked:
|
||||
search_domain_filter = allowed + [f"-{item}" for item in blocked]
|
||||
verbose_logger.debug("WebSearchInterception: Applying domain filter %s", search_domain_filter)
|
||||
search_kwargs: Final = MappingProxyType(
|
||||
{**configured_search_kwargs, **parent_correlation.as_search_kwargs()}
|
||||
)
|
||||
result: Final = (
|
||||
await litellm.asearch(
|
||||
query=query_arg, search_provider=search_provider, **_NO_ASEARCH_NAMED, **search_kwargs
|
||||
query=query_arg,
|
||||
search_provider=search_provider,
|
||||
search_domain_filter=search_domain_filter,
|
||||
**_NO_ASEARCH_NAMED,
|
||||
**search_kwargs,
|
||||
)
|
||||
if search_metadata is None
|
||||
else await litellm.asearch(
|
||||
query=query_arg,
|
||||
search_provider=search_provider,
|
||||
search_domain_filter=search_domain_filter,
|
||||
litellm_metadata=search_metadata,
|
||||
**_NO_ASEARCH_NAMED,
|
||||
**search_kwargs,
|
||||
|
|
|
|||
|
|
@ -160,13 +160,18 @@ class DeepSeekChatConfig(OpenAIGPTConfig):
|
|||
|
||||
def _is_vision_forwardable_content(self, message: AllMessageValues, content: Sequence[object]) -> bool:
|
||||
"""
|
||||
True only for a user message whose content list holds well-formed
|
||||
text and image_url blocks with at least one image; a block missing
|
||||
its payload falls back to the string collapse instead of crashing
|
||||
or reaching the wire malformed. The model capability gate lives in
|
||||
the caller.
|
||||
True only for a user or tool message whose content list holds
|
||||
well-formed text and image_url blocks with at least one image; a block
|
||||
missing its payload falls back to the string collapse instead of
|
||||
crashing or reaching the wire malformed. The model capability gate
|
||||
lives in the caller.
|
||||
|
||||
``role="tool"`` is forwardable: the DeepSeek Chat Completions API
|
||||
accepts and reads image_url blocks in tool results (verified directly
|
||||
against the API; bug #44211) — agent loops that screenshot inside a
|
||||
tool rely on it.
|
||||
"""
|
||||
if message.get("role") != "user":
|
||||
if message.get("role") not in ("user", "tool"):
|
||||
return False
|
||||
if not all(self._is_forwardable_block(block) for block in content):
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -3611,11 +3611,22 @@ class Router:
|
|||
initial_kwargs["original_function"] = router_self._completion
|
||||
initial_kwargs["messages"] = messages
|
||||
router_self._update_kwargs_before_fallbacks(model=model_group, kwargs=initial_kwargs)
|
||||
fallback_response = router_self.function_with_fallbacks(
|
||||
**initial_kwargs,
|
||||
# Pass the MidStreamFallbackError through the common fallback utils, like the
|
||||
# async twin does. Calling function_with_fallbacks() here instead re-runs the
|
||||
# original (failing) group first, and because each nested Router.completion()
|
||||
# wraps its stream in this same iterator, every retry fails again only when
|
||||
# the caller iterates it — recursing until the stack runs out.
|
||||
fallback_response = run_async_function(
|
||||
router_self.async_function_with_fallbacks_common_utils,
|
||||
e,
|
||||
disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs),
|
||||
fallbacks=fallbacks,
|
||||
context_window_fallbacks=context_window_fallbacks,
|
||||
content_policy_fallbacks=content_policy_fallbacks,
|
||||
model_group=model_group,
|
||||
args=(),
|
||||
kwargs=initial_kwargs,
|
||||
include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True,
|
||||
)
|
||||
|
||||
if hasattr(fallback_response, "__iter__"):
|
||||
|
|
|
|||
|
|
@ -0,0 +1,132 @@
|
|||
"""
|
||||
Tests for domain-limit passthrough in web search interception.
|
||||
|
||||
Covers bug #44188: ``allowed_domains`` / ``blocked_domains`` set on an
|
||||
Anthropic-native ``web_search_*`` tool must survive the conversion to the
|
||||
standard LiteLLM web search tool and be applied to the downstream
|
||||
``litellm.asearch()`` call as ``search_domain_filter``.
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.websearch_interception.handler import (
|
||||
WEBSEARCH_DOMAIN_FILTER_KEY,
|
||||
WebSearchInterceptionLogger,
|
||||
_extract_web_search_domain_filters,
|
||||
)
|
||||
from litellm.llms.base_llm.search.transformation import SearchResponse, SearchResult
|
||||
|
||||
|
||||
def _make_search_response() -> SearchResponse:
|
||||
return SearchResponse(
|
||||
results=[
|
||||
SearchResult(
|
||||
title="LiteLLM Docs",
|
||||
url="https://docs.litellm.ai/",
|
||||
snippet="Unified interface for LLMs.",
|
||||
date=None,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
class TestExtractWebSearchDomainFilters:
|
||||
def test_collects_allowed_and_blocked(self):
|
||||
tools = [
|
||||
{
|
||||
"type": "web_search_20250305",
|
||||
"name": "web_search",
|
||||
"allowed_domains": ["docs.litellm.ai"],
|
||||
"blocked_domains": ["twitter.com", "x.com"],
|
||||
}
|
||||
]
|
||||
assert _extract_web_search_domain_filters(tools) == {
|
||||
"allowed_domains": ["docs.litellm.ai"],
|
||||
"blocked_domains": ["twitter.com", "x.com"],
|
||||
}
|
||||
|
||||
def test_returns_none_without_domain_limits(self):
|
||||
tools = [{"type": "web_search_20250305", "name": "web_search", "max_uses": 5}]
|
||||
assert _extract_web_search_domain_filters(tools) is None
|
||||
|
||||
def test_ignores_domains_on_non_web_search_tools(self):
|
||||
tools = [
|
||||
{"name": "bash", "allowed_domains": ["example.com"]},
|
||||
]
|
||||
assert _extract_web_search_domain_filters(tools) is None
|
||||
|
||||
def test_ignores_non_string_entries(self):
|
||||
tools = [
|
||||
{
|
||||
"type": "web_search_20250305",
|
||||
"name": "web_search",
|
||||
"allowed_domains": ["docs.litellm.ai", 42, None],
|
||||
}
|
||||
]
|
||||
assert _extract_web_search_domain_filters(tools) == {"allowed_domains": ["docs.litellm.ai"]}
|
||||
|
||||
|
||||
class TestDeploymentHookStashesDomainFilters:
|
||||
@pytest.mark.asyncio
|
||||
async def test_stashes_filters_for_native_tool(self):
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
kwargs = {
|
||||
"tools": [
|
||||
{
|
||||
"type": "web_search_20250305",
|
||||
"name": "web_search",
|
||||
"allowed_domains": ["docs.litellm.ai"],
|
||||
"blocked_domains": ["twitter.com"],
|
||||
}
|
||||
],
|
||||
"litellm_params": {"custom_llm_provider": "bedrock"},
|
||||
}
|
||||
out = await logger.async_pre_call_deployment_hook(kwargs, None)
|
||||
assert out is not None
|
||||
assert kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] == {
|
||||
"allowed_domains": ["docs.litellm.ai"],
|
||||
"blocked_domains": ["twitter.com"],
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_stash_without_domain_limits(self):
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
kwargs = {
|
||||
"tools": [{"type": "web_search_20250305", "name": "web_search"}],
|
||||
"litellm_params": {"custom_llm_provider": "bedrock"},
|
||||
}
|
||||
await logger.async_pre_call_deployment_hook(kwargs, None)
|
||||
assert WEBSEARCH_DOMAIN_FILTER_KEY not in kwargs
|
||||
|
||||
|
||||
class TestExecuteSearchAppliesDomainFilter:
|
||||
@pytest.mark.asyncio
|
||||
async def test_forwards_search_domain_filter(self):
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
kwargs = {
|
||||
WEBSEARCH_DOMAIN_FILTER_KEY: {
|
||||
"allowed_domains": ["docs.litellm.ai"],
|
||||
"blocked_domains": ["twitter.com"],
|
||||
}
|
||||
}
|
||||
asearch = AsyncMock(return_value=_make_search_response())
|
||||
with patch("litellm.asearch", asearch):
|
||||
await logger._execute_search("what is litellm", kwargs=kwargs)
|
||||
|
||||
assert asearch.await_count == 1
|
||||
assert asearch.await_args.kwargs.get("search_domain_filter") == [
|
||||
"docs.litellm.ai",
|
||||
"-twitter.com",
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_filter_when_kwargs_empty(self):
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
asearch = AsyncMock(return_value=_make_search_response())
|
||||
with patch("litellm.asearch", asearch):
|
||||
await logger._execute_search("what is litellm", kwargs={})
|
||||
|
||||
assert asearch.await_count == 1
|
||||
assert asearch.await_args.kwargs.get("search_domain_filter") is None
|
||||
|
|
@ -0,0 +1,67 @@
|
|||
"""
|
||||
Tests for DeepSeek vision-content forwarding in role=tool messages.
|
||||
|
||||
Bug #44211: ``_is_vision_forwardable_content`` rejected every non-user role,
|
||||
so a ``role=tool`` message carrying image_url blocks was silently collapsed
|
||||
to text before the request left LiteLLM — while the DeepSeek API itself
|
||||
accepts and reads tool-result images.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.deepseek.chat.transformation import DeepSeekChatConfig
|
||||
|
||||
|
||||
TOOL_MESSAGE_WITH_IMAGE = {
|
||||
"role": "tool",
|
||||
"tool_call_id": "abc",
|
||||
"content": [
|
||||
{"type": "text", "text": "screenshot:"},
|
||||
{"type": "image_url", "image_url": {"url": "data:image/png;base64,aGVsbG8="}},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
class TestVisionForwardableContent:
|
||||
def test_tool_message_with_image_is_forwardable(self):
|
||||
config = DeepSeekChatConfig()
|
||||
assert config._is_vision_forwardable_content(
|
||||
message=TOOL_MESSAGE_WITH_IMAGE,
|
||||
content=TOOL_MESSAGE_WITH_IMAGE["content"],
|
||||
)
|
||||
|
||||
def test_user_message_with_image_stays_forwardable(self):
|
||||
config = DeepSeekChatConfig()
|
||||
message = {"role": "user", "content": TOOL_MESSAGE_WITH_IMAGE["content"]}
|
||||
assert config._is_vision_forwardable_content(message=message, content=message["content"])
|
||||
|
||||
def test_assistant_message_stays_collapsed(self):
|
||||
config = DeepSeekChatConfig()
|
||||
message = {"role": "assistant", "content": TOOL_MESSAGE_WITH_IMAGE["content"]}
|
||||
assert not config._is_vision_forwardable_content(message=message, content=message["content"])
|
||||
|
||||
def test_image_missing_payload_falls_back(self):
|
||||
config = DeepSeekChatConfig()
|
||||
message = {
|
||||
"role": "tool",
|
||||
"content": [
|
||||
{"type": "text", "text": "screenshot:"},
|
||||
{"type": "image_url", "image_url": {"url": ""}},
|
||||
],
|
||||
}
|
||||
assert not config._is_vision_forwardable_content(message=message, content=message["content"])
|
||||
|
||||
|
||||
class TestForwardOrCollapseContent:
|
||||
def test_tool_message_image_content_is_not_collapsed(self):
|
||||
config = DeepSeekChatConfig()
|
||||
out = config._forward_or_collapse_content(message=TOOL_MESSAGE_WITH_IMAGE, forward_images=True)
|
||||
assert isinstance(out.get("content"), list)
|
||||
blocks = out["content"]
|
||||
assert any(isinstance(block, dict) and block.get("type") == "image_url" for block in blocks)
|
||||
|
||||
def test_tool_message_text_only_still_collapsed(self):
|
||||
config = DeepSeekChatConfig()
|
||||
message = {"role": "tool", "content": [{"type": "text", "text": "plain"}]}
|
||||
out = config._forward_or_collapse_content(message=message, forward_images=True)
|
||||
assert out.get("content") == "plain"
|
||||
|
|
@ -187,7 +187,9 @@ def test_response_cost_calculator_keeps_optional_params_out_of_hidden_params():
|
|||
assert optional_params["aws_session_token"] == "session-secret"
|
||||
|
||||
|
||||
def test_embedding_success_logging_and_spend_log_carry_no_forwarded_credentials(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_embedding_success_logging_and_spend_log_carry_no_forwarded_credentials(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.spend_tracking.spend_tracking_utils import _get_proxy_server_request_for_spend_logs_payload
|
||||
|
||||
|
|
@ -236,10 +238,6 @@ def test_embedding_success_logging_and_spend_log_carry_no_forwarded_credentials(
|
|||
assert logging_obj.optional_params["extra_headers"] == {"x-goog-api-key": "goog-secret"}
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_realtime_stream_combines_text_and_audio_token_details():
|
||||
"""Realtime response.done usage with input_token_details / output_token_details."""
|
||||
from litellm.cost_calculator import RealtimeAPITokenUsageProcessor
|
||||
|
|
@ -1354,8 +1352,6 @@ def test_bedrock_cost_calculator_comparison_with_without_cache():
|
|||
print(f"Cost with cache: {cost_with_cache}")
|
||||
|
||||
|
||||
|
||||
|
||||
def test_gemini_25_explicit_caching_cost_direct_usage():
|
||||
"""
|
||||
Test that Gemini 2.5 models correctly calculate costs with explicit caching.
|
||||
|
|
@ -1924,8 +1920,6 @@ def test_cost_margin_with_discount(monkeypatch):
|
|||
print(f" - Expected: ${expected_cost:.6f}")
|
||||
|
||||
|
||||
|
||||
|
||||
def test_completion_cost_extracts_service_tier_from_response(_local_model_cost_map):
|
||||
"""Test that completion_cost extracts service_tier from completion_response object."""
|
||||
from litellm import completion_cost
|
||||
|
|
@ -2675,8 +2669,6 @@ def test_gemini_without_cache_tokens_details():
|
|||
print("✅ Gemini without cacheTokensDetails works correctly")
|
||||
|
||||
|
||||
|
||||
|
||||
def test_additional_costs_only_for_azure_ai(_local_model_cost_map):
|
||||
"""
|
||||
Test that _get_additional_costs is only called for azure_ai provider.
|
||||
|
|
@ -3221,9 +3213,7 @@ def test_cost_per_token_resolves_per_second_rate_precedence(
|
|||
|
||||
model: Final = "test-chat-per-second-rate-precedence"
|
||||
entry: Final = {**pricing_fields, "litellm_provider": "together_ai", "mode": "chat"}
|
||||
litellm.register_model(
|
||||
model_cost={model: entry}
|
||||
)
|
||||
litellm.register_model(model_cost={model: entry})
|
||||
|
||||
assert cost_per_token(
|
||||
model=model,
|
||||
|
|
@ -3648,6 +3638,42 @@ def test_combine_usage_objects_sums_mirrored_cache_write_fields_once():
|
|||
assert combined_pair.prompt_tokens_details.cache_creation_tokens == 100
|
||||
|
||||
|
||||
def test_select_model_name_selects_character_priced_deployment(_local_model_cost_map):
|
||||
"""
|
||||
A deployment whose only rate is input_cost_per_character must be selected
|
||||
by router_model_id: the aspeech cost path resolves its price through the
|
||||
deployment entry, and a character-only entry failing the "prices anything"
|
||||
check silently produced spend = 0 (issue #44200).
|
||||
"""
|
||||
from litellm.cost_calculator import _select_model_name_for_cost_calc
|
||||
|
||||
router_model_id = "openai/qwen-audio-3.1-tts-flash-uuid"
|
||||
litellm.model_cost[router_model_id] = {
|
||||
"input_cost_per_character": 1e-8,
|
||||
"output_cost_per_character": 0.0,
|
||||
"litellm_provider": "openai",
|
||||
}
|
||||
|
||||
selected = _select_model_name_for_cost_calc(
|
||||
model="qwen-audio-3.1-tts-flash",
|
||||
completion_response=None,
|
||||
custom_pricing=True,
|
||||
custom_llm_provider="openai",
|
||||
router_model_id=router_model_id,
|
||||
)
|
||||
|
||||
assert selected == router_model_id
|
||||
|
||||
|
||||
def test_cost_map_entry_prices_anything_recognizes_character_rates():
|
||||
from litellm.cost_calculator import _cost_map_entry_prices_anything
|
||||
|
||||
assert _cost_map_entry_prices_anything({"input_cost_per_character": 1e-8}) is True
|
||||
assert _cost_map_entry_prices_anything({"output_cost_per_character": 0.0}) is True
|
||||
assert _cost_map_entry_prices_anything({"input_cost_per_token": 1e-6}) is True
|
||||
assert _cost_map_entry_prices_anything({"mode": "audio_speech"}) is False
|
||||
|
||||
|
||||
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.
|
||||
|
|
@ -4815,9 +4841,7 @@ def test_xai_batch_tier_discounts_the_long_context_rate_like_the_flat_batch_rate
|
|||
assert info[f"{prefix}_above_200k_tokens_batches"] < info[f"{prefix}_above_200k_tokens"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("prompt_tokens", "tier"), [(200_000, "_above_200k_tokens_batches"), (199_999, "_batches")]
|
||||
)
|
||||
@pytest.mark.parametrize(("prompt_tokens", "tier"), [(200_000, "_above_200k_tokens_batches"), (199_999, "_batches")])
|
||||
def test_xai_batch_cost_calculator_bills_the_200k_batch_tier_inclusively(
|
||||
_local_model_cost_map: None, prompt_tokens: int, tier: str
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -3507,7 +3507,7 @@ def test_completion_streaming_iterator_adopts_the_deployment_that_served_a_neste
|
|||
}
|
||||
return chunk
|
||||
|
||||
with patch.object(router, "function_with_fallbacks", return_value=NestedFallbackStream()):
|
||||
with patch.object(router, "async_function_with_fallbacks_common_utils", return_value=NestedFallbackStream()):
|
||||
result = router._completion_streaming_iterator(
|
||||
model_response=FailedStream(),
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
|
|
@ -3685,7 +3685,7 @@ def test_completion_streaming_iterator_adopts_fallback_response_headers():
|
|||
def __iter__(self):
|
||||
return iter([])
|
||||
|
||||
with patch.object(router, "function_with_fallbacks", return_value=FallbackStream()):
|
||||
with patch.object(router, "async_function_with_fallbacks_common_utils", return_value=FallbackStream()):
|
||||
result = router._completion_streaming_iterator(
|
||||
model_response=FailedStream(),
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
|
|
@ -3752,7 +3752,7 @@ def test_completion_streaming_iterator_fallback_on_429():
|
|||
|
||||
with patch.object(
|
||||
router,
|
||||
"function_with_fallbacks",
|
||||
"async_function_with_fallbacks_common_utils",
|
||||
return_value=mock_fallback_response,
|
||||
) as mock_fallback:
|
||||
result = router._completion_streaming_iterator(
|
||||
|
|
@ -3765,10 +3765,99 @@ def test_completion_streaming_iterator_fallback_on_429():
|
|||
|
||||
assert mock_fallback.called
|
||||
call_kwargs = mock_fallback.call_args
|
||||
assert mock_fallback.call_args.args[0] is rate_limit_error
|
||||
# Pre-first-chunk: should use original messages, no continuation prompt
|
||||
assert call_kwargs.kwargs.get("messages") == messages
|
||||
assert call_kwargs.kwargs.get("kwargs", {}).get("messages") == messages
|
||||
# Verify original_function is _completion (sync)
|
||||
assert call_kwargs.kwargs.get("original_function") == router._completion
|
||||
assert call_kwargs.kwargs.get("kwargs", {}).get("original_function") == router._completion
|
||||
|
||||
|
||||
def test_completion_streaming_iterator_routes_mid_stream_fallback_through_common_utils():
|
||||
"""Regression (#43945): the sync mid-stream fallback re-entry must hand the
|
||||
MidStreamFallbackError to async_function_with_fallbacks_common_utils, like the async
|
||||
twin does. Calling function_with_fallbacks() instead re-runs the original (failing)
|
||||
group first, and because each nested Router.completion() wraps its own stream in this
|
||||
same iterator, every retry fails again only while being iterated — recursing until
|
||||
the stack runs out (measured: 478 requests to the failing group, then
|
||||
InternalServerError, with a healthy fallback configured)."""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "gpt-4", "api_key": "fake-key"},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
messages = [{"role": "user", "content": "Test"}]
|
||||
initial_kwargs = {"model": "gpt-4", "stream": True}
|
||||
|
||||
pre_first_chunk_error = MidStreamFallbackError(
|
||||
message="upstream died before the first chunk",
|
||||
model="gpt-4",
|
||||
llm_provider="openai",
|
||||
generated_content="",
|
||||
is_pre_first_chunk=True,
|
||||
)
|
||||
|
||||
class SyncIteratorImmediateError:
|
||||
def __init__(self):
|
||||
self.model = "gpt-4"
|
||||
self.custom_llm_provider = "openai"
|
||||
self.logging_obj = MagicMock()
|
||||
self.chunks = []
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
raise pre_first_chunk_error
|
||||
|
||||
class FallbackStream:
|
||||
def __init__(self):
|
||||
self._chunks = iter(
|
||||
[
|
||||
litellm.ModelResponseStream(
|
||||
choices=[{"index": 0, "delta": {"content": "from the fallback"}}]
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
return next(self._chunks)
|
||||
|
||||
with patch.object(
|
||||
router,
|
||||
"async_function_with_fallbacks_common_utils",
|
||||
return_value=FallbackStream(),
|
||||
) as mock_utils:
|
||||
with patch.object(router, "function_with_fallbacks") as mock_function_with_fallbacks:
|
||||
result = router._completion_streaming_iterator(
|
||||
model_response=SyncIteratorImmediateError(),
|
||||
messages=messages,
|
||||
initial_kwargs=initial_kwargs,
|
||||
)
|
||||
|
||||
collected_chunks = list(result)
|
||||
|
||||
assert mock_utils.called
|
||||
# the triggering error must reach the common utils, so cooldowns apply and
|
||||
# the walk starts from the fallback list, not the failing group
|
||||
assert mock_utils.call_args.args[0] is pre_first_chunk_error
|
||||
assert mock_utils.call_args.kwargs.get("kwargs", {}).get("messages") == messages
|
||||
assert not mock_function_with_fallbacks.called, (
|
||||
"re-running the original group is what recurses; common utils already "
|
||||
"excludes the deployment that raised"
|
||||
)
|
||||
|
||||
assert len(collected_chunks) == 1
|
||||
|
||||
|
||||
def test_completion_streaming_iterator_preserves_hidden_params():
|
||||
|
|
@ -3980,7 +4069,7 @@ def test_completion_streaming_iterator_reraises_mid_chunk_error_with_no_text_con
|
|||
|
||||
mock_response = SyncIteratorNoTextChunkError()
|
||||
|
||||
with patch.object(router, "function_with_fallbacks") as mock_fallback:
|
||||
with patch.object(router, "async_function_with_fallbacks_common_utils") as mock_fallback:
|
||||
result = router._completion_streaming_iterator(
|
||||
model_response=mock_response,
|
||||
messages=messages,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue