Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 23 additions & 9 deletions burr/integrations/opentelemetry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = {}
Expand All @@ -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


Expand Down Expand Up @@ -412,6 +415,7 @@ def pre_run_step(
),
partition_key=partition_key,
app_id=app_id,
tracker=self.burr_tracker,
),
)

Expand Down Expand Up @@ -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",
Expand All @@ -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,
Expand All @@ -512,19 +523,22 @@ 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,
span_dependencies=[], # TODO -- log
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,
Expand Down
173 changes: 170 additions & 3 deletions tests/integrations/test_burr_opentelemetry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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",
[
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -105,17 +120,20 @@ 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:
# This should not raise an error even though tracker is None
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():
Expand Down Expand Up @@ -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
]
Loading