This commit is contained in:
mubashir1osmani 2026-04-27 21:54:43 -04:00
parent d778be1f24
commit caca527cf8
No known key found for this signature in database
GPG key ID: AB055FF67D0B4D9A
2 changed files with 100 additions and 0 deletions

View file

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

View file

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