Feat: add guardrail tracing to OTEL, Arize phoenix (#10896)

* feat: add guardrail tracing to OTEL, Arize phoenix

* fix: code qa check

* test: trace guard on OTEL

* fix: linting
This commit is contained in:
Ishaan Jaff 2025-05-16 13:38:11 -07:00 • committed by GitHub
parent 0942a9d51d
commit 1a7932a262
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 157 additions and 12 deletions

View file

@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.types.services import ServiceLoggerPayload
from litellm.types.utils import (
ChatCompletionMessageToolCall,
@ -16,6 +17,7 @@ from litellm.types.utils import (
if TYPE_CHECKING:
from opentelemetry.sdk.trace.export import SpanExporter as _SpanExporter
from opentelemetry.trace import Context as _Context
from opentelemetry.trace import Span as _Span
from litellm.proxy._types import (
@ -24,6 +26,7 @@ if TYPE_CHECKING:
from litellm.proxy.proxy_server import UserAPIKeyAuth as _UserAPIKeyAuth
Span = Union[_Span, Any]
Context = Union[_Context, Any]
SpanExporter = Union[_SpanExporter, Any]
UserAPIKeyAuth = Union[_UserAPIKeyAuth, Any]
ManagementEndpointLoggingPayload = Union[_ManagementEndpointLoggingPayload, Any]
@ -32,7 +35,7 @@ else:
SpanExporter = Any
UserAPIKeyAuth = Any
ManagementEndpointLoggingPayload = Any
Context = Any
LITELLM_TRACER_NAME = os.getenv("OTEL_TRACER_NAME", "litellm")
LITELLM_RESOURCE: Dict[Any, Any] = {
@ -63,9 +66,13 @@ class OpenTelemetryConfig:
InMemorySpanExporter,
)
exporter=os.getenv("OTEL_EXPORTER_OTLP_PROTOCOL", os.getenv("OTEL_EXPORTER", "console"))
endpoint=os.getenv("OTEL_EXPORTER_OTLP_ENDPOINT", os.getenv("OTEL_ENDPOINT"))
headers=os.getenv("OTEL_EXPORTER_OTLP_HEADERS", os.getenv("OTEL_HEADERS")) # example: OTEL_HEADERS=x-honeycomb-team=B85YgLm96***"
exporter = os.getenv(
"OTEL_EXPORTER_OTLP_PROTOCOL", os.getenv("OTEL_EXPORTER", "console")
)
endpoint = os.getenv("OTEL_EXPORTER_OTLP_ENDPOINT", os.getenv("OTEL_ENDPOINT"))
headers = os.getenv(
"OTEL_EXPORTER_OTLP_HEADERS", os.getenv("OTEL_HEADERS")
) # example: OTEL_HEADERS=x-honeycomb-team=B85YgLm96***"
if exporter == "in_memory":
return cls(exporter=InMemorySpanExporter())
@ -340,9 +347,72 @@ class OpenTelemetry(CustomLogger):
span.end(end_time=self._to_ns(end_time))
# Create span for guardrail information
self._create_guardrail_span(kwargs=kwargs, context=_parent_context)
if parent_otel_span is not None:
parent_otel_span.end(end_time=self._to_ns(datetime.now()))
def _create_guardrail_span(
self, kwargs: Optional[dict], context: Optional[Context]
):
"""
Creates a span for Guardrail, if any guardrail information is present in standard_logging_object
"""
# Create span for guardrail information
kwargs = kwargs or {}
standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object"
)
if standard_logging_payload is None:
return
guardrail_information = standard_logging_payload.get("guardrail_information")
if guardrail_information is None:
return
start_time_float = guardrail_information.get("start_time")
end_time_float = guardrail_information.get("end_time")
start_time_datetime = datetime.now()
if start_time_float is not None:
start_time_datetime = datetime.fromtimestamp(start_time_float)
end_time_datetime = datetime.now()
if end_time_float is not None:
end_time_datetime = datetime.fromtimestamp(end_time_float)
guardrail_span = self.tracer.start_span(
name="guardrail",
start_time=self._to_ns(start_time_datetime),
context=context,
)
self.safe_set_attribute(
span=guardrail_span,
key="guardrail_name",
value=guardrail_information.get("guardrail_name"),
)
self.safe_set_attribute(
span=guardrail_span,
key="guardrail_mode",
value=guardrail_information.get("guardrail_mode"),
)
# Set masked_entity_count directly without conversion
masked_entity_count = guardrail_information.get("masked_entity_count")
if masked_entity_count is not None:
guardrail_span.set_attribute(
"masked_entity_count", safe_dumps(masked_entity_count)
)
self.safe_set_attribute(
span=guardrail_span,
key="guardrail_response",
value=guardrail_information.get("guardrail_response"),
)
guardrail_span.end(end_time=self._to_ns(end_time_datetime))
def _add_dynamic_span_processor_if_needed(self, kwargs):
"""
Helper method to add a span processor with dynamic headers if needed.
@ -402,6 +472,9 @@ class OpenTelemetry(CustomLogger):
self.set_attributes(span, kwargs, response_obj)
span.end(end_time=self._to_ns(end_time))
# Create span for guardrail information
self._create_guardrail_span(kwargs=kwargs, context=_parent_context)
if parent_otel_span is not None:
parent_otel_span.end(end_time=self._to_ns(datetime.now()))
@ -856,7 +929,11 @@ class OpenTelemetry(CustomLogger):
self.OTEL_EXPORTER,
)
return BatchSpanProcessor(ConsoleSpanExporter())
elif self.OTEL_EXPORTER == "otlp_http" or self.OTEL_EXPORTER == "http/protobuf" or self.OTEL_EXPORTER == "http/json":
elif (
self.OTEL_EXPORTER == "otlp_http"
or self.OTEL_EXPORTER == "http/protobuf"
or self.OTEL_EXPORTER == "http/json"
):
verbose_logger.debug(
"OpenTelemetry: intiializing http exporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,

View file

@ -5,11 +5,11 @@ model_list:
api_key: any_key
api_base: https://exampleopenaiendpoint-production.up.railway.app/
litellm_settings:
callbacks:
- "langfuse"
- "arize_phoenix"
guardrails:
- guardrail_name: "bedrock-pre-guard"
litellm_params:
guardrail: bedrock # supported values: "aporia", "bedrock", "lakera"
mode: "during_call"
guardrailIdentifier: ff6ujrregl1q
guardrailVersion: "DRAFT"
general_settings:
store_prompts_in_spend_logs: true

View file

@ -0,0 +1,68 @@
import os
import sys
import unittest
from unittest.mock import MagicMock, patch
# Adds the grandparent directory to sys.path to allow importing project modules
sys.path.insert(0, os.path.abspath("../.."))
from litellm.integrations.opentelemetry import OpenTelemetry
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
class TestOpenTelemetry(unittest.TestCase):
@patch("litellm.integrations.opentelemetry.datetime")
def test_create_guardrail_span_with_valid_info(self, mock_datetime):
# Setup
otel = OpenTelemetry()
otel.tracer = MagicMock()
mock_span = MagicMock()
otel.tracer.start_span.return_value = mock_span
# Create guardrail information
guardrail_info = {
"guardrail_name": "test_guardrail",
"guardrail_mode": "input",
"masked_entity_count": {"CREDIT_CARD": 2},
"guardrail_response": "filtered_content",
"start_time": 1609459200.0,
"end_time": 1609459201.0,
}
# Create a kwargs dict with standard_logging_object containing guardrail information
kwargs = {"standard_logging_object": {"guardrail_information": guardrail_info}}
# Call the method
otel._create_guardrail_span(kwargs=kwargs, context=None)
# Assertions
otel.tracer.start_span.assert_called_once()
# print all calls to mock_span.set_attribute
print("Calls to mock_span.set_attribute:")
for call in mock_span.set_attribute.call_args_list:
print(call)
# Check that the span has the correct attributes set
mock_span.set_attribute.assert_any_call("guardrail_name", "test_guardrail")
mock_span.set_attribute.assert_any_call("guardrail_mode", "input")
mock_span.set_attribute.assert_any_call(
"guardrail_response", "filtered_content"
)
mock_span.set_attribute.assert_any_call(
"masked_entity_count", safe_dumps({"CREDIT_CARD": 2})
)
# Verify that the span was ended
mock_span.end.assert_called_once()
def test_create_guardrail_span_with_no_info(self):
# Setup
otel = OpenTelemetry()
otel.tracer = MagicMock()
# Test with no guardrail information
kwargs = {"standard_logging_object": {}}
otel._create_guardrail_span(kwargs=kwargs, context=None)
# Verify that start_span was never called
otel.tracer.start_span.assert_not_called()