mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
[Feat] QA - Arize Team based logging (#12331)
* add _get_tracer_with_dynamic_headers * fix construct_dynamic_arize_headers * [Feat] UI - Allow Viewing/Editing Team Based Callbacks (#12329) * add logging settings view on UI * fix change ordering * add construct_dynamic_otel_headers for arize * refactor common code * test_construct_dynamic_arize_headers * otel unit tests * test_arize_dynamic_params * test_arize_dynamic_headers_in_grpc_requests
This commit is contained in:
parent
cf8ce11fb8
commit
090f847bd9
6 changed files with 396 additions and 72 deletions
|
|
@ -105,28 +105,39 @@ class ArizeLogger(OpenTelemetry):
|
|||
pass
|
||||
|
||||
|
||||
@staticmethod
|
||||
def construct_dynamic_arize_headers(
|
||||
def construct_dynamic_otel_headers(
|
||||
self,
|
||||
standard_callback_dynamic_params: StandardCallbackDynamicParams
|
||||
):
|
||||
) -> Optional[dict]:
|
||||
"""
|
||||
Construct dynamic Arize headers from standard callback dynamic params
|
||||
|
||||
This is used for team/key based logging.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary of dynamic Arize headers
|
||||
"""
|
||||
dynamic_headers = {}
|
||||
|
||||
#########################################################
|
||||
# `arize-space-id` handling
|
||||
# the suggested param is `arize_space_key`
|
||||
#########################################################
|
||||
if standard_callback_dynamic_params.get("arize_space_id"):
|
||||
dynamic_headers["arize-space-id"] = standard_callback_dynamic_params.get(
|
||||
"arize_space_id"
|
||||
)
|
||||
if standard_callback_dynamic_params.get("arize_space_key"):
|
||||
dynamic_headers["space_key"] = standard_callback_dynamic_params.get(
|
||||
dynamic_headers["arize-space-id"] = standard_callback_dynamic_params.get(
|
||||
"arize_space_key"
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# `api_key` handling
|
||||
#########################################################
|
||||
if standard_callback_dynamic_params.get("arize_api_key"):
|
||||
dynamic_headers["api_key"] = standard_callback_dynamic_params.get(
|
||||
"arize_api_key"
|
||||
)
|
||||
|
||||
if standard_callback_dynamic_params.get("arize_space_id"):
|
||||
dynamic_headers["arize-space-id"] = standard_callback_dynamic_params.get(
|
||||
"arize_space_id"
|
||||
)
|
||||
return dynamic_headers
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ 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 opentelemetry.trace import Tracer as _Tracer
|
||||
|
||||
from litellm.proxy._types import (
|
||||
ManagementEndpointLoggingPayload as _ManagementEndpointLoggingPayload,
|
||||
|
|
@ -26,12 +27,14 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy.proxy_server import UserAPIKeyAuth as _UserAPIKeyAuth
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
Tracer = Union[_Tracer, Any]
|
||||
Context = Union[_Context, Any]
|
||||
SpanExporter = Union[_SpanExporter, Any]
|
||||
UserAPIKeyAuth = Union[_UserAPIKeyAuth, Any]
|
||||
ManagementEndpointLoggingPayload = Union[_ManagementEndpointLoggingPayload, Any]
|
||||
else:
|
||||
Span = Any
|
||||
Tracer = Any
|
||||
SpanExporter = Any
|
||||
UserAPIKeyAuth = Any
|
||||
ManagementEndpointLoggingPayload = Any
|
||||
|
|
@ -313,6 +316,71 @@ class OpenTelemetry(CustomLogger):
|
|||
|
||||
# End Parent OTEL Sspan
|
||||
parent_otel_span.end(end_time=self._to_ns(datetime.now()))
|
||||
|
||||
#########################################################
|
||||
# Team/Key Based Logging Control Flow
|
||||
#########################################################
|
||||
def get_tracer_to_use_for_request(self, kwargs: dict) -> Tracer:
|
||||
"""
|
||||
Get the tracer to use for this request
|
||||
|
||||
If dynamic headers are present, a temporary tracer is created with the dynamic headers.
|
||||
Otherwise, the default tracer is used.
|
||||
|
||||
Returns:
|
||||
Tracer: The tracer to use for this request
|
||||
"""
|
||||
dynamic_headers = self._get_dynamic_otel_headers_from_kwargs(kwargs)
|
||||
|
||||
if dynamic_headers is not None:
|
||||
# Create spans using a temporary tracer with dynamic headers
|
||||
tracer_to_use = self._get_tracer_with_dynamic_headers(dynamic_headers)
|
||||
verbose_logger.debug("Using dynamic headers for this request: %s", dynamic_headers)
|
||||
else:
|
||||
tracer_to_use = self.tracer
|
||||
|
||||
return tracer_to_use
|
||||
|
||||
def _get_dynamic_otel_headers_from_kwargs(self, kwargs) -> Optional[dict]:
|
||||
"""Extract dynamic headers from kwargs if available."""
|
||||
standard_callback_dynamic_params: Optional[
|
||||
StandardCallbackDynamicParams
|
||||
] = kwargs.get("standard_callback_dynamic_params")
|
||||
|
||||
if not standard_callback_dynamic_params:
|
||||
return None
|
||||
|
||||
dynamic_headers = self.construct_dynamic_otel_headers(
|
||||
standard_callback_dynamic_params=standard_callback_dynamic_params
|
||||
)
|
||||
|
||||
return dynamic_headers if dynamic_headers else None
|
||||
|
||||
def _get_tracer_with_dynamic_headers(self, dynamic_headers: dict):
|
||||
"""Create a temporary tracer with dynamic headers for this request only."""
|
||||
from opentelemetry.sdk.resources import Resource
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
|
||||
# Create a temporary tracer provider with dynamic headers
|
||||
temp_provider = TracerProvider(resource=Resource(attributes=LITELLM_RESOURCE))
|
||||
temp_provider.add_span_processor(self._get_span_processor(dynamic_headers=dynamic_headers))
|
||||
|
||||
return temp_provider.get_tracer(LITELLM_TRACER_NAME)
|
||||
|
||||
def construct_dynamic_otel_headers(self, standard_callback_dynamic_params: StandardCallbackDynamicParams) -> Optional[dict]:
|
||||
"""
|
||||
Construct dynamic headers from standard callback dynamic params
|
||||
|
||||
Note: You just need to override this method in Arize, Langfuse Otel if you want to allow team/key based logging.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary of dynamic headers
|
||||
"""
|
||||
return None
|
||||
|
||||
#########################################################
|
||||
# End of Team/Key Based Logging Control Flow
|
||||
#########################################################
|
||||
|
||||
def _handle_sucess(self, kwargs, response_obj, start_time, end_time):
|
||||
from opentelemetry import trace
|
||||
|
|
@ -323,12 +391,11 @@ class OpenTelemetry(CustomLogger):
|
|||
kwargs,
|
||||
self.config,
|
||||
)
|
||||
|
||||
_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(
|
||||
# Span 1: Request sent to litellm SDK
|
||||
otel_tracer: Tracer = self.get_tracer_to_use_for_request(kwargs)
|
||||
span = otel_tracer.start_span(
|
||||
name=self._get_span_name(kwargs),
|
||||
start_time=self._to_ns(start_time),
|
||||
context=_parent_context,
|
||||
|
|
@ -342,7 +409,7 @@ class OpenTelemetry(CustomLogger):
|
|||
pass
|
||||
else:
|
||||
# Span 2: Raw Request / Response to LLM
|
||||
raw_request_span = self.tracer.start_span(
|
||||
raw_request_span = otel_tracer.start_span(
|
||||
name=RAW_REQUEST_SPAN_NAME,
|
||||
start_time=self._to_ns(start_time),
|
||||
context=trace.set_span_in_context(span),
|
||||
|
|
@ -387,7 +454,8 @@ class OpenTelemetry(CustomLogger):
|
|||
if end_time_float is not None:
|
||||
end_time_datetime = datetime.fromtimestamp(end_time_float)
|
||||
|
||||
guardrail_span = self.tracer.start_span(
|
||||
otel_tracer: Tracer = self.get_tracer_to_use_for_request(kwargs)
|
||||
guardrail_span = otel_tracer.start_span(
|
||||
name="guardrail",
|
||||
start_time=self._to_ns(start_time_datetime),
|
||||
context=context,
|
||||
|
|
@ -420,39 +488,6 @@ class OpenTelemetry(CustomLogger):
|
|||
|
||||
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.
|
||||
|
||||
This allows for per-request configuration of telemetry exporters by
|
||||
extracting headers from standard_callback_dynamic_params.
|
||||
"""
|
||||
from opentelemetry import trace
|
||||
|
||||
from litellm.integrations.arize.arize import ArizeLogger
|
||||
|
||||
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
|
||||
dynamic_headers = ArizeLogger.construct_dynamic_arize_headers(standard_callback_dynamic_params=standard_callback_dynamic_params)
|
||||
|
||||
# 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
|
||||
|
|
@ -465,7 +500,8 @@ class OpenTelemetry(CustomLogger):
|
|||
_parent_context, parent_otel_span = self._get_span_context(kwargs)
|
||||
|
||||
# Span 1: Requst sent to litellm SDK
|
||||
span = self.tracer.start_span(
|
||||
otel_tracer: Tracer = self.get_tracer_to_use_for_request(kwargs)
|
||||
span = otel_tracer.start_span(
|
||||
name=self._get_span_name(kwargs),
|
||||
start_time=self._to_ns(start_time),
|
||||
context=_parent_context,
|
||||
|
|
|
|||
|
|
@ -6,12 +6,3 @@ model_list:
|
|||
litellm_params:
|
||||
model: openai/*
|
||||
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: "bedrock-post-guard"
|
||||
litellm_params:
|
||||
guardrail: bedrock # supported values: "aporia", "bedrock", "lakera"
|
||||
mode: "post_call"
|
||||
guardrailIdentifier: wf0hkdb5x07f # your guardrail ID on bedrock
|
||||
guardrailVersion: "DRAFT" # your guardrail version on bedrock
|
||||
default_on: true
|
||||
167
tests/test_litellm/integrations/arize/test_arize.py
Normal file
167
tests/test_litellm/integrations/arize/test_arize.py
Normal file
|
|
@ -0,0 +1,167 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from typing import Optional
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
# Adds the grandparent directory to sys.path to allow importing project modules
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.arize.arize import ArizeLogger
|
||||
from litellm.integrations.opentelemetry import OpenTelemetryConfig
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_arize_dynamic_params():
|
||||
"""Test that the OpenTelemetry logger uses the correct dynamic headers for each Arize request."""
|
||||
|
||||
# Create ArizeLogger instance
|
||||
arize_logger = ArizeLogger()
|
||||
|
||||
# Capture the get_tracer_to_use_for_request calls
|
||||
tracer_calls = []
|
||||
original_get_tracer = arize_logger.get_tracer_to_use_for_request
|
||||
|
||||
def mock_get_tracer_to_use_for_request(kwargs):
|
||||
# Capture the kwargs to see what dynamic headers are being used
|
||||
tracer_calls.append(kwargs)
|
||||
# Return the default tracer
|
||||
return arize_logger.tracer
|
||||
|
||||
# Mock the get_tracer_to_use_for_request method
|
||||
arize_logger.get_tracer_to_use_for_request = mock_get_tracer_to_use_for_request
|
||||
|
||||
# Set up callbacks
|
||||
litellm.callbacks = [arize_logger]
|
||||
|
||||
# First request with team1 credentials
|
||||
await litellm.acompletion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi test from arize dynamic config"}],
|
||||
temperature=0.1,
|
||||
mock_response="test_response",
|
||||
arize_api_key="team1_key",
|
||||
arize_space_id="team1_space_id"
|
||||
)
|
||||
|
||||
# Second request with team2 credentials
|
||||
await litellm.acompletion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi test from arize dynamic config"}],
|
||||
temperature=0.1,
|
||||
mock_response="test_response",
|
||||
arize_api_key="team2_key",
|
||||
arize_space_id="team2_space_id"
|
||||
)
|
||||
|
||||
# Allow some time for async processing
|
||||
await asyncio.sleep(5)
|
||||
|
||||
# Assertions
|
||||
print(f"Tracer calls: {len(tracer_calls)}")
|
||||
|
||||
# We should have captured calls for both requests
|
||||
assert len(tracer_calls) >= 2, f"Expected at least 2 tracer calls, got {len(tracer_calls)}"
|
||||
|
||||
# Check that we have the expected dynamic params in the kwargs
|
||||
team1_found = False
|
||||
team2_found = False
|
||||
|
||||
print("args to tracer calls", tracer_calls)
|
||||
|
||||
for call_kwargs in tracer_calls:
|
||||
dynamic_params = call_kwargs.get("standard_callback_dynamic_params", {})
|
||||
if dynamic_params.get("arize_api_key") == "team1_key":
|
||||
team1_found = True
|
||||
assert dynamic_params.get("arize_space_id") == "team1_space_id"
|
||||
elif dynamic_params.get("arize_api_key") == "team2_key":
|
||||
team2_found = True
|
||||
assert dynamic_params.get("arize_space_id") == "team2_space_id"
|
||||
|
||||
# Verify both teams were found
|
||||
assert team1_found, "team1 dynamic params not found"
|
||||
assert team2_found, "team2 dynamic params not found"
|
||||
|
||||
print("✅ All assertions passed - OpenTelemetry logger correctly received dynamic params")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_arize_dynamic_headers_in_grpc_requests():
|
||||
"""Test that dynamic Arize params are passed as headers to the gRPC/HTTP exporter."""
|
||||
|
||||
# Track all exporter calls and their headers
|
||||
exporter_headers = []
|
||||
|
||||
def mock_otlp_http_exporter(*args, **kwargs):
|
||||
# Capture the headers passed to the HTTP exporter
|
||||
headers = kwargs.get('headers', {})
|
||||
exporter_headers.append(headers)
|
||||
|
||||
# Return a mock exporter
|
||||
mock_exporter = MagicMock()
|
||||
mock_exporter.export = MagicMock(return_value=None)
|
||||
return mock_exporter
|
||||
|
||||
# Patch the HTTP exporter (Arize uses HTTP by default)
|
||||
with patch('opentelemetry.exporter.otlp.proto.http.trace_exporter.OTLPSpanExporter', mock_otlp_http_exporter):
|
||||
|
||||
# Create ArizeLogger with HTTP configuration
|
||||
config = OpenTelemetryConfig(
|
||||
exporter="otlp_http",
|
||||
endpoint="https://otlp.arize.com/v1"
|
||||
)
|
||||
arize_logger = ArizeLogger(config=config)
|
||||
litellm.callbacks = [arize_logger]
|
||||
|
||||
# Request 1: team1 dynamic params
|
||||
await litellm.acompletion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi from team1"}],
|
||||
mock_response="response1",
|
||||
arize_api_key="team1_api_key",
|
||||
arize_space_id="team1_space_id"
|
||||
)
|
||||
|
||||
# Request 2: team2 dynamic params
|
||||
await litellm.acompletion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi from team2"}],
|
||||
mock_response="response2",
|
||||
arize_api_key="team2_api_key",
|
||||
arize_space_id="team2_space_id"
|
||||
)
|
||||
|
||||
# Allow time for async processing
|
||||
await asyncio.sleep(3)
|
||||
|
||||
# Assertions
|
||||
print(f"Captured exporter headers: {exporter_headers}")
|
||||
|
||||
# Should have multiple exporter calls (default + dynamic)
|
||||
assert len(exporter_headers) >= 2, f"Expected at least 2 exporter calls, got {len(exporter_headers)}"
|
||||
|
||||
# Find team1 and team2 headers
|
||||
team1_found = False
|
||||
team2_found = False
|
||||
|
||||
for headers in exporter_headers:
|
||||
if headers.get('api_key') == 'team1_api_key' and headers.get('arize-space-id') == 'team1_space_id':
|
||||
team1_found = True
|
||||
print(f"✅ Found team1 headers: {headers}")
|
||||
elif headers.get('api_key') == 'team2_api_key' and headers.get('arize-space-id') == 'team2_space_id':
|
||||
team2_found = True
|
||||
print(f"✅ Found team2 headers: {headers}")
|
||||
|
||||
# Verify both dynamic header sets were used
|
||||
assert team1_found, "team1 dynamic headers not found in exporter calls"
|
||||
assert team2_found, "team2 dynamic headers not found in exporter calls"
|
||||
|
||||
print("✅ Test passed - Dynamic Arize params correctly passed to gRPC/HTTP exporter")
|
||||
|
||||
|
||||
|
|
@ -239,35 +239,46 @@ def test_construct_dynamic_arize_headers():
|
|||
Test the construct_dynamic_arize_headers method with various input scenarios.
|
||||
Ensures that dynamic Arize headers are properly constructed from callback parameters.
|
||||
"""
|
||||
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
||||
# Test with all parameters present
|
||||
dynamic_params_full = {
|
||||
"arize_space_key": "test_space_key",
|
||||
"arize_api_key": "test_api_key",
|
||||
"arize_space_id": "test_space_id"
|
||||
}
|
||||
dynamic_params_full = StandardCallbackDynamicParams(
|
||||
arize_api_key="test_api_key",
|
||||
arize_space_id="test_space_id"
|
||||
)
|
||||
arize_logger = ArizeLogger()
|
||||
|
||||
headers = ArizeLogger.construct_dynamic_arize_headers(dynamic_params_full)
|
||||
headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_full)
|
||||
expected_headers = {
|
||||
"space_key": "test_space_key",
|
||||
"api_key": "test_api_key",
|
||||
"arize-space-id": "test_space_id"
|
||||
}
|
||||
assert headers == expected_headers
|
||||
|
||||
# Test with only space_id
|
||||
dynamic_params_space_id_only = {
|
||||
"arize_space_id": "test_space_id"
|
||||
}
|
||||
dynamic_params_space_id_only = StandardCallbackDynamicParams(
|
||||
arize_space_id="test_space_id"
|
||||
)
|
||||
|
||||
headers = ArizeLogger.construct_dynamic_arize_headers(dynamic_params_space_id_only)
|
||||
headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_space_id_only)
|
||||
expected_headers = {
|
||||
"arize-space-id": "test_space_id"
|
||||
}
|
||||
assert headers == expected_headers
|
||||
|
||||
# Test with empty parameters dict
|
||||
dynamic_params_empty = {}
|
||||
dynamic_params_empty = StandardCallbackDynamicParams()
|
||||
|
||||
headers = ArizeLogger.construct_dynamic_arize_headers(dynamic_params_empty)
|
||||
headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_empty)
|
||||
assert headers == {}
|
||||
|
||||
# test with space key and api key
|
||||
dynamic_params_space_key_and_api_key = StandardCallbackDynamicParams(
|
||||
arize_space_key="test_space_key",
|
||||
arize_api_key="test_api_key"
|
||||
)
|
||||
headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_space_key_and_api_key)
|
||||
expected_headers = {
|
||||
"arize-space-id": "test_space_key",
|
||||
"api_key": "test_api_key"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -66,3 +66,111 @@ class TestOpenTelemetry(unittest.TestCase):
|
|||
|
||||
# Verify that start_span was never called
|
||||
otel.tracer.start_span.assert_not_called()
|
||||
|
||||
def test_get_tracer_to_use_for_request_with_dynamic_headers(self):
|
||||
"""Test that get_tracer_to_use_for_request returns a dynamic tracer when dynamic headers are present."""
|
||||
# Setup
|
||||
otel = OpenTelemetry()
|
||||
otel.tracer = MagicMock()
|
||||
|
||||
# Mock the dynamic header extraction and tracer creation
|
||||
with patch.object(otel, '_get_dynamic_otel_headers_from_kwargs') as mock_get_headers, \
|
||||
patch.object(otel, '_get_tracer_with_dynamic_headers') as mock_get_tracer:
|
||||
|
||||
# Test case 1: With dynamic headers
|
||||
mock_get_headers.return_value = {"arize-space-id": "test-space", "api_key": "test-key"}
|
||||
mock_dynamic_tracer = MagicMock()
|
||||
mock_get_tracer.return_value = mock_dynamic_tracer
|
||||
|
||||
kwargs = {"standard_callback_dynamic_params": {"arize_space_key": "test-space"}}
|
||||
result = otel.get_tracer_to_use_for_request(kwargs)
|
||||
|
||||
# Assertions
|
||||
mock_get_headers.assert_called_once_with(kwargs)
|
||||
mock_get_tracer.assert_called_once_with({"arize-space-id": "test-space", "api_key": "test-key"})
|
||||
self.assertEqual(result, mock_dynamic_tracer)
|
||||
|
||||
def test_get_tracer_to_use_for_request_without_dynamic_headers(self):
|
||||
"""Test that get_tracer_to_use_for_request returns the default tracer when no dynamic headers are present."""
|
||||
# Setup
|
||||
otel = OpenTelemetry()
|
||||
otel.tracer = MagicMock()
|
||||
|
||||
# Mock the dynamic header extraction to return None
|
||||
with patch.object(otel, '_get_dynamic_otel_headers_from_kwargs') as mock_get_headers:
|
||||
mock_get_headers.return_value = None
|
||||
|
||||
kwargs = {}
|
||||
result = otel.get_tracer_to_use_for_request(kwargs)
|
||||
|
||||
# Assertions
|
||||
mock_get_headers.assert_called_once_with(kwargs)
|
||||
self.assertEqual(result, otel.tracer)
|
||||
|
||||
def test_get_dynamic_otel_headers_from_kwargs(self):
|
||||
"""Test that _get_dynamic_otel_headers_from_kwargs correctly extracts dynamic headers from kwargs."""
|
||||
# Setup
|
||||
otel = OpenTelemetry()
|
||||
|
||||
# Mock the construct_dynamic_otel_headers method
|
||||
with patch.object(otel, 'construct_dynamic_otel_headers') as mock_construct:
|
||||
# Test case 1: With standard_callback_dynamic_params
|
||||
mock_construct.return_value = {"arize-space-id": "test-space", "api_key": "test-key"}
|
||||
|
||||
standard_params = {
|
||||
"arize_space_key": "test-space",
|
||||
"arize_api_key": "test-key"
|
||||
}
|
||||
kwargs = {"standard_callback_dynamic_params": standard_params}
|
||||
|
||||
result = otel._get_dynamic_otel_headers_from_kwargs(kwargs)
|
||||
|
||||
# Assertions
|
||||
mock_construct.assert_called_once_with(standard_callback_dynamic_params=standard_params)
|
||||
self.assertEqual(result, {"arize-space-id": "test-space", "api_key": "test-key"})
|
||||
|
||||
# Test case 2: Without standard_callback_dynamic_params
|
||||
kwargs_empty = {}
|
||||
result_empty = otel._get_dynamic_otel_headers_from_kwargs(kwargs_empty)
|
||||
|
||||
# Should return None when no dynamic params
|
||||
self.assertIsNone(result_empty)
|
||||
|
||||
# Test case 3: With empty construct result
|
||||
mock_construct.return_value = {}
|
||||
result_empty_construct = otel._get_dynamic_otel_headers_from_kwargs(kwargs)
|
||||
|
||||
# Should return None when construct returns empty dict
|
||||
self.assertIsNone(result_empty_construct)
|
||||
|
||||
@patch("opentelemetry.sdk.trace.TracerProvider")
|
||||
@patch("opentelemetry.sdk.resources.Resource")
|
||||
def test_get_tracer_with_dynamic_headers(self, mock_resource, mock_tracer_provider):
|
||||
"""Test that _get_tracer_with_dynamic_headers creates a temporary tracer with dynamic headers."""
|
||||
# Setup
|
||||
otel = OpenTelemetry()
|
||||
|
||||
# Mock the span processor creation
|
||||
with patch.object(otel, '_get_span_processor') as mock_get_span_processor:
|
||||
mock_span_processor = MagicMock()
|
||||
mock_get_span_processor.return_value = mock_span_processor
|
||||
|
||||
# Mock the tracer provider and its methods
|
||||
mock_provider_instance = MagicMock()
|
||||
mock_tracer_provider.return_value = mock_provider_instance
|
||||
mock_tracer = MagicMock()
|
||||
mock_provider_instance.get_tracer.return_value = mock_tracer
|
||||
|
||||
# Mock the resource
|
||||
mock_resource_instance = MagicMock()
|
||||
mock_resource.return_value = mock_resource_instance
|
||||
|
||||
# Test
|
||||
dynamic_headers = {"arize-space-id": "test-space", "api_key": "test-key"}
|
||||
result = otel._get_tracer_with_dynamic_headers(dynamic_headers)
|
||||
|
||||
# Assertions
|
||||
mock_get_span_processor.assert_called_once_with(dynamic_headers=dynamic_headers)
|
||||
mock_provider_instance.add_span_processor.assert_called_once_with(mock_span_processor)
|
||||
mock_provider_instance.get_tracer.assert_called_once_with("litellm")
|
||||
self.assertEqual(result, mock_tracer)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue