Merge pull request #9353 from BerriAI/litellm_arize_dynamic_logging

[Feat] - API - Allow using dynamic Arize AI Spaces on LiteLLM
This commit is contained in:
Ishaan Jaff 2025-03-18 23:35:28 -07:00 • committed by GitHub
commit e32aee9124
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 222 additions and 51 deletions

View file

@ -3,31 +3,38 @@ arize AI is OTEL compatible
this file has Arize ai specific helper functions
"""
import os
from typing import TYPE_CHECKING, Any
import os
from datetime import datetime
from typing import TYPE_CHECKING, Any, Optional, Union
from litellm.integrations.arize import _utils
from litellm.integrations.opentelemetry import OpenTelemetry
from litellm.types.integrations.arize import ArizeConfig
from litellm.types.services import ServiceLoggerPayload
if TYPE_CHECKING:
from litellm.types.integrations.arize import Protocol as _Protocol
from opentelemetry.trace import Span as _Span
from litellm.types.integrations.arize import Protocol as _Protocol
Protocol = _Protocol
Span = _Span
else:
Protocol = Any
Span = Any
class ArizeLogger:
class ArizeLogger(OpenTelemetry):
def set_attributes(self, span: Span, kwargs, response_obj: Optional[Any]):
ArizeLogger.set_arize_attributes(span, kwargs, response_obj)
return
@staticmethod
def set_arize_attributes(span: Span, kwargs, response_obj):
_utils.set_attributes(span, kwargs, response_obj)
return
@staticmethod
def get_arize_config() -> ArizeConfig:
@ -43,11 +50,6 @@ class ArizeLogger:
space_key = os.environ.get("ARIZE_SPACE_KEY")
api_key = os.environ.get("ARIZE_API_KEY")
if not space_key:
raise ValueError("ARIZE_SPACE_KEY not found in environment variables")
if not api_key:
raise ValueError("ARIZE_API_KEY not found in environment variables")
grpc_endpoint = os.environ.get("ARIZE_ENDPOINT")
http_endpoint = os.environ.get("ARIZE_HTTP_ENDPOINT")
@ -55,13 +57,13 @@ class ArizeLogger:
protocol: Protocol = "otlp_grpc"
if grpc_endpoint:
protocol="otlp_grpc"
endpoint=grpc_endpoint
protocol = "otlp_grpc"
endpoint = grpc_endpoint
elif http_endpoint:
protocol="otlp_http"
endpoint=http_endpoint
protocol = "otlp_http"
endpoint = http_endpoint
else:
protocol="otlp_grpc"
protocol = "otlp_grpc"
endpoint = "https://otlp.arize.com/v1"
return ArizeConfig(
@ -71,4 +73,33 @@ class ArizeLogger:
endpoint=endpoint,
)
async def async_service_success_hook(
self,
payload: ServiceLoggerPayload,
parent_otel_span: Optional[Span] = None,
start_time: Optional[Union[datetime, float]] = None,
end_time: Optional[Union[datetime, float]] = None,
event_metadata: Optional[dict] = None,
):
"""Arize is used mainly for LLM I/O tracing, sending router+caching metrics adds bloat to arize logs"""
pass
async def async_service_failure_hook(
self,
payload: ServiceLoggerPayload,
error: Optional[str] = "",
parent_otel_span: Optional[Span] = None,
start_time: Optional[Union[datetime, float]] = None,
end_time: Optional[Union[float, datetime]] = None,
event_metadata: Optional[dict] = None,
):
"""Arize is used mainly for LLM I/O tracing, sending router+caching metrics adds bloat to arize logs"""
pass
def create_litellm_proxy_request_started_span(
self,
start_time: datetime,
headers: dict,
):
"""Arize is used mainly for LLM I/O tracing, sending Proxy Server Request adds bloat to arize logs"""
pass

View file

@ -10,6 +10,7 @@ from litellm.types.services import ServiceLoggerPayload
from litellm.types.utils import (
ChatCompletionMessageToolCall,
Function,
StandardCallbackDynamicParams,
StandardLoggingPayload,
)
@ -311,6 +312,8 @@ class OpenTelemetry(CustomLogger):
)
_parent_context, parent_otel_span = self._get_span_context(kwargs)
self._add_dynamic_span_processor_if_needed(kwargs)
# Span 1: Requst sent to litellm SDK
span = self.tracer.start_span(
name=self._get_span_name(kwargs),
@ -341,6 +344,45 @@ class OpenTelemetry(CustomLogger):
if parent_otel_span is not None:
parent_otel_span.end(end_time=self._to_ns(datetime.now()))
def _add_dynamic_span_processor_if_needed(self, kwargs):
"""
Helper method to add a span processor with dynamic headers if needed.
This allows for per-request configuration of telemetry exporters by
extracting headers from standard_callback_dynamic_params.
"""
from opentelemetry import trace
standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = (
kwargs.get("standard_callback_dynamic_params")
)
if not standard_callback_dynamic_params:
return
# Extract headers from dynamic params
dynamic_headers = {}
# Handle Arize headers
if standard_callback_dynamic_params.get("arize_space_key"):
dynamic_headers["space_key"] = standard_callback_dynamic_params.get(
"arize_space_key"
)
if standard_callback_dynamic_params.get("arize_api_key"):
dynamic_headers["api_key"] = standard_callback_dynamic_params.get(
"arize_api_key"
)
# Only create a span processor if we have headers to use
if len(dynamic_headers) > 0:
from opentelemetry.sdk.trace import TracerProvider
provider = trace.get_tracer_provider()
if isinstance(provider, TracerProvider):
span_processor = self._get_span_processor(
dynamic_headers=dynamic_headers
)
provider.add_span_processor(span_processor)
def _handle_failure(self, kwargs, response_obj, start_time, end_time):
from opentelemetry.trace import Status, StatusCode
@ -443,14 +485,12 @@ class OpenTelemetry(CustomLogger):
self, span: Span, kwargs, response_obj: Optional[Any]
):
try:
if self.callback_name == "arize":
from litellm.integrations.arize.arize import ArizeLogger
ArizeLogger.set_arize_attributes(span, kwargs, response_obj)
return
elif self.callback_name == "arize_phoenix":
if self.callback_name == "arize_phoenix":
from litellm.integrations.arize.arize_phoenix import ArizePhoenixLogger
ArizePhoenixLogger.set_arize_phoenix_attributes(span, kwargs, response_obj)
ArizePhoenixLogger.set_arize_phoenix_attributes(
span, kwargs, response_obj
)
return
elif self.callback_name == "langtrace":
from litellm.integrations.langtrace import LangtraceAttributes
@ -779,7 +819,7 @@ class OpenTelemetry(CustomLogger):
carrier = {"traceparent": traceparent}
return TraceContextTextMapPropagator().extract(carrier=carrier), None
def _get_span_processor(self):
def _get_span_processor(self, dynamic_headers: Optional[dict] = None):
from opentelemetry.exporter.otlp.proto.grpc.trace_exporter import (
OTLPSpanExporter as OTLPSpanExporterGRPC,
)
@ -799,10 +839,9 @@ class OpenTelemetry(CustomLogger):
self.OTEL_ENDPOINT,
self.OTEL_HEADERS,
)
_split_otel_headers = {}
if self.OTEL_HEADERS is not None and isinstance(self.OTEL_HEADERS, str):
_split_otel_headers = self.OTEL_HEADERS.split("=")
_split_otel_headers = {_split_otel_headers[0]: _split_otel_headers[1]}
_split_otel_headers = OpenTelemetry._get_headers_dictionary(
headers=dynamic_headers or self.OTEL_HEADERS
)
if isinstance(self.OTEL_EXPORTER, SpanExporter):
verbose_logger.debug(
@ -844,6 +883,25 @@ class OpenTelemetry(CustomLogger):
)
return BatchSpanProcessor(ConsoleSpanExporter())
@staticmethod
def _get_headers_dictionary(headers: Optional[Union[str, dict]]) -> Dict[str, str]:
"""
Convert a string or dictionary of headers into a dictionary of headers.
"""
_split_otel_headers: Dict[str, str] = {}
if headers:
if isinstance(headers, str):
# when passed HEADERS="x-honeycomb-team=B85YgLm96******"
# Split only on first '=' occurrence
parts = headers.split("=", 1)
if len(parts) == 2:
_split_otel_headers = {parts[0]: parts[1]}
else:
_split_otel_headers = {}
elif isinstance(headers, dict):
_split_otel_headers = headers
return _split_otel_headers
async def async_management_endpoint_success_hook(
self,
logging_payload: ManagementEndpointLoggingPayload,
@ -948,3 +1006,18 @@ class OpenTelemetry(CustomLogger):
)
management_endpoint_span.set_status(Status(StatusCode.ERROR))
management_endpoint_span.end(end_time=_end_time_ns)
def create_litellm_proxy_request_started_span(
self,
start_time: datetime,
headers: dict,
) -> Optional[Span]:
"""
Create a span for the received proxy server request.
"""
return self.tracer.start_span(
name="Received Proxy Server Request",
start_time=self._to_ns(start_time),
context=self.get_traceparent_from_header(headers=headers),
kind=self.span_kind.SERVER,
)

View file

@ -29,6 +29,7 @@ from litellm.batches.batch_utils import _handle_completed_batch
from litellm.caching.caching import DualCache, InMemoryCache
from litellm.caching.caching_handler import LLMCachingHandler
from litellm.cost_calculator import _select_model_name_for_cost_calc
from litellm.integrations.arize.arize import ArizeLogger
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.mlflow import MlflowLogger
@ -76,7 +77,6 @@ from litellm.types.utils import (
from litellm.utils import _get_base_model_from_metadata, executor, print_verbose
from ..integrations.argilla import ArgillaLogger
from ..integrations.arize.arize import ArizeLogger
from ..integrations.arize.arize_phoenix import ArizePhoenixLogger
from ..integrations.athina import AthinaLogger
from ..integrations.azure_storage.azure_storage import AzureBlobStorageLogger
@ -2658,13 +2658,13 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
)
for callback in _in_memory_loggers:
if (
isinstance(callback, OpenTelemetry)
isinstance(callback, ArizeLogger)
and callback.callback_name == "arize"
):
return callback # type: ignore
_otel_logger = OpenTelemetry(config=otel_config, callback_name="arize")
_in_memory_loggers.append(_otel_logger)
return _otel_logger # type: ignore
_arize_otel_logger = ArizeLogger(config=otel_config, callback_name="arize")
_in_memory_loggers.append(_arize_otel_logger)
return _arize_otel_logger # type: ignore
elif logging_integration == "arize_phoenix":
from litellm.integrations.opentelemetry import (
OpenTelemetry,
@ -2897,15 +2897,13 @@ def get_custom_logger_compatible_class( # noqa: PLR0915
if isinstance(callback, OpenTelemetry):
return callback
elif logging_integration == "arize":
from litellm.integrations.opentelemetry import OpenTelemetry
if "ARIZE_SPACE_KEY" not in os.environ:
raise ValueError("ARIZE_SPACE_KEY not found in environment variables")
if "ARIZE_API_KEY" not in os.environ:
raise ValueError("ARIZE_API_KEY not found in environment variables")
for callback in _in_memory_loggers:
if (
isinstance(callback, OpenTelemetry)
isinstance(callback, ArizeLogger)
and callback.callback_name == "arize"
):
return callback

View file

@ -51,7 +51,7 @@ from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.service_account_checks import service_account_checks
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.utils import PrismaClient, ProxyLogging, _to_ns
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.types.services import ServiceTypes
user_api_key_service_logger_obj = ServiceLogging() # used for tracking latency on OTEL
@ -370,14 +370,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
)
if open_telemetry_logger is not None:
parent_otel_span = open_telemetry_logger.tracer.start_span(
name="Received Proxy Server Request",
start_time=_to_ns(start_time),
context=open_telemetry_logger.get_traceparent_from_header(
headers=request.headers
),
kind=open_telemetry_logger.span_kind.SERVER,
parent_otel_span = (
open_telemetry_logger.create_litellm_proxy_request_started_span(
start_time=start_time,
headers=dict(request.headers),
)
)
### USER-DEFINED AUTH FUNCTION ###

View file

@ -33,9 +33,12 @@ from litellm.types.utils import (
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from litellm.integrations.opentelemetry import OpenTelemetry
Span = _Span
else:
Span = Any
OpenTelemetry = Any
def showwarning(message, category, filename, lineno, file=None, line=None):
@ -777,7 +780,7 @@ disable_spend_logs = False
jwt_handler = JWTHandler()
prompt_injection_detection_obj: Optional[_OPTIONAL_PromptInjectionDetection] = None
store_model_in_db: bool = False
open_telemetry_logger: Optional[Any] = None
open_telemetry_logger: Optional[OpenTelemetry] = None
### INITIALIZE GLOBAL LOGGING OBJECT ###
proxy_logging_obj = ProxyLogging(
user_api_key_cache=user_api_key_cache, premium_user=premium_user

View file

@ -1,4 +1,4 @@
from typing import TYPE_CHECKING, Literal, Any
from typing import TYPE_CHECKING, Any, Literal, Optional
from pydantic import BaseModel
@ -6,9 +6,10 @@ if TYPE_CHECKING:
Protocol = Literal["otlp_grpc", "otlp_http"]
else:
Protocol = Any
class ArizeConfig(BaseModel):
space_key: str
api_key: str
space_key: Optional[str] = None
api_key: Optional[str] = None
protocol: Protocol
endpoint: str

View file

@ -1766,6 +1766,10 @@ class StandardCallbackDynamicParams(TypedDict, total=False):
# Humanloop dynamic params
humanloop_api_key: Optional[str]
# Arize dynamic params
arize_api_key: Optional[str]
arize_space_key: Optional[str]
# Logging settings
turn_off_message_logging: Optional[bool] # when true will not log messages

View file

@ -35,6 +35,26 @@ async def test_async_otel_callback():
await asyncio.sleep(2)
@pytest.mark.asyncio()
async def test_async_dynamic_arize_config():
litellm.set_verbose = True
verbose_proxy_logger.setLevel(logging.DEBUG)
verbose_logger.setLevel(logging.DEBUG)
litellm.success_callback = ["arize"]
await litellm.acompletion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hi test from arize dynamic config"}],
temperature=0.1,
user="OTEL_USER",
arize_api_key=os.getenv("ARIZE_SPACE_2_API_KEY"),
arize_space_key=os.getenv("ARIZE_SPACE_2_KEY"),
)
await asyncio.sleep(2)
@pytest.fixture
def mock_env_vars(monkeypatch):
monkeypatch.setenv("ARIZE_SPACE_KEY", "test_space_key")

View file

@ -3,16 +3,20 @@ import sys
import time
from unittest.mock import Mock, patch
import json
from litellm.main import completion
import opentelemetry.exporter.otlp.proto.grpc.trace_exporter
from typing import Optional
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system-path
from litellm.integrations._types.open_inference import SpanAttributes
from litellm.integrations.arize.arize import ArizeConfig, ArizeLogger
from litellm.integrations.custom_logger import CustomLogger
from litellm.main import completion
import litellm
from litellm.types.utils import Choices
from litellm.types.utils import Choices, StandardCallbackDynamicParams
import pytest
import asyncio
def test_arize_set_attributes():
@ -65,3 +69,43 @@ def test_arize_set_attributes():
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_TOTAL, 100)
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, 60)
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 40)
class TestArizeLogger(CustomLogger):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.standard_callback_dynamic_params: Optional[
StandardCallbackDynamicParams
] = None
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
print("logged kwargs", json.dumps(kwargs, indent=4, default=str))
self.standard_callback_dynamic_params = kwargs.get(
"standard_callback_dynamic_params"
)
@pytest.mark.asyncio
async def test_arize_dynamic_params():
"""verify arize ai dynamic params are recieved by a callback"""
test_arize_logger = TestArizeLogger()
litellm.callbacks = [test_arize_logger]
await litellm.acompletion(
model="gpt-4o",
messages=[{"role": "user", "content": "basic arize test"}],
mock_response="test",
arize_api_key="test_api_key_dynamic",
arize_space_key="test_space_key_dynamic",
)
await asyncio.sleep(2)
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"
)