diff --git a/litellm/integrations/arize/_utils.py b/litellm/integrations/arize/_utils.py index 3ad7dda0796..ca9d432b4da 100644 --- a/litellm/integrations/arize/_utils.py +++ b/litellm/integrations/arize/_utils.py @@ -938,6 +938,24 @@ def _set_request_attributes( if kwargs.get("model"): safe_set_attribute(span, span_attrs.LLM_MODEL_NAME, kwargs.get("model")) + # Track fallback context: which attempt number in the fallback chain + fallback_attempt = kwargs.get("fallback_depth") + if fallback_attempt is not None: + safe_set_attribute(span, "llm.fallback.attempt_number", fallback_attempt) + # Original model that was attempted before fallback + original_model = kwargs.get("original_model") + if original_model: + safe_set_attribute(span, "llm.fallback.original_model", original_model) + # Error that triggered fallback + original_exception = kwargs.get("original_exception") + if original_exception is not None: + status_code = getattr(original_exception, "status_code", None) + if status_code is not None: + safe_set_attribute(span, "llm.fallback.error_status_code", status_code) + exception_class = getattr(original_exception, "__class__.__name__", None) + if exception_class: + safe_set_attribute(span, "llm.fallback.error_class", exception_class) + safe_set_attribute( span, "llm.request.type", standard_logging_payload.get("call_type") ) diff --git a/litellm/integrations/arize/arize_phoenix.py b/litellm/integrations/arize/arize_phoenix.py index 563d1b493c1..7e9b8a4fdd0 100644 --- a/litellm/integrations/arize/arize_phoenix.py +++ b/litellm/integrations/arize/arize_phoenix.py @@ -392,3 +392,85 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore "status": "healthy", "message": "Arize-Phoenix credentials are configured properly", } + + async def log_success_fallback_event( + self, + original_model_group: str, + kwargs: dict, + original_exception: Exception, + ): + """ + Log a successful fallback event by creating a span that captures the fallback chain. + + When a fallback succeeds, this creates a span showing: + - Which model originally failed (original_model_group) + - Which model it fell back to (from kwargs["model"]) + - The error that triggered the fallback + """ + from opentelemetry.trace import Status, StatusCode + + original_model = original_model_group + fallback_model = kwargs.get("model") + fallback_depth = kwargs.get("fallback_depth", 1) + + status_code = getattr(original_exception, "status_code", None) + exception_class = original_exception.__class__.__name__ + + span_name = f"fallback: {original_model} -> {fallback_model}" + span = self.tracer.start_span(name=span_name) + + def _safe_set(s, k, v): + if hasattr(s, "set_attribute"): + s.set_attribute(k, v) + + _safe_set(span, "llm.fallback.event_type", "success") + _safe_set(span, "llm.fallback.original_model", original_model) + _safe_set(span, "llm.fallback.fallback_model", fallback_model) + _safe_set(span, "llm.fallback.attempt_number", fallback_depth) + if status_code is not None: + _safe_set(span, "llm.fallback.error_status_code", status_code) + _safe_set(span, "llm.fallback.error_class", exception_class) + + span.set_status(Status(StatusCode.OK)) + span.end() + + async def log_failure_fallback_event( + self, + original_model_group: str, + kwargs: dict, + original_exception: Exception, + ): + """ + Log a failed fallback event by creating a span that captures the failure. + + When all fallbacks fail, this creates a span showing: + - Which model originally failed + - Which fallback models were attempted + - The chain of errors + """ + from opentelemetry.trace import Status, StatusCode + + original_model = original_model_group + fallback_attempted = kwargs.get("model") + max_fallbacks = kwargs.get("max_fallbacks", 0) + + exception_class = original_exception.__class__.__name__ + status_code = getattr(original_exception, "status_code", None) + + span_name = f"fallback_failed: {original_model}" + span = self.tracer.start_span(name=span_name) + + def _safe_set(s, k, v): + if hasattr(s, "set_attribute"): + s.set_attribute(k, v) + + _safe_set(span, "llm.fallback.event_type", "failure") + _safe_set(span, "llm.fallback.original_model", original_model) + _safe_set(span, "llm.fallback.attempted_model", fallback_attempted) + _safe_set(span, "llm.fallback.max_fallbacks", max_fallbacks) + if status_code is not None: + _safe_set(span, "llm.fallback.error_status_code", status_code) + _safe_set(span, "llm.fallback.error_class", exception_class) + + span.set_status(Status(StatusCode.ERROR)) + span.end()