This commit is contained in:
JingHao-Leon 2026-10-04 12:47:45 -07:00 • committed by GitHub
commit a359693602
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 660 additions and 181 deletions

View file

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

View file

@ -897,6 +897,14 @@ def _set_response_cost_attr(span: "Span", standard_logging_payload) -> None:
cost: Final = standard_logging_payload.get("response_cost")
if cost is None:
return
# A 0 cost is ambiguous between "genuinely free" and "pricing failed".
# When the pricing-failure debug info is present, omit the attributes
# rather than reporting a fake $0 that cost backends cannot distinguish
# from a free model (issue #44186) — e.g. Langfuse prefers an ingested
# cost over its own price knowledge, so a fake 0 hides real spend there.
failure_debug: Final = standard_logging_payload.get("response_cost_failure_debug_info")
if failure_debug:
return
try:
cost_value: Final = float(cost)
except (TypeError, ValueError):

View file

@ -2128,6 +2128,20 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
guardrail_span.end(end_time=self._to_ns(end_time_datetime))
def _get_error_status_description(self, kwargs) -> str | None:
"""Mirror the error.message attribute text for the ERROR status
description, so backends that render the span status show the actual
error instead of an empty message (issue #44184)."""
standard_logging_payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object")
if standard_logging_payload is None:
return None
error_information: Final = standard_logging_payload.get("error_information")
if error_information is None:
error_str: Final = standard_logging_payload.get("error_str")
return error_str if isinstance(error_str, str) and error_str else None
message: Final = error_information.get("error_message")
return message if isinstance(message, str) and message else None
def _handle_failure(self, kwargs, response_obj, start_time, end_time):
from opentelemetry.trace import Status, StatusCode
@ -2157,6 +2171,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
parent_otel_span = None # Ignore parent spans from other providers
_parent_context = None
# Mirror the error.message attribute text onto the ERROR status, so
# backends that render the span status don't show an empty error
# (issue #44184).
error_status_description: Final = self._get_error_status_description(kwargs)
# Decide whether to create a primary span
# Always create if no parent span exists (backward compatibility)
# OR if USE_OTEL_LITELLM_REQUEST_SPAN is explicitly enabled
@ -2174,7 +2193,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
if self._gen_ai_semconv_latest_experimental:
span_kwargs["kind"] = self.span_kind.CLIENT
span = otel_tracer.start_span(**span_kwargs)
span.set_status(Status(StatusCode.ERROR))
span.set_status(Status(status_code=StatusCode.ERROR, description=error_status_description))
self.set_attributes(span, kwargs, response_obj)
# Record exception information using OTEL standard method
@ -2187,7 +2206,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
# Only set attributes if the span is still recording (not closed)
# Note: parent_otel_span is guaranteed to be not None here
if parent_otel_span.is_recording():
parent_otel_span.set_status(Status(StatusCode.ERROR))
parent_otel_span.set_status(Status(status_code=StatusCode.ERROR, description=error_status_description))
self.set_attributes(parent_otel_span, kwargs, response_obj)
self._record_exception_on_span(span=parent_otel_span, kwargs=kwargs)

View file

@ -479,7 +479,14 @@ class LLMCallSpanData:
usage=LLMUsage.from_standard_logging_payload(payload),
finish_reasons=finish_reasons,
error=_parse_error(payload),
response_cost=as_float(payload.get("response_cost")),
# response_cost = 0 is ambiguous between a free model and a
# pricing failure; when the pricing-failure debug info is
# present, report None (unknown) instead of a fake 0 (issue
# #44186) — e.g. Langfuse prefers an ingested cost over its own
# price knowledge, so a fake 0 hides real spend there.
response_cost=(
as_float(payload.get("response_cost")) if not payload.get("response_cost_failure_debug_info") else None
),
cost=LLMCost.from_breakdown(cast("Mapping[str, object] | None", payload.get("cost_breakdown"))),
server=ServerInfo.from_api_base(context.api_base),
identity=context.identity,
@ -574,7 +581,10 @@ class MCPToolCallSpanData:
_json_or_none(meta.get("result")) if capture_content and meta.get("result") is not None else None
),
error=_parse_error(payload),
response_cost=as_float(payload.get("response_cost")),
# Same 0-vs-pricing-failure ambiguity as the LLM span above.
response_cost=(
as_float(payload.get("response_cost")) if not payload.get("response_cost_failure_debug_info") else None
),
identity=RequestContext.from_standard_logging_payload(payload).identity,
)

