fix arize obsv bugs

This commit is contained in:
mubashir1osmani 2026-04-25 17:49:34 -04:00
parent e9e86ed956
commit 07e42bc871
No known key found for this signature in database
GPG key ID: AB055FF67D0B4D9A
2 changed files with 558 additions and 45 deletions

View file

@ -1,6 +1,9 @@
import json
import hashlib
import os
from typing import TYPE_CHECKING, Any, Dict, Optional, Type
from pydantic import BaseModel
from typing_extensions import override
from litellm._logging import verbose_logger
@ -14,15 +17,19 @@ from litellm.types.utils import StandardLoggingPayload
if TYPE_CHECKING:
from opentelemetry.trace import Span
from litellm.integrations._types.open_inference import (
MessageAttributes,
ImageAttributes,
SpanAttributes,
AudioAttributes,
EmbeddingAttributes,
ImageAttributes,
MessageAttributes,
MessageContentAttributes,
OpenInferenceSpanKindValues,
SpanAttributes,
)
DEFAULT_MAX_INLINE_IMAGE_BYTES = 32 * 1024
class ArizeOTELAttributes(BaseLLMObsOTELAttributes):
@staticmethod
@override
@ -92,17 +99,113 @@ class ArizeOTELAttributes(BaseLLMObsOTELAttributes):
def _set_response_attributes(span: "Span", response_obj):
"""Helper to set response output and token usage attributes on span."""
# Pydantic responses (ImageResponse, ResponsesAPIResponse, ...) need to be
# dict-coerced so the dict-keyed setters below see the data.
if isinstance(response_obj, BaseModel):
try:
response_obj = response_obj.model_dump()
except Exception:
return
if not hasattr(response_obj, "get"):
return
_set_choice_outputs(span, response_obj, MessageAttributes, SpanAttributes)
_set_image_outputs(span, response_obj, ImageAttributes, SpanAttributes)
_set_audio_outputs(span, response_obj, AudioAttributes, SpanAttributes)
_set_embedding_outputs(span, response_obj, EmbeddingAttributes, SpanAttributes)
_set_image_outputs(
span,
response_obj,
MessageAttributes,
MessageContentAttributes,
ImageAttributes,
SpanAttributes,
)
_set_structured_outputs(span, response_obj, MessageAttributes, SpanAttributes)
_set_usage_outputs(span, response_obj, SpanAttributes)
def _set_image_outputs(
span: "Span",
response_obj,
msg_attrs,
content_attrs,
image_attrs,
span_attrs,
):
"""
Render generated images on the LLM span.
Phoenix renders an image inline when the assistant message has an
``image``-type content entry whose ``image.url`` is either a public URL or
a ``data:image/<type>;base64,<...>`` URI. Mirrors the structure produced
by ``openinference.instrumentation.ImageMessageContent``.
Iterates ``response_obj["data"]`` (image-gen / image-edit shape). Skips
items that lack both ``url`` and ``b64_json`` so embedding payloads (which
also live under ``data``) aren't matched.
"""
data = response_obj.get("data")
if not isinstance(data, list) or not data:
return
role_set = False
content_idx = 0
for image_item in data:
if isinstance(image_item, BaseModel):
image_item = image_item.model_dump()
if not hasattr(image_item, "get"):
continue
if not image_item.get("url") and not image_item.get("b64_json"):
continue # not an image item (could be embedding)
image_url, mime, omission_notice = _get_image_trace_payload(
image_item, response_obj
)
if not image_url and not omission_notice:
continue
if not role_set:
safe_set_attribute(
span,
f"{span_attrs.LLM_OUTPUT_MESSAGES}.0.{msg_attrs.MESSAGE_ROLE}",
"assistant",
)
# First image also drives the span's top-level Output preview.
if image_url:
safe_set_attribute(span, span_attrs.OUTPUT_VALUE, image_url)
safe_set_attribute(span, span_attrs.OUTPUT_MIME_TYPE, mime)
else:
safe_set_attribute(span, span_attrs.OUTPUT_VALUE, omission_notice)
safe_set_attribute(span, span_attrs.OUTPUT_MIME_TYPE, "text/plain")
role_set = True
prefix = (
f"{span_attrs.LLM_OUTPUT_MESSAGES}.0."
f"{msg_attrs.MESSAGE_CONTENTS}.{content_idx}"
)
if image_url:
safe_set_attribute(
span, f"{prefix}.{content_attrs.MESSAGE_CONTENT_TYPE}", "image"
)
safe_set_attribute(
span,
f"{prefix}.{content_attrs.MESSAGE_CONTENT_IMAGE}.{image_attrs.IMAGE_URL}",
image_url,
)
elif omission_notice:
safe_set_attribute(
span, f"{prefix}.{content_attrs.MESSAGE_CONTENT_TYPE}", "text"
)
safe_set_attribute(
span,
f"{prefix}.{content_attrs.MESSAGE_CONTENT_TEXT}",
omission_notice,
)
content_idx += 1
def _set_choice_outputs(span: "Span", response_obj, msg_attrs, span_attrs):
for idx, choice in enumerate(response_obj.get("choices", [])):
response_message = choice.get("message", {})
@ -124,22 +227,6 @@ def _set_choice_outputs(span: "Span", response_obj, msg_attrs, span_attrs):
)
def _set_image_outputs(span: "Span", response_obj, image_attrs, span_attrs):
images = response_obj.get("data", [])
for i, image in enumerate(images):
img_url = image.get("url")
if img_url is None and image.get("b64_json"):
img_url = f"data:image/png;base64,{image.get('b64_json')}"
if not img_url:
continue
if i == 0:
safe_set_attribute(span, span_attrs.OUTPUT_VALUE, img_url)
safe_set_attribute(span, f"{image_attrs.IMAGE_URL}.{i}", img_url)
def _set_audio_outputs(span: "Span", response_obj, audio_attrs, span_attrs):
audio = response_obj.get("audio", [])
for i, audio_item in enumerate(audio):
@ -194,25 +281,37 @@ def _set_structured_outputs(span: "Span", response_obj, msg_attrs, span_attrs):
output_items = response_obj.get("output", [])
for i, item in enumerate(output_items):
prefix = f"{span_attrs.LLM_OUTPUT_MESSAGES}.{i}"
if not hasattr(item, "type"):
# Items can arrive as Pydantic models (SDK path) or dicts (proxy path).
# Read both via a uniform accessor so dict-shaped output[] arrays from
# Responses API don't get silently skipped.
def _get(obj, key, default=None):
if hasattr(obj, "get"):
return obj.get(key, default)
return getattr(obj, key, default)
item_type = _get(item, "type")
if item_type is None:
continue
item_type = item.type
if item_type == "reasoning" and hasattr(item, "summary"):
for summary in item.summary:
if hasattr(summary, "text"):
safe_set_attribute(
span,
f"{prefix}.{msg_attrs.MESSAGE_REASONING_SUMMARY}",
summary.text,
)
elif item_type == "message" and hasattr(item, "content"):
if item_type == "reasoning":
summary = _get(item, "summary")
if isinstance(summary, list):
for s in summary:
text = _get(s, "text")
if text:
safe_set_attribute(
span,
f"{prefix}.{msg_attrs.MESSAGE_REASONING_SUMMARY}",
text,
)
elif item_type == "message":
content_list = _get(item, "content") or []
message_content = ""
content_list = item.content
if content_list and len(content_list) > 0:
if content_list:
first_content = content_list[0]
message_content = getattr(first_content, "text", "")
message_role = getattr(item, "role", "assistant")
message_content = _get(first_content, "text", "") or ""
message_role = _get(item, "role", "assistant") or "assistant"
safe_set_attribute(span, span_attrs.OUTPUT_VALUE, message_content)
safe_set_attribute(
span, f"{prefix}.{msg_attrs.MESSAGE_CONTENT}", message_content
@ -419,6 +518,388 @@ def set_attributes(
span.record_exception(e)
def _resolve_image_mime_type(image, response_obj) -> str:
"""
Pick the MIME type for an image in a generation/edit response.
Phoenix's image renderer keys off the data-URI mime, so a wrong value
(e.g. ``image/png`` for a JPEG) breaks inline preview. Resolve in order:
per-image ``mime_type`` → per-image ``output_format`` → response
``output_format`` → ``image/png``.
"""
mime: Optional[str] = None
if hasattr(image, "get"):
mime = image.get("mime_type") or image.get("output_format")
if not mime and hasattr(response_obj, "get"):
mime = response_obj.get("output_format")
if not mime:
return "image/png"
mime = str(mime).lower()
if mime == "jpg":
mime = "jpeg"
if "/" not in mime:
mime = f"image/{mime}"
return mime
def _estimate_b64_decoded_bytes(b64_payload: str) -> int:
"""Estimate decoded byte size from base64 payload without decoding."""
payload = b64_payload.strip()
padding = payload.count("=")
return max(0, (len(payload) * 3) // 4 - padding)
def _get_max_inline_image_bytes() -> Optional[int]:
"""
Resolve max inline bytes for image payloads.
- Default is 32KB to keep Phoenix spans exportable.
- Set LITELLM_ARIZE_MAX_INLINE_IMAGE_BYTES<=0 to disable the cap.
"""
raw_value = os.getenv("LITELLM_ARIZE_MAX_INLINE_IMAGE_BYTES")
if raw_value is None or raw_value == "":
return DEFAULT_MAX_INLINE_IMAGE_BYTES
try:
parsed = int(raw_value)
if parsed <= 0:
return None
return parsed
except (TypeError, ValueError):
return DEFAULT_MAX_INLINE_IMAGE_BYTES
def _get_image_trace_payload(
image_item, response_obj
) -> tuple[Optional[str], str, Optional[str]]:
"""
Build image payload tuple:
(image_url_or_data_uri, mime_type, omission_notice).
"""
url = image_item.get("url")
mime = _resolve_image_mime_type(image_item, response_obj)
if url:
return url, mime, None
b64 = image_item.get("b64_json")
if not b64:
return None, mime, None
max_inline_bytes = _get_max_inline_image_bytes()
if max_inline_bytes is not None:
decoded_bytes = _estimate_b64_decoded_bytes(str(b64))
if decoded_bytes > max_inline_bytes:
digest = hashlib.sha256(str(b64).encode("utf-8")).hexdigest()[:12]
omission_notice = (
f"[image omitted from trace: {decoded_bytes} bytes exceeds "
f"{max_inline_bytes} byte inline limit, sha256={digest}]"
)
return None, mime, omission_notice
return f"data:{mime};base64,{b64}", mime, None
def _extract_responses_api_text(output_items) -> Optional[str]:
"""Concatenate ``output[*].content[*].text`` for Responses API message items."""
if not isinstance(output_items, list) or not output_items:
return None
def _g(obj, key, default=None):
if hasattr(obj, "get"):
return obj.get(key, default)
return getattr(obj, key, default)
texts = []
for item in output_items:
if _g(item, "type") != "message":
continue
content_list = _g(item, "content")
if not isinstance(content_list, list):
continue
for c in content_list:
text = _g(c, "text")
if text:
texts.append(text)
return "\n\n".join(texts) if texts else None
def _coerce_text(value: Any) -> Optional[str]:
"""Reduce a heterogeneous prompt/response value to a renderable string."""
if value is None:
return None
if isinstance(value, str):
return value
if isinstance(value, (list, tuple)):
# Embedding `input` is often List[str]; show the first entry plus a
# count rather than serialising every vector input.
text_items = [item for item in value if isinstance(item, str)]
if text_items:
head = text_items[0]
if len(text_items) == 1:
return head
return f"{head}\n\n[+ {len(text_items) - 1} more]"
return safe_dumps(value)
if isinstance(value, dict):
return safe_dumps(value)
return str(value)
def _extract_chain_input(
kwargs, standard_logging_payload: Optional[StandardLoggingPayload]
) -> tuple[Optional[str], Optional[str]]:
"""
- LLM chat history → last user message content
- LLM completion / image gen → ``prompt``
- Embedding → ``input`` (string or list of strings)
- Reranker / Retriever → ``query`` (+ document count)
- Tool / MCP → ``{name, arguments}`` JSON
- Agent (a2a / responses) → ``input`` or last message
- Guardrail → ``text``
- Anything else → ``standard_logging_payload["messages"]`` as JSON
"""
# 1. Tool calls — distinct shape, render as JSON arguments
tool_name = kwargs.get("name") or kwargs.get("tool_name")
tool_args = kwargs.get("arguments") or kwargs.get("tool_arguments")
if tool_name and tool_args is not None:
return (
safe_dumps({"name": tool_name, "arguments": tool_args}),
"application/json",
)
# 2. Chat-style messages (LLM, agents that use messages)
messages = kwargs.get("messages")
if isinstance(messages, list) and messages:
last = messages[-1]
if isinstance(last, dict):
content = last.get("content")
text = _coerce_text(content)
if text:
return text, None
# 3. Reranker / Retriever — query plus optional document count summary
query = kwargs.get("query")
if query:
text = _coerce_text(query)
documents = kwargs.get("documents")
if isinstance(documents, list) and documents:
text = f"{text}\n\n[+ {len(documents)} documents]"
return text, None
# 4. Single-text inputs (embedding/image/completion/guardrail)
for field in ("input", "prompt", "text"):
text = _coerce_text(kwargs.get(field))
if text:
return text, None
# 5. Universal fallback — every call type populates this
if isinstance(standard_logging_payload, dict):
msgs = standard_logging_payload.get("messages")
text = _coerce_text(msgs)
if text:
mime = "application/json" if isinstance(msgs, (list, dict)) else None
return text, mime
return None, None
def _extract_chain_output(
response_obj, standard_logging_payload: Optional[StandardLoggingPayload] = None
) -> tuple[Optional[str], Optional[str]]:
"""
Build ``(output.value, output.mime_type)`` for the parent CHAIN span.
Tries the rich response_obj shapes first, then falls back to
``standard_logging_payload["response"]`` so TOOL / AGENT / RETRIEVER /
RERANKER / GUARDRAIL spans still get a meaningful output.
- Chat → first choice message content (or tool_calls JSON if no content)
- Image gen → first image URL or ``data:`` URI (matches Phoenix's renderer)
- Embedding → ``"<n> embeddings"`` summary
- Tool result → ``content`` text or full result JSON
- Reranker → ``results`` JSON
- Anything else → ``standard_logging_payload["response"]``
"""
if isinstance(response_obj, BaseModel):
try:
response_obj = response_obj.model_dump()
except Exception:
pass
response_obj_dict = response_obj if isinstance(response_obj, dict) else None
if response_obj_dict is not None:
# 1a. Responses API — output[*].content[*].text on message items.
# Matches what _set_structured_outputs writes on the LLM child span,
# so the parent CHAIN renders the same readable assistant text.
responses_text = _extract_responses_api_text(response_obj_dict.get("output"))
if responses_text:
return responses_text, None
# 1. Chat completion choices
choices = response_obj_dict.get("choices") or []
if isinstance(choices, list) and choices:
first = choices[0]
message = first.get("message") if hasattr(first, "get") else None
if isinstance(message, dict):
content = _coerce_text(message.get("content"))
if content:
return content, None
# Tool-calling assistant turn — no content but tool_calls present
tool_calls = message.get("tool_calls")
if tool_calls:
return safe_dumps(tool_calls), "application/json"
# 2. Image / Embedding ``data`` array
data = response_obj_dict.get("data")
if isinstance(data, list) and data:
first = data[0]
if isinstance(first, BaseModel):
first = first.model_dump()
if isinstance(first, dict):
url = first.get("url")
if url:
return url, _resolve_image_mime_type(first, response_obj_dict)
b64 = first.get("b64_json")
if b64:
image_url, mime, omission_notice = _get_image_trace_payload(
first, response_obj_dict
)
if image_url:
return image_url, mime
if omission_notice:
return omission_notice, "text/plain"
if first.get("embedding") is not None:
return f"{len(data)} embeddings", "text/plain"
# 3. Tool / MCP result shapes — content array, output_text, result
for field in ("output_text", "result"):
text = _coerce_text(response_obj_dict.get(field))
if text:
mime = (
"application/json"
if isinstance(response_obj_dict.get(field), (list, dict))
else None
)
return text, mime
content = response_obj_dict.get("content")
if content:
return _coerce_text(content), (
"application/json" if isinstance(content, (list, dict)) else None
)
# 4. Reranker results
results = response_obj_dict.get("results")
if results:
return safe_dumps(results), "application/json"
# 5. Universal fallback for any call type
if isinstance(standard_logging_payload, dict):
resp = standard_logging_payload.get("response")
text = _coerce_text(resp)
if text:
mime = "application/json" if isinstance(resp, (list, dict)) else None
return text, mime
return None, None
def set_parent_span_attributes(span: "Span", kwargs, response_obj):
"""
set parent level span attr (model, call type) to prevent token / cost duplication from child spans
"""
try:
litellm_params = kwargs.get("litellm_params", {}) or {}
standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object"
)
safe_set_attribute(
span,
SpanAttributes.OPENINFERENCE_SPAN_KIND,
OpenInferenceSpanKindValues.CHAIN.value,
)
if kwargs.get("model"):
safe_set_attribute(span, SpanAttributes.LLM_MODEL_NAME, kwargs.get("model"))
safe_set_attribute(
span,
SpanAttributes.LLM_PROVIDER,
litellm_params.get("custom_llm_provider", "Unknown"),
)
if standard_logging_payload is not None:
call_type = standard_logging_payload.get("call_type")
if call_type:
safe_set_attribute(span, "llm.request.type", call_type)
metadata = standard_logging_payload.get("metadata")
_set_metadata_attributes(span, metadata, SpanAttributes)
model_params = standard_logging_payload.get("model_parameters") or {}
user_id = (
model_params.get("user") if isinstance(model_params, dict) else None
)
if user_id is not None:
safe_set_attribute(span, SpanAttributes.USER_ID, user_id)
optional_params = _sanitize_optional_params(kwargs.get("optional_params"))
safe_set_attribute(
span, "llm.is_streaming", str(optional_params.get("stream", False))
)
# pass provider / litellm call id
response_id = _resolve_response_id(response_obj, standard_logging_payload)
if response_id is not None:
safe_set_attribute(span, "llm.response.id", response_id)
# sanitize input / output for parent span
input_value, input_mime = _extract_chain_input(kwargs, standard_logging_payload)
if input_value:
safe_set_attribute(span, SpanAttributes.INPUT_VALUE, input_value)
if input_mime:
safe_set_attribute(span, SpanAttributes.INPUT_MIME_TYPE, input_mime)
output_value, output_mime = _extract_chain_output(
response_obj, standard_logging_payload
)
if output_value:
safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, output_value)
if output_mime:
safe_set_attribute(span, SpanAttributes.OUTPUT_MIME_TYPE, output_mime)
except Exception as e:
verbose_logger.error(
f"[Arize/Phoenix] Failed to set parent span attributes: {e}"
)
if hasattr(span, "record_exception"):
span.record_exception(e)
def _resolve_response_id(
response_obj, standard_logging_payload: Optional[StandardLoggingPayload]
) -> Optional[str]:
"""
for completions / responses endpoint, pass the response id issued by provider.
(chatcmpl-*)
other endpoints (embeddings, images) pass the standard x-litellm-call-id
to match litellm UI logs with phoenix
"""
if isinstance(response_obj, dict):
provider_id = response_obj.get("id")
if provider_id:
return str(provider_id)
if isinstance(standard_logging_payload, dict):
litellm_call_id = standard_logging_payload.get("litellm_call_id")
if litellm_call_id:
return str(litellm_call_id)
return None
def _sanitize_optional_params(optional_params: Optional[dict]) -> dict:
if not isinstance(optional_params, dict):
return {}
@ -483,8 +964,9 @@ def _set_request_attributes(
if optional_params.get("user"):
safe_set_attribute(span, "llm.user", optional_params.get("user"))
if response_obj and response_obj.get("id"):
safe_set_attribute(span, "llm.response.id", response_obj.get("id"))
response_id = _resolve_response_id(response_obj, standard_logging_payload)
if response_id is not None:
safe_set_attribute(span, "llm.response.id", response_id)
if response_obj and response_obj.get("model"):
safe_set_attribute(span, "llm.response.model", response_obj.get("model"))

View file

@ -5,6 +5,7 @@ from litellm._logging import verbose_logger
from litellm.integrations.arize import _utils
from litellm.integrations.arize._utils import ArizeOTELAttributes
from litellm.types.integrations.arize_phoenix import ArizePhoenixConfig
from opentelemetry import trace as _trace
if TYPE_CHECKING:
from opentelemetry.sdk.trace import TracerProvider
@ -87,6 +88,18 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
"""
pass
async def async_post_call_success_hook(
self,
data,
user_api_key_dict,
response,
):
"""
skipping this hook prevents both the orphan and the duplicate guardrail spans.
already handled in global tracer provider
"""
return response
def set_attributes(self, span: Span, kwargs, response_obj: Optional[Any]):
ArizePhoenixLogger.set_arize_phoenix_attributes(span, kwargs, response_obj)
return
@ -221,15 +234,26 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
# Raw-request sub-span (if enabled) — must be created before
# ending the parent span so the hierarchy is valid.
self._maybe_log_raw_request(kwargs, response_obj, start_time, end_time, span)
# Guardrail context: in proxy mode it's a sibling of litellm_request
# under litellm_proxy_request. In SDK mode there is no proxy parent,
# so we parent it to the litellm_request span to avoid an orphan.
guardrail_ctx = (
ctx if parent_span is not None else _trace.set_span_in_context(span)
)
span.end(end_time=self._to_ns(end_time))
# Guardrail span
self._create_guardrail_span(kwargs=kwargs, context=ctx)
# Guardrail span (always parented — no orphan roots)
self._create_guardrail_span(kwargs=kwargs, context=guardrail_ctx)
# Annotate and close our proxy parent span
# Annotate and close our proxy parent span.
# Only session/request metadata goes on the parent — the child span
# already carries the full LLM payload (messages, tokens, response).
# Duplicating them here double-counts tokens.
if parent_span is not None:
parent_span.set_status(Status(StatusCode.OK))
self.set_attributes(parent_span, kwargs, response_obj)
_utils.set_parent_span_attributes(parent_span, kwargs, response_obj)
parent_span.end(end_time=self._to_ns(end_time))
# Metrics & cost recording
@ -263,15 +287,22 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
span.set_status(Status(StatusCode.ERROR))
self.set_attributes(span, kwargs, response_obj)
self._record_exception_on_span(span=span, kwargs=kwargs)
# See _handle_success for guardrail context rationale.
guardrail_ctx = (
ctx if parent_span is not None else _trace.set_span_in_context(span)
)
span.end(end_time=self._to_ns(end_time))
# Guardrail span
self._create_guardrail_span(kwargs=kwargs, context=ctx)
# Guardrail span (always parented, no orphan roots)
self._create_guardrail_span(kwargs=kwargs, context=guardrail_ctx)
# Annotate and close our proxy parent span
# Annotate and close our proxy parent span (see _handle_success for
# why only parent-scoped attributes go here).
if parent_span is not None:
parent_span.set_status(Status(StatusCode.ERROR))
self.set_attributes(parent_span, kwargs, response_obj)
_utils.set_parent_span_attributes(parent_span, kwargs, response_obj)
self._record_exception_on_span(span=parent_span, kwargs=kwargs)
parent_span.end(end_time=self._to_ns(end_time))