mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 36f0a0fe6d into 461a58c40a
This commit is contained in:
commit
12c5ca50ba
14 changed files with 1328 additions and 981 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",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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__"):
|
||||
|
|
@ -5569,10 +5580,12 @@ class Router:
|
|||
model, initial_kwargs
|
||||
)
|
||||
buffered_lifecycle_chunks: tuple[bytes, ...] = () # rebind-ok: flushed once committed or on decline
|
||||
chunks_sent_to_client = 0 # any frame that reached the client blocks a same-group retry
|
||||
try:
|
||||
async for chunk in source_iterator:
|
||||
if _anthropic_stream_forwards_ping_live(chunk, has_generated_content):
|
||||
yield chunk
|
||||
chunks_sent_to_client += 1
|
||||
continue
|
||||
if _anthropic_stream_commits_now(chunk, has_generated_content, len(buffered_lifecycle_chunks)):
|
||||
has_generated_content = True
|
||||
|
|
@ -5627,9 +5640,36 @@ class Router:
|
|||
yield buffered_chunk
|
||||
buffered_lifecycle_chunks = ()
|
||||
yield chunk
|
||||
chunks_sent_to_client += 1
|
||||
for buffered_chunk in buffered_lifecycle_chunks:
|
||||
yield buffered_chunk
|
||||
except Exception as stream_error: # noqa: BLE001 # any raised provider error must reach the fallback gate
|
||||
# A transport drop before ANY frame reached the client never
|
||||
# reaches the retry machinery — the error is raised from the
|
||||
# stream iterator, outside the retried call boundary — and
|
||||
# surfaced as a 500 after a single upstream attempt (issue
|
||||
# #44238). With no fallbacks configured there is nothing to
|
||||
# hand the failure to, so retry the same group once; once any
|
||||
# frame reached the client (or fallbacks are configured, whose
|
||||
# recovery owns the failure) keep the existing behavior.
|
||||
if (
|
||||
chunks_sent_to_client == 0
|
||||
and not fallbacks_disabled_for_request(initial_kwargs)
|
||||
and not initial_kwargs.get("fallbacks", self.fallbacks)
|
||||
):
|
||||
verbose_router_logger.info(
|
||||
"Anthropic messages stream dropped before first content; retrying the same group once"
|
||||
)
|
||||
retry_kwargs: Final = {
|
||||
**initial_kwargs,
|
||||
"original_function": self._ageneric_api_call_with_fallbacks_anthropic_messages_attempt,
|
||||
}
|
||||
retry_stream = await self._ageneric_api_call_with_fallbacks_anthropic_messages_attempt(
|
||||
**retry_kwargs
|
||||
)
|
||||
async for item in retry_stream:
|
||||
yield item
|
||||
return
|
||||
async for item in self._aanthropic_messages_recover_stream_error(
|
||||
stream_error,
|
||||
has_generated_content,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
Loading…
Add table
Reference in a new issue