View file

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

View file

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

View file

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

View file

@ -87,20 +87,16 @@ class TestOpentelemetryUnitTests(BaseLoggingCallbackTest):
detected_context, detected_span = otel_integration._get_span_context(kwargs)
# Assert: Should detect the active span
assert (
detected_span is not None
), "Should detect active span from global context"
assert (
detected_span is parent_span
), "Detected span should be the active parent span"
assert detected_span is not None, "Should detect active span from global context"
assert detected_span is parent_span, "Detected span should be the active parent span"
detected_span_context = detected_span.get_span_context()
assert (
detected_span_context.trace_id == parent_span_context.trace_id
), "Detected span should have same trace_id as parent"
assert (
detected_span_context.span_id == parent_span_context.span_id
), "Detected span should have same span_id as parent"
assert detected_span_context.trace_id == parent_span_context.trace_id, (
"Detected span should have same trace_id as parent"
)
assert detected_span_context.span_id == parent_span_context.span_id, (
"Detected span should have same span_id as parent"
)
def test_record_exception_on_span(self):
"""
@ -162,9 +158,9 @@ class TestOpentelemetryUnitTests(BaseLoggingCallbackTest):
actual_calls = [call.args for call in mock_span.set_attribute.call_args_list]
for expected_call in expected_calls:
assert (
expected_call in actual_calls
), f"Expected set_attribute call {expected_call} not found in actual calls: {actual_calls}"
assert expected_call in actual_calls, (
f"Expected set_attribute call {expected_call} not found in actual calls: {actual_calls}"
)
def test_record_exception_on_span_with_fallback(self):
"""
@ -205,6 +201,42 @@ class TestOpentelemetryUnitTests(BaseLoggingCallbackTest):
mock_span.record_exception.assert_called_once_with(test_exception)
# Assert: error.message should be set from error_str using ErrorAttributes constant
mock_span.set_attribute.assert_called_with(
ErrorAttributes.ERROR_MESSAGE, "Fallback error message"
)
mock_span.set_attribute.assert_called_with(ErrorAttributes.ERROR_MESSAGE, "Fallback error message")
def test_get_error_status_description_from_error_information(self):
"""The ERROR status description mirrors the error_message attribute
text, so backends that render the span status show the actual error
(issue #44184)."""
from litellm.integrations.opentelemetry import OpenTelemetry
otel_integration = OpenTelemetry()
kwargs = {
"standard_logging_object": {
"error_information": {
"error_code": "500",
"error_class": "ValueError",
"error_message": "Test error message",
},
"error_str": "Test error message",
},
}
assert otel_integration._get_error_status_description(kwargs) == "Test error message"
def test_get_error_status_description_falls_back_to_error_str(self):
from litellm.integrations.opentelemetry import OpenTelemetry
otel_integration = OpenTelemetry()
kwargs = {
"standard_logging_object": {
"error_information": None,
"error_str": "Fallback error message",
},
}
assert otel_integration._get_error_status_description(kwargs) == "Fallback error message"
def test_get_error_status_description_none_when_no_error(self):
from litellm.integrations.opentelemetry import OpenTelemetry
otel_integration = OpenTelemetry()
assert otel_integration._get_error_status_description({"standard_logging_object": None}) is None
assert otel_integration._get_error_status_description({}) is None

View file

@ -70,9 +70,7 @@ def test_arize_set_attributes():
# Simulated LLM response object
response_obj = ModelResponse(
usage={"total_tokens": 100, "completion_tokens": 60, "prompt_tokens": 40},
choices=[
Choices(message={"role": "assistant", "content": "Basic Response Content"})
],
choices=[Choices(message={"role": "assistant", "content": "Basic Response Content"})],
model="gpt-4o",
id="chatcmpl-ID",
)
@ -89,9 +87,7 @@ def test_arize_set_attributes():
assert span.set_attribute.call_count == 26
# Metadata attached to the span
span.set_attribute.assert_any_call(
SpanAttributes.METADATA, json.dumps({"key_1": "value_1", "key_2": None})
)
span.set_attribute.assert_any_call(SpanAttributes.METADATA, json.dumps({"key_1": "value_1", "key_2": None}))
# Basic LLM information
span.set_attribute.assert_any_call(SpanAttributes.LLM_MODEL_NAME, "gpt-4o")
@ -114,16 +110,12 @@ def test_arize_set_attributes():
span.set_attribute.assert_any_call(SpanAttributes.OPENINFERENCE_SPAN_KIND, "LLM")
# And TOOL must never be written for an LLM chat completion call.
span_kind_writes = [
c.args[1]
for c in span.set_attribute.call_args_list
if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND
c.args[1] for c in span.set_attribute.call_args_list if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND
]
assert "TOOL" not in span_kind_writes
# Request message content and metadata
span.set_attribute.assert_any_call(
SpanAttributes.INPUT_VALUE, "Basic Request Content"
)
span.set_attribute.assert_any_call(SpanAttributes.INPUT_VALUE, "Basic Request Content")
span.set_attribute.assert_any_call(
f"{SpanAttributes.LLM_INPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_ROLE}",
"user",
@ -134,9 +126,7 @@ def test_arize_set_attributes():
)
# Tool call definitions and function names
span.set_attribute.assert_any_call(
f"{SpanAttributes.LLM_TOOLS}.0.name", "get_weather"
)
span.set_attribute.assert_any_call(f"{SpanAttributes.LLM_TOOLS}.0.name", "get_weather")
span.set_attribute.assert_any_call(
f"{SpanAttributes.LLM_TOOLS}.0.description",
"Fetches weather details.",
@ -146,26 +136,20 @@ def test_arize_set_attributes():
json.dumps(
{
"type": "object",
"properties": {
"location": {"type": "string", "description": "City name"}
},
"properties": {"location": {"type": "string", "description": "City name"}},
"required": ["location"],
}
),
)
# Invocation parameters
span.set_attribute.assert_any_call(
SpanAttributes.LLM_INVOCATION_PARAMETERS, '{"user": "test_user"}'
)
span.set_attribute.assert_any_call(SpanAttributes.LLM_INVOCATION_PARAMETERS, '{"user": "test_user"}')
# User ID
span.set_attribute.assert_any_call(SpanAttributes.USER_ID, "test_user")
# Output message content
span.set_attribute.assert_any_call(
SpanAttributes.OUTPUT_VALUE, "Basic Response Content"
)
span.set_attribute.assert_any_call(SpanAttributes.OUTPUT_VALUE, "Basic Response Content")
span.set_attribute.assert_any_call(
f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_ROLE}",
"assistant",
@ -228,9 +212,7 @@ def test_arize_set_attributes_responses_api():
ResponseReasoningItem(
id="reasoning-001",
type="reasoning",
summary=[
Summary(text="First, I need to analyze...", type="summary_text")
],
summary=[Summary(text="First, I need to analyze...", type="summary_text")],
),
ResponseOutputMessage(
id="msg-001",
@ -277,9 +259,7 @@ def test_arize_set_attributes_responses_api():
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_TOTAL, 370)
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, 250)
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 120)
span.set_attribute.assert_any_call(
SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180
)
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180)
def test_set_usage_outputs_pydantic_completion_usage():
@ -327,9 +307,7 @@ def test_set_usage_outputs_pydantic_completion_usage():
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 40)
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, 60)
# reasoning_tokens for chat completions live in completion_tokens_details
span.set_attribute.assert_any_call(
SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 25
)
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 25)
def test_set_usage_outputs_pydantic_response_api_usage():
@ -362,9 +340,7 @@ def test_set_usage_outputs_pydantic_response_api_usage():
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_TOTAL, 370)
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 120)
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, 250)
span.set_attribute.assert_any_call(
SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180
)
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180)
class TestArizeLogger(CustomLogger):
@ -375,16 +351,12 @@ class TestArizeLogger(CustomLogger):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.standard_callback_dynamic_params: Optional[
StandardCallbackDynamicParams
] = None
self.standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = None
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
# Capture dynamic params and print them for verification
print("logged kwargs", json.dumps(kwargs, indent=4, default=str))
self.standard_callback_dynamic_params = kwargs.get(
"standard_callback_dynamic_params"
)
self.standard_callback_dynamic_params = kwargs.get("standard_callback_dynamic_params")
@pytest.mark.asyncio
@ -410,14 +382,8 @@ async def test_arize_dynamic_params():
# Assert dynamic parameters were received in the callback
assert test_arize_logger.standard_callback_dynamic_params is not None
assert (
test_arize_logger.standard_callback_dynamic_params.get("arize_api_key")
== "test_api_key_dynamic"
)
assert (
test_arize_logger.standard_callback_dynamic_params.get("arize_space_key")
== "test_space_key_dynamic"
)
assert test_arize_logger.standard_callback_dynamic_params.get("arize_api_key") == "test_api_key_dynamic"
assert test_arize_logger.standard_callback_dynamic_params.get("arize_space_key") == "test_space_key_dynamic"
def test_construct_dynamic_arize_headers():
@ -428,9 +394,7 @@ def test_construct_dynamic_arize_headers():
from litellm.types.utils import StandardCallbackDynamicParams
# Test with all parameters present
dynamic_params_full = StandardCallbackDynamicParams(
arize_api_key="test_api_key", arize_space_id="test_space_id"
)
dynamic_params_full = StandardCallbackDynamicParams(arize_api_key="test_api_key", arize_space_id="test_space_id")
arize_logger = ArizeLogger()
headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_full)
@ -438,9 +402,7 @@ def test_construct_dynamic_arize_headers():
assert headers == expected_headers
# Test with only space_id
dynamic_params_space_id_only = StandardCallbackDynamicParams(
arize_space_id="test_space_id"
)
dynamic_params_space_id_only = StandardCallbackDynamicParams(arize_space_id="test_space_id")
headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_space_id_only)
expected_headers = {"arize-space-id": "test_space_id"}
@ -456,9 +418,7 @@ def test_construct_dynamic_arize_headers():
dynamic_params_space_key_and_api_key = StandardCallbackDynamicParams(
arize_space_key="test_space_key", arize_api_key="test_api_key"
)
headers = arize_logger.construct_dynamic_otel_headers(
dynamic_params_space_key_and_api_key
)
headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_space_key_and_api_key)
expected_headers = {"arize-space-id": "test_space_key", "api_key": "test_api_key"}
@ -528,9 +488,7 @@ def test_arize_emits_no_cache_tokens_when_absent():
from litellm.integrations.arize._utils import _set_usage_outputs
span = MagicMock()
response_obj = {
"usage": {"total_tokens": 10, "completion_tokens": 4, "prompt_tokens": 6}
}
response_obj = {"usage": {"total_tokens": 10, "completion_tokens": 4, "prompt_tokens": 6}}
_set_usage_outputs(span, response_obj, SpanAttributes)
attrs = _collect_calls(span)
assert SpanAttributes.LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_READ not in attrs
@ -542,14 +500,8 @@ def test_passthrough_call_type_resolves_to_llm_span_kind():
from litellm.integrations._types.open_inference import OpenInferenceSpanKindValues
from litellm.integrations.arize._utils import _infer_open_inference_span_kind
assert (
_infer_open_inference_span_kind("allm_passthrough_route")
== OpenInferenceSpanKindValues.LLM.value
)
assert (
_infer_open_inference_span_kind("llm_passthrough_route")
== OpenInferenceSpanKindValues.LLM.value
)
assert _infer_open_inference_span_kind("allm_passthrough_route") == OpenInferenceSpanKindValues.LLM.value
assert _infer_open_inference_span_kind("llm_passthrough_route") == OpenInferenceSpanKindValues.LLM.value
def test_arize_chat_completion_with_tools_stays_llm_span_kind():
@ -605,9 +557,7 @@ def test_arize_chat_completion_with_tools_stays_llm_span_kind():
ArizeLogger.set_arize_attributes(span, kwargs, response_obj)
span_kind_writes = [
c.args[1]
for c in span.set_attribute.call_args_list
if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND
c.args[1] for c in span.set_attribute.call_args_list if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND
]
assert span_kind_writes, "span.kind must be written"
assert all(v == "LLM" for v in span_kind_writes)
@ -659,13 +609,8 @@ def test_arize_emits_assistant_tool_calls_on_output_message():
attrs = _collect_calls(span)
base = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_TOOL_CALLS}.0"
assert attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_ID}"] == "call_abc"
assert (
attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_NAME}"] == "get_weather"
)
assert (
attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_ARGUMENTS_JSON}"]
== '{"location": "SF"}'
)
assert attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_NAME}"] == "get_weather"
assert attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_ARGUMENTS_JSON}"] == '{"location": "SF"}'
def test_arize_output_value_falls_back_to_tool_calls_summary():
@ -818,9 +763,7 @@ def test_arize_emits_tool_call_id_and_name_on_input_tool_message():
assert attrs[f"{assistant_base}.{ToolCallAttributes.TOOL_CALL_ID}"] == "call_abc"
# Tool message at index 2
tool_prefix = f"{SpanAttributes.LLM_INPUT_MESSAGES}.2"
assert (
attrs[f"{tool_prefix}.{MessageAttributes.MESSAGE_TOOL_CALL_ID}"] == "call_abc"
)
assert attrs[f"{tool_prefix}.{MessageAttributes.MESSAGE_TOOL_CALL_ID}"] == "call_abc"
assert attrs[f"{tool_prefix}.{MessageAttributes.MESSAGE_NAME}"] == "get_weather"
@ -866,10 +809,7 @@ def test_arize_emits_multimodal_input_contents():
assert attrs[f"{base}.0.message_content.type"] == "text"
assert attrs[f"{base}.0.message_content.text"] == "What is in this image?"
assert attrs[f"{base}.1.message_content.type"] == "image"
assert (
attrs[f"{base}.1.message_content.image.image.url"]
== "https://example.com/cat.png"
)
assert attrs[f"{base}.1.message_content.image.image.url"] == "https://example.com/cat.png"
def test_arize_emits_session_and_user_attrs_from_metadata():
@ -974,11 +914,7 @@ def test_arize_does_not_overwrite_user_id_from_optional_params():
id="r2",
)
ArizeLogger.set_arize_attributes(span, kwargs, response_obj)
user_id_writes = [
c.args[1]
for c in span.set_attribute.call_args_list
if c.args[0] == SpanAttributes.USER_ID
]
user_id_writes = [c.args[1] for c in span.set_attribute.call_args_list if c.args[0] == SpanAttributes.USER_ID]
assert "from_metadata" not in user_id_writes
@ -1013,6 +949,72 @@ def test_arize_emits_response_cost():
assert attrs["llm.response.cost"] == 0.0012345 # legacy key still emitted
def test_arize_omits_cost_when_pricing_failed():
"""response_cost = 0 plus pricing-failure debug info means LiteLLM could
not price the call — omit the cost attributes entirely: backends cannot
tell a fake 0 from a free model (issue #44186)."""
from unittest.mock import MagicMock
from litellm.types.utils import Choices, ModelResponse
span = MagicMock()
kwargs = {
"model": "gpt-4o",
"messages": [{"role": "user", "content": "hi"}],
"standard_logging_object": {
"model_parameters": {},
"metadata": {},
"call_type": "completion",
"response_cost": 0.0,
"response_cost_failure_debug_info": {"error_str": "model not in cost map"},
},
"optional_params": {},
"litellm_params": {"custom_llm_provider": "openai"},
}
response_obj = ModelResponse(
usage={"total_tokens": 4, "completion_tokens": 2, "prompt_tokens": 2},
choices=[Choices(message={"role": "assistant", "content": "hello"})],
model="gpt-4o",
id="r4",
)
ArizeLogger.set_arize_attributes(span, kwargs, response_obj)
attrs = _collect_calls(span)
assert "llm.cost.total" not in attrs
assert "llm.response.cost" not in attrs
def test_arize_emits_zero_cost_for_free_model():
"""A genuine 0 (free model, no pricing-failure info) must still be
emitted — the guard only fires on the pricing-failure marker."""
from unittest.mock import MagicMock
from litellm.types.utils import Choices, ModelResponse
span = MagicMock()
kwargs = {
"model": "gpt-4o",
"messages": [{"role": "user", "content": "hi"}],
"standard_logging_object": {
"model_parameters": {},
"metadata": {},
"call_type": "completion",
"response_cost": 0.0,
},
"optional_params": {},
"litellm_params": {"custom_llm_provider": "openai"},
}
response_obj = ModelResponse(
usage={"total_tokens": 4, "completion_tokens": 2, "prompt_tokens": 2},
choices=[Choices(message={"role": "assistant", "content": "hello"})],
model="gpt-4o",
id="r5",
)
ArizeLogger.set_arize_attributes(span, kwargs, response_obj)
attrs = _collect_calls(span)
assert attrs["llm.cost.total"] == 0.0
assert attrs["llm.response.cost"] == 0.0
def test_arize_passthrough_bedrock_anthropic_normalization():
"""Bedrock-Anthropic passthrough: input/output text must be set so the
span renders something other than raw provider attrs."""
@ -1048,9 +1050,7 @@ def test_arize_passthrough_bedrock_anthropic_normalization():
"complete_input_dict": {
"anthropic_version": "bedrock-2023-05-31",
"max_tokens": 64,
"messages": [
{"role": "user", "content": "What is the capital of France?"}
],
"messages": [{"role": "user", "content": "What is the capital of France?"}],
}
},
"standard_logging_object": {
@ -1068,19 +1068,13 @@ def test_arize_passthrough_bedrock_anthropic_normalization():
assert attrs[SpanAttributes.INPUT_VALUE] == "What is the capital of France?"
msg0 = f"{SpanAttributes.LLM_INPUT_MESSAGES}.0"
assert attrs[f"{msg0}.{MessageAttributes.MESSAGE_ROLE}"] == "user"
assert (
attrs[f"{msg0}.{MessageAttributes.MESSAGE_CONTENT}"]
== "What is the capital of France?"
)
assert attrs[f"{msg0}.{MessageAttributes.MESSAGE_CONTENT}"] == "What is the capital of France?"
# Output rendering (Anthropic content[].text)
assert attrs[SpanAttributes.OUTPUT_VALUE] == "The capital of France is Paris."
out0 = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0"
assert attrs[f"{out0}.{MessageAttributes.MESSAGE_ROLE}"] == "assistant"
assert (
attrs[f"{out0}.{MessageAttributes.MESSAGE_CONTENT}"]
== "The capital of France is Paris."
)
assert attrs[f"{out0}.{MessageAttributes.MESSAGE_CONTENT}"] == "The capital of France is Paris."
# Token counts (Bedrock input_tokens/output_tokens) — extracted via
# coercion of the non-dict response.
@ -1089,9 +1083,7 @@ def test_arize_passthrough_bedrock_anthropic_normalization():
# Span kind defended even though the call_type is a passthrough variant.
span_kind_writes = [
c.args[1]
for c in span.set_attribute.call_args_list
if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND
c.args[1] for c in span.set_attribute.call_args_list if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND
]
assert span_kind_writes # at least one
assert all(v == "LLM" for v in span_kind_writes)
@ -1109,11 +1101,7 @@ def test_arize_passthrough_call_type_does_not_run_on_chat_completion():
span = MagicMock()
_maybe_normalize_passthrough(
span,
{
"additional_args": {
"complete_input_dict": {"messages": [{"role": "user", "content": "x"}]}
}
},
{"additional_args": {"complete_input_dict": {"messages": [{"role": "user", "content": "x"}]}}},
{"choices": [{"message": {"role": "assistant", "content": "y"}}]},
{"choices": [{"message": {"role": "assistant", "content": "y"}}]},
{"call_type": "completion"},
@ -1133,11 +1121,7 @@ def test_arize_passthrough_skipped_when_message_redaction_enabled():
span = MagicMock()
kwargs = {
"additional_args": {
"complete_input_dict": {
"messages": [
{"role": "user", "content": "Patient John Doe, SSN 123-45-6789"}
]
}
"complete_input_dict": {"messages": [{"role": "user", "content": "Patient John Doe, SSN 123-45-6789"}]}
},
# Enables redaction via the dynamic-param path inside
# should_redact_message_logging(), without touching globals.
@ -1211,9 +1195,7 @@ def test_arize_mcp_call_tool_result_does_not_break_attribute_setting():
"optional_params": {},
"litellm_params": {"custom_llm_provider": "mcp"},
}
response_obj = CallToolResult(
content=[TextContent(type="text", text="sunny, 21C")], isError=False
)
response_obj = CallToolResult(content=[TextContent(type="text", text="sunny, 21C")], isError=False)
ArizeLogger.set_arize_attributes(span, kwargs, response_obj)
@ -1295,9 +1277,7 @@ def test_arize_mcp_tool_span_renders_name_input_and_output():
from mcp.types import CallToolResult, TextContent
span = MagicMock()
response_obj = CallToolResult(
content=[TextContent(type="text", text="sunny, 21C")], isError=False
)
response_obj = CallToolResult(content=[TextContent(type="text", text="sunny, 21C")], isError=False)
ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj)
@ -1336,9 +1316,7 @@ def test_arize_mcp_tool_span_respects_message_redaction():
from mcp.types import CallToolResult, TextContent
span = MagicMock()
response_obj = CallToolResult(
content=[TextContent(type="text", text="SSN 123-45-6789")], isError=False
)
response_obj = CallToolResult(content=[TextContent(type="text", text="SSN 123-45-6789")], isError=False)
ArizeLogger.set_arize_attributes(
span,

View file

@ -90,6 +90,22 @@ def test_llm_call_span_cost_breakdown():
assert f"{LiteLLM.COST_PREFIX}margin_total_amount" not in a
def test_llm_call_span_pricing_failure_reports_unknown_cost():
"""A 0 response_cost alongside pricing-failure debug info means LiteLLM
could not price the call — report unknown (None), not a fake 0 that cost
backends cannot distinguish from a free model (issue #44186)."""
data = LLMCallSpanData.from_standard_logging_payload(
_payload(response_cost=0.0, response_cost_failure_debug_info={"error_str": "model not in cost map"})
)
assert data.response_cost is None
def test_llm_call_span_zero_cost_without_failure_info_is_kept():
"""A genuine 0 (free model) without pricing-failure info stays 0."""
data = LLMCallSpanData.from_standard_logging_payload(_payload(response_cost=0.0))
assert data.response_cost == 0.0
def test_tracer_scope_carries_litellm_version():
from litellm._version import version as litellm_version

View file

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

View file

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

View file

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

View file

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