From d44750517c53770738e087c3cebc5d0f49206470 Mon Sep 17 00:00:00 2001 From: Adam Hitchcock Date: Fri, 2 Oct 2026 12:14:23 -0700 Subject: [PATCH] fix(types): parameterize Span's context manager and annotate otel span processor `Span` subclassed a bare `contextlib.AbstractContextManager`, so `with start_span(...) as span` typed `span` as Unknown under pyright strict. `start_span`'s `parent` used a bare `dict`, and `add_braintrust_span_processor` left `tracer_provider` and `custom_filter` unannotated, which made both functions partially unknown to consumers that ship strict type checks. Co-Authored-By: Claude Opus 5.5 (1M context) --- py/src/braintrust/logger.py | 22 +++++++++---------- py/src/braintrust/otel/__init__.py | 12 +++++++--- .../braintrust/type_tests/test_span_types.py | 13 +++++++++++ 3 files changed, 33 insertions(+), 14 deletions(-) create mode 100644 py/src/braintrust/type_tests/test_span_types.py diff --git a/py/src/braintrust/logger.py b/py/src/braintrust/logger.py index 21155a2dd..dde37e4f1 100644 --- a/py/src/braintrust/logger.py +++ b/py/src/braintrust/logger.py @@ -272,7 +272,7 @@ def export(self) -> str: """Return a serialized representation of the object that can be used to start subspans in other places. See `Span.start_span` for more details.""" -class Span(Exportable, contextlib.AbstractContextManager, ABC): +class Span(Exportable, contextlib.AbstractContextManager["Span"], ABC): """ A Span encapsulates logged data and metrics for a unit of work. This interface is shared by all span implementations. @@ -311,7 +311,7 @@ def start_span( span_attributes: SpanAttributes | Mapping[str, Any] | None = None, start_time: float | None = None, set_current: bool | None = None, - parent: str | dict | None = None, + parent: str | dict[str, str] | None = None, internal: SpanInternalOptions | None = None, **event: Any, ) -> "Span": @@ -456,7 +456,7 @@ def start_span( span_attributes: SpanAttributes | Mapping[str, Any] | None = None, start_time: float | None = None, set_current: bool | None = None, - parent: str | dict | None = None, + parent: str | dict[str, str] | None = None, internal: SpanInternalOptions | None = None, **event: Any, ): @@ -1044,7 +1044,7 @@ def flush(self, batch_size: int | None = None): # Track upload attempts (don't actually call upload() in tests) self.upload_attempts.extend(attachments) - def pop(self): + def pop(self) -> list[dict[str, Any]]: with self.lock: logs = [record for item in self.logs if (record := item.get()) is not None] self.logs = [] @@ -2788,7 +2788,7 @@ def inject_trace_context(carrier: dict | None = None, span: "Span | None" = None return carrier -def extract_trace_context(headers: dict) -> dict | None: +def extract_trace_context(headers: dict) -> dict[str, str] | None: """Extract an opaque W3C trace-context from inbound request headers. This is the receive-side counterpart of `Span.inject` / @@ -3075,7 +3075,7 @@ def start_span( span_attributes: SpanAttributes | Mapping[str, Any] | None = None, start_time: float | None = None, set_current: bool | None = None, - parent: str | dict | None = None, + parent: str | dict[str, str] | None = None, propagated_event: dict[str, Any] | None = None, state: BraintrustState | None = None, internal: SpanInternalOptions | None = None, @@ -4519,7 +4519,7 @@ def start_span( span_attributes: SpanAttributes | Mapping[str, Any] | None = None, start_time: float | None = None, set_current: bool | None = None, - parent: str | dict | None = None, + parent: str | dict[str, str] | None = None, propagated_event: dict[str, Any] | None = None, internal: SpanInternalOptions | None = None, **event: Any, @@ -4662,7 +4662,7 @@ def _start_span_impl( span_attributes: SpanAttributes | Mapping[str, Any] | None = None, start_time: float | None = None, set_current: bool | None = None, - parent: str | dict | None = None, + parent: str | dict[str, str] | None = None, propagated_event: dict[str, Any] | None = None, lookup_span_parent: bool = True, internal: SpanInternalOptions | None = None, @@ -5010,7 +5010,7 @@ def start_span( span_attributes: SpanAttributes | Mapping[str, Any] | None = None, start_time: float | None = None, set_current: bool | None = None, - parent: str | dict | None = None, + parent: str | dict[str, str] | None = None, propagated_event: dict[str, Any] | None = None, internal: SpanInternalOptions | None = None, **event: Any, @@ -6195,7 +6195,7 @@ def start_span( span_attributes: SpanAttributes | Mapping[str, Any] | None = None, start_time: float | None = None, set_current: bool | None = None, - parent: str | dict | None = None, + parent: str | dict[str, str] | None = None, propagated_event: dict[str, Any] | None = None, span_id: str | None = None, root_span_id: str | None = None, @@ -6247,7 +6247,7 @@ def _start_span_impl( span_attributes: SpanAttributes | Mapping[str, Any] | None = None, start_time: float | None = None, set_current: bool | None = None, - parent: str | dict | None = None, + parent: str | dict[str, str] | None = None, propagated_event: dict[str, Any] | None = None, span_id: str | None = None, root_span_id: str | None = None, diff --git a/py/src/braintrust/otel/__init__.py b/py/src/braintrust/otel/__init__.py index 85ac21599..2ef361161 100644 --- a/py/src/braintrust/otel/__init__.py +++ b/py/src/braintrust/otel/__init__.py @@ -3,12 +3,18 @@ import os import threading import warnings +from collections.abc import Callable +from typing import TYPE_CHECKING from urllib.parse import urljoin from braintrust.env import BraintrustEnv from braintrust.span_origin import SpanOriginEnvironment, detect_environment, merge_span_origin_context +if TYPE_CHECKING: + from opentelemetry.sdk.trace import ReadableSpan, TracerProvider + + INSTALL_ERR_MSG = ( "OpenTelemetry packages are not installed. " "Install optional OpenTelemetry dependencies with: pip install braintrust[otel]" @@ -310,15 +316,15 @@ def shutdown(self): def add_braintrust_span_processor( - tracer_provider, + tracer_provider: "TracerProvider", api_key: str | None = None, parent: str | None = None, api_url: str | None = None, filter_ai_spans: bool = False, - custom_filter=None, + custom_filter: "Callable[[ReadableSpan], bool | None] | None" = None, headers: dict[str, str] | None = None, environment: SpanOriginEnvironment | None = None, -): +) -> None: processor = BraintrustSpanProcessor( api_key=api_key, parent=parent, diff --git a/py/src/braintrust/type_tests/test_span_types.py b/py/src/braintrust/type_tests/test_span_types.py new file mode 100644 index 000000000..6650dfb13 --- /dev/null +++ b/py/src/braintrust/type_tests/test_span_types.py @@ -0,0 +1,13 @@ +import braintrust +from braintrust.logger import Span +from typing_extensions import assert_type + + +def span_context_manager_keeps_span_type() -> None: + with braintrust.start_span(name="typed") as span: + assert_type(span, Span) + span.log(output="ok") + + +def start_span_accepts_extracted_trace_context(headers: dict[str, str]) -> Span: + return braintrust.start_span(name="child", parent=braintrust.extract_trace_context(headers))