diff --git a/burr/integrations/opentelemetry.py b/burr/integrations/opentelemetry.py index 7a6970e88..a09092247 100644 --- a/burr/integrations/opentelemetry.py +++ b/burr/integrations/opentelemetry.py @@ -66,6 +66,9 @@ class FullSpanContext: action_span: ActionSpan partition_key: str app_id: str + # The tracker that owns this span. Spans can end on a thread where tracker_context is unset, + # so we carry the tracker with the span instead of relying only on the context var. + tracker: Optional[SyncTrackingClient] = None span_map = {} @@ -77,7 +80,7 @@ def cache_span(span: Span, context: FullSpanContext) -> Span: def uncache_span(span: Span) -> Span: - del span_map[span.get_span_context().span_id] + span_map.pop(span.get_span_context().span_id, None) return span @@ -412,6 +415,7 @@ def pre_run_step( ), partition_key=partition_key, app_id=app_id, + tracker=self.burr_tracker, ), ) @@ -476,10 +480,15 @@ class BurrTrackingSpanProcessor(SpanProcessor): def tracker(self): """Quick trick to get closer to the right tracker. This is suboptimal as we don't really have guarentees that we'll be *in* the right context when it gets logged, but the way OpenTel - is implemented we will (with the immediate span processor). TODO -- track a map of span ID -> tracker + is implemented we will (with the immediate span processor). When the context is not set + (e.g. a span ending on a worker thread), we fall back to the tracker cached with the span. """ return tracker_context.get() + def _tracker_for(self, cached_span: FullSpanContext) -> Optional[SyncTrackingClient]: + tracker = self.tracker + return tracker if tracker is not None else cached_span.tracker + def on_start( self, span: "Span", @@ -491,16 +500,18 @@ def on_start( parent_span = get_cached_span(span.parent.span_id) # If it exists, we can spawn a new span and cache that if parent_span is not None: + tracker = self._tracker_for(parent_span) cache_span( span, context := FullSpanContext( action_span=parent_span.action_span.spawn(span.name), partition_key=parent_span.partition_key, app_id=parent_span.app_id, + tracker=tracker, ), ) - if self.tracker is not None: - self.tracker.pre_start_span( + if tracker is not None: + tracker.pre_start_span( action=context.action_span.action, action_sequence_id=context.action_span.action_sequence_id, span=context.action_span, @@ -512,9 +523,13 @@ def on_start( def on_end(self, span: "Span") -> None: cached_span = get_cached_span(span.get_span_context().span_id) # If this is none it means we're outside of the burr context - if cached_span is not None and self.tracker is not None: - # TODO -- get tracker context to work - self.tracker.post_end_span( + if cached_span is None: + return + # Always drop the entry first so it cannot leak, even if no tracker is found or logging fails + uncache_span(span) + tracker = self._tracker_for(cached_span) + if tracker is not None: + tracker.post_end_span( action=cached_span.action_span.action, action_sequence_id=cached_span.action_span.action_sequence_id, span=cached_span.action_span, @@ -522,9 +537,8 @@ def on_end(self, span: "Span") -> None: app_id=cached_span.app_id, partition_key=cached_span.partition_key, ) - uncache_span(span) if len(span.attributes) > 0: - self.tracker.do_log_attributes( + tracker.do_log_attributes( attributes=dict(**span.attributes), action=cached_span.action_span.action, action_sequence_id=cached_span.action_span.action_sequence_id, diff --git a/tests/integrations/test_burr_opentelemetry.py b/tests/integrations/test_burr_opentelemetry.py index 9c5f640c3..f2ed8ee87 100644 --- a/tests/integrations/test_burr_opentelemetry.py +++ b/tests/integrations/test_burr_opentelemetry.py @@ -15,19 +15,25 @@ # specific language governing permissions and limitations # under the License. +import contextvars import json +from concurrent.futures import ThreadPoolExecutor from unittest.mock import Mock, patch import pydantic import pytest -from opentelemetry.sdk.trace import Span +from opentelemetry import context as otel_context +from opentelemetry.sdk.trace import Span, TracerProvider from opentelemetry.trace import SpanContext -from burr.core import serde +from burr.core import State, serde +from burr.core.action import Action from burr.integrations.opentelemetry import ( BurrTrackingSpanProcessor, FullSpanContext, + OpenTelemetryTracker, convert_to_otel_attribute, + span_map, tracker_context, ) from burr.tracking.base import SyncTrackingClient @@ -39,6 +45,14 @@ class SampleModel(pydantic.BaseModel): bar: bool +@pytest.fixture(autouse=True) +def clear_span_map(): + """span_map is module-global, so keep a failing test from leaking entries into the next one.""" + span_map.clear() + yield + span_map.clear() + + @pytest.mark.parametrize( "value, expected", [ @@ -73,6 +87,7 @@ def test_burr_tracking_span_processor_on_start_with_none_tracker(): mock_parent_context.action_span.spawn = Mock(return_value=mock_spawned_span) mock_parent_context.partition_key = "test_partition" mock_parent_context.app_id = "test_app" + mock_parent_context.tracker = None mock_get_cached.return_value = mock_parent_context # Mock cache_span @@ -105,10 +120,11 @@ def test_burr_tracking_span_processor_on_end_with_none_tracker(): mock_cached_span.action_span.action_sequence_id = 1 mock_cached_span.app_id = "test_app" mock_cached_span.partition_key = "test_partition" + mock_cached_span.tracker = None mock_get_cached.return_value = mock_cached_span # Mock uncache_span - with patch("burr.integrations.opentelemetry.uncache_span"): + with patch("burr.integrations.opentelemetry.uncache_span") as mock_uncache: # Set tracker_context to None (simulating no tracker in context) token = tracker_context.set(None) try: @@ -116,6 +132,8 @@ def test_burr_tracking_span_processor_on_end_with_none_tracker(): processor.on_end(mock_span) finally: tracker_context.reset(token) + # The span is still removed from the cache when no tracker is found + mock_uncache.assert_called_once_with(mock_span) def test_burr_tracking_span_processor_on_start_with_valid_tracker(): @@ -192,3 +210,152 @@ def test_burr_tracking_span_processor_on_end_with_valid_tracker(): assert mock_tracker.post_end_span.called finally: tracker_context.reset(token) + + +def _tracker_with_local_provider(): + provider = TracerProvider() + provider.add_span_processor(BurrTrackingSpanProcessor()) + tracer = provider.get_tracer("test") + burr_tracker = Mock(spec=SyncTrackingClient) + with patch("burr.integrations.opentelemetry.initialize_tracer"): + otel_tracker = OpenTelemetryTracker(burr_tracker=burr_tracker) + otel_tracker.tracer = tracer + action = Mock(spec=Action) + action.name = "work" + return tracer, burr_tracker, otel_tracker, action + + +def _pre_run_step(otel_tracker, action): + otel_tracker.pre_run_step( + app_id="test_app", + partition_key="test_partition", + sequence_id=0, + state=State({}), + action=action, + inputs={}, + ) + + +def _post_run_step(otel_tracker, action): + otel_tracker.post_run_step( + app_id="test_app", + partition_key="test_partition", + sequence_id=0, + state=State({}), + action=action, + result={}, + exception=None, + ) + + +@pytest.mark.parametrize("on_worker_thread", [False, True]) +def test_burr_tracking_span_processor_child_span_reaches_tracker_and_is_uncached(on_worker_thread): + """Spans started inside an action are logged to the Burr tracker and removed from the span + cache, whether they end on the action's thread or on a worker thread that only carries the + OpenTelemetry context (thread pools don't copy context vars, so tracker_context is unset there). + """ + tracer, burr_tracker, otel_tracker, action = _tracker_with_local_provider() + child_span_ids = [] + + def child(): + with tracer.start_as_current_span("llm_call") as span: + span.set_attribute("model", "test-model") + child_span_ids.append(span.get_span_context().span_id) + assert child_span_ids[-1] in span_map + + def run_step(): + _pre_run_step(otel_tracker, action) + if on_worker_thread: + ctx = otel_context.get_current() + + def run_with_otel_context(): + token = otel_context.attach(ctx) + try: + child() + finally: + otel_context.detach(token) + + with ThreadPoolExecutor(max_workers=1) as executor: + executor.submit(run_with_otel_context).result() + else: + child() + _post_run_step(otel_tracker, action) + + contextvars.copy_context().run(run_step) + + assert len(child_span_ids) == 1 + assert child_span_ids[0] not in span_map + assert len(span_map) == 0 + burr_tracker.pre_start_span.assert_called_once() + assert burr_tracker.pre_start_span.call_args.kwargs["span"].name == "llm_call" + # one for the child span, one for the action span itself + assert burr_tracker.post_end_span.call_count == 2 + burr_tracker.do_log_attributes.assert_called_once() + assert burr_tracker.do_log_attributes.call_args.kwargs["attributes"] == {"model": "test-model"} + + +def test_burr_tracking_span_processor_child_span_ending_after_step_is_uncached(): + """A span started during an action but ended after post_run_step (which clears + tracker_context) is still logged to the owning tracker and removed from the span cache.""" + tracer, burr_tracker, otel_tracker, action = _tracker_with_local_provider() + + def run_step(): + _pre_run_step(otel_tracker, action) + span = tracer.start_span("background_call") + _post_run_step(otel_tracker, action) + assert tracker_context.get() is None + span.end() + return span.get_span_context().span_id + + span_id = contextvars.copy_context().run(run_step) + + assert span_id not in span_map + assert len(span_map) == 0 + assert burr_tracker.post_end_span.call_args.kwargs["span"].name == "background_call" + + +def test_burr_tracking_span_processor_uncaches_span_when_tracker_raises(): + """A failing tracker callback in on_end must not leave the span in the cache.""" + tracer, burr_tracker, otel_tracker, action = _tracker_with_local_provider() + burr_tracker.post_end_span.side_effect = RuntimeError("tracker failed") + span_ids = [] + + def run_step(): + _pre_run_step(otel_tracker, action) + span = tracer.start_span("llm_call") + span_ids.append(span.get_span_context().span_id) + with pytest.raises(RuntimeError, match="tracker failed"): + span.end() + burr_tracker.post_end_span.side_effect = None + _post_run_step(otel_tracker, action) + + contextvars.copy_context().run(run_step) + + assert span_ids[0] not in span_map + assert len(span_map) == 0 + + +def test_burr_tracking_span_processor_span_after_nested_app_step_reaches_outer_tracker(): + """A sub-application stepped inside an action clears tracker_context in its post_run_step. + Spans the outer action opens afterwards still reach the outer tracker and leave the cache.""" + tracer, outer_tracker, outer_otel, outer_action = _tracker_with_local_provider() + _, inner_tracker, inner_otel, inner_action = _tracker_with_local_provider() + inner_otel.tracer = tracer + + def run_step(): + _pre_run_step(outer_otel, outer_action) + _pre_run_step(inner_otel, inner_action) + _post_run_step(inner_otel, inner_action) + assert tracker_context.get() is None + with tracer.start_as_current_span("after_inner"): + pass + _post_run_step(outer_otel, outer_action) + + contextvars.copy_context().run(run_step) + + assert len(span_map) == 0 + ended = [c.kwargs["span"].name for c in outer_tracker.post_end_span.call_args_list] + assert "after_inner" in ended + assert "after_inner" not in [ + c.kwargs["span"].name for c in inner_tracker.post_end_span.call_args_list + ]