diff --git a/python/scenario/_tracing/setup.py b/python/scenario/_tracing/setup.py index 2b9d4ee00..f4f7fd638 100644 --- a/python/scenario/_tracing/setup.py +++ b/python/scenario/_tracing/setup.py @@ -10,7 +10,7 @@ import logging import os -from typing import List, Optional, Sequence +from typing import List, Optional, Protocol, Sequence, cast from opentelemetry import trace from opentelemetry.sdk.trace import TracerProvider, SpanProcessor @@ -25,6 +25,10 @@ _initialized = False +class _SpanProcessorProvider(Protocol): + def add_span_processor(self, span_processor: SpanProcessor) -> None: ... + + def setup_scenario_tracing( *, span_filter: Optional[SpanFilter] = None, @@ -88,7 +92,7 @@ def ensure_tracing_initialized(observability: Optional[dict] = None) -> None: _initialized = True -def _get_concrete_provider(provider) -> Optional[TracerProvider]: +def _get_concrete_provider(provider) -> Optional[_SpanProcessorProvider]: """Returns the concrete TracerProvider if one exists. Checks the provider itself and one level of delegation @@ -97,6 +101,12 @@ def _get_concrete_provider(provider) -> Optional[TracerProvider]: if isinstance(provider, TracerProvider): return provider + # OpenTelemetry's public provider interface does not require an SDK + # TracerProvider subclass. Providers such as Temporal's replay-safe wrapper + # expose the processor hook directly and are safe to configure in place. + if callable(getattr(provider, "add_span_processor", None)): + return cast(_SpanProcessorProvider, provider) + # Check delegation pattern delegate = None if hasattr(provider, "get_delegate"): @@ -122,16 +132,17 @@ def _do_setup( concrete = _get_concrete_provider(existing_provider) if concrete is not None: - _attach_to_existing(concrete, span_filter, span_processors, trace_exporter) + _attach_to_existing(concrete, span_filter, span_processors, trace_exporter, instrumentors) else: _full_setup(span_filter, span_processors, trace_exporter, instrumentors) def _attach_to_existing( - provider: TracerProvider, + provider: _SpanProcessorProvider, span_filter: Optional[SpanFilter], span_processors: Optional[List[SpanProcessor]], trace_exporter: Optional[SpanExporter], + instrumentors: Optional[Sequence], ) -> None: """Attach processors to an existing TracerProvider.""" provider.add_span_processor(judge_span_collector) @@ -150,6 +161,10 @@ def _attach_to_existing( # Add LangWatch exporter _add_langwatch_exporter(provider, span_filter) + if instrumentors: + for instrumentor in instrumentors: + instrumentor.instrument(tracer_provider=provider) + def _full_setup( span_filter: Optional[SpanFilter], diff --git a/python/tests/test_tracing_setup.py b/python/tests/test_tracing_setup.py index 69339dacd..dae3e2857 100644 --- a/python/tests/test_tracing_setup.py +++ b/python/tests/test_tracing_setup.py @@ -8,6 +8,7 @@ setup_scenario_tracing, ensure_tracing_initialized, _get_concrete_provider, + _do_setup, _reset_tracing_for_tests, ) @@ -92,6 +93,17 @@ def test_returns_tracer_provider_directly(self) -> None: assert result is provider + def test_returns_public_provider_with_span_processor_support(self) -> None: + class PublicProvider: + def add_span_processor(self, span_processor) -> None: + del span_processor + + provider = PublicProvider() + + result = _get_concrete_provider(provider) + + assert result is provider + def test_returns_delegate_from_get_delegate(self) -> None: concrete = TracerProvider() @@ -129,6 +141,20 @@ class FakeProxy: assert result is None +def test_instruments_an_existing_public_provider() -> None: + class PublicProvider: + def add_span_processor(self, span_processor) -> None: + del span_processor + + provider = PublicProvider() + instrumentor = MagicMock() + + with patch("scenario._tracing.setup.trace.get_tracer_provider", return_value=provider): + _do_setup(instrumentors=[instrumentor]) + + instrumentor.instrument.assert_called_once_with(tracer_provider=provider) + + class TestResetTracingForTests: """Tests for _reset_tracing_for_tests."""