diff --git a/.codespellrc b/.codespellrc
index e20054be5..8d99d935b 100644
--- a/.codespellrc
+++ b/.codespellrc
@@ -4,4 +4,4 @@
check-hidden = true
# skipping auto generated folders
skip = ./.git,./.tox,./.venv,./.mypy_cache,./docs/_build,./target,*/LICENSE,./venv,*/cassettes
-ignore-words-list = ot
+ignore-words-list = ot,asend
diff --git a/.github/instructions/instrumentation.instructions.md b/.github/instructions/instrumentation.instructions.md
index 98c6664c6..513f637ba 100644
--- a/.github/instructions/instrumentation.instructions.md
+++ b/.github/instructions/instrumentation.instructions.md
@@ -55,7 +55,11 @@ prefer opt-in or additive. Breaking changes need explicit justification in the P
capture path — never as unconditional span/log attributes.
- Adding attributes to invocations produced by the util is fine.
- Streaming responses must be instrumented by subclassing the util's `SyncStreamWrapper` /
- `AsyncStreamWrapper` (`opentelemetry.util.genai.stream`). Flag hand-rolled stream wrappers.
+ `AsyncStreamWrapper` (`opentelemetry.util.genai.stream`). Flag hand-rolled stream wrappers, and
+ invocation-backed wrappers that do not pass the invocation to `super().__init__(stream, invocation)`
+ or do not `invocation.suspend()` before the stream is returned. Flag an `_execution_context()`
+ override on an invocation-backed wrapper: the base wrapper already activates a suspended
+ invocation during reads, and an override bypasses that guard.
- Instrumentation should not change what a call returns or when its work happens. Flag: work the SDK
didn't do (materializing a result early to build telemetry — stay lazy); a changed return type
(`isinstance`/`__class__` should still resolve to the original; `wrapt.ObjectProxy` is the usual
diff --git a/AGENTS.md b/AGENTS.md
index 0f57d49b7..52b5a6189 100644
--- a/AGENTS.md
+++ b/AGENTS.md
@@ -229,13 +229,19 @@ as the reference:
A streamed response only finishes once the caller has drained the stream, so the invocation must
stay open until then. Do **not** call `invocation.stop()` when the SDK returns the stream — the
-span would close before any chunks arrive.
+span would close before any chunks arrive. The invocation's span must also not stay current in
+the caller's context while the stream is unconsumed: call `invocation.suspend()` before returning
+the stream, and re-activate the invocation only while a chunk is being read.
Instrument streams by subclassing `SyncStreamWrapper` / `AsyncStreamWrapper` from
`opentelemetry.util.genai.stream` (the public, supported helpers). The base class proxies the
underlying SDK stream, drives iteration, and finalizes telemetry exactly once on success, error,
-or `close()`. Subclasses pass the SDK stream to `super().__init__(stream)` and implement three
-hooks:
+or `close()`. Subclasses pass the SDK stream and the invocation to
+`super().__init__(stream, invocation)`, which also marks the invocation as a streamed request
+(`gen_ai.request.stream`), call `invocation.suspend()` before returning the stream, and implement
+three hooks. The base `_execution_context()` makes a suspended invocation current again for each
+read and cleanup operation and restores the context before the chunk is returned, so
+invocation-backed wrappers don't override it.
- `_process_chunk(chunk)` — accumulate per-chunk state (e.g. response model, finish reasons,
token usage, streamed content) onto the invocation.
@@ -246,8 +252,8 @@ hooks:
```python
class MyStreamWrapper(SyncStreamWrapper[Chunk]):
def __init__(self, stream, invocation, capture_content):
- super().__init__(stream)
- self._self_invocation = invocation
+ super().__init__(stream, invocation)
+ invocation.suspend()
...
def _process_chunk(self, chunk): ... # accumulate state
diff --git a/instrumentation/opentelemetry-instrumentation-genai-openai/.changelog/817.fixed b/instrumentation/opentelemetry-instrumentation-genai-openai/.changelog/817.fixed
new file mode 100644
index 000000000..ccdb683f8
--- /dev/null
+++ b/instrumentation/opentelemetry-instrumentation-genai-openai/.changelog/817.fixed
@@ -0,0 +1 @@
+Suspend the chat completion and Responses API invocations while a streamed response is unconsumed and re-activate them only while a chunk is read or the stream is cleaned up: the inference span is no longer current in the caller's context between reads, so spans a caller creates between chunks are siblings of the inference span, not children.
diff --git a/instrumentation/opentelemetry-instrumentation-genai-openai/pyproject.toml b/instrumentation/opentelemetry-instrumentation-genai-openai/pyproject.toml
index 7931def5a..c17b50c64 100644
--- a/instrumentation/opentelemetry-instrumentation-genai-openai/pyproject.toml
+++ b/instrumentation/opentelemetry-instrumentation-genai-openai/pyproject.toml
@@ -28,7 +28,7 @@ dependencies = [
"opentelemetry-api ~= 1.43",
"opentelemetry-instrumentation >= 0.64b0, <1",
"opentelemetry-semantic-conventions >= 0.64b0, <1",
- "opentelemetry-util-genai >= 1.2b0, <2",
+ "opentelemetry-util-genai >= 1.3b0.dev, <2",
]
[project.optional-dependencies]
diff --git a/instrumentation/opentelemetry-instrumentation-genai-openai/src/opentelemetry/instrumentation/genai/openai/_raw_response.py b/instrumentation/opentelemetry-instrumentation-genai-openai/src/opentelemetry/instrumentation/genai/openai/_raw_response.py
index 6577e63c1..52150f0c7 100644
--- a/instrumentation/opentelemetry-instrumentation-genai-openai/src/opentelemetry/instrumentation/genai/openai/_raw_response.py
+++ b/instrumentation/opentelemetry-instrumentation-genai-openai/src/opentelemetry/instrumentation/genai/openai/_raw_response.py
@@ -239,6 +239,13 @@ def __call__(
) -> object: ...
+def _stop_inside(invocation: _StreamingInvocation) -> None:
+ # A response closed without parse() ends the span inside it, as a parsed
+ # stream's cleanup does.
+ with invocation.activate():
+ invocation.stop()
+
+
def wrap_stream_result(
wrapper_cls: StreamWrapperFactory,
result: RawResponseLike | AnyStream,
@@ -256,9 +263,13 @@ def wrap_stream_result(
if served_model:
if hasattr(invocation, "response_model_name"):
setattr(invocation, "response_model_name", served_model)
+ # The stream wrapper is built only on ``parse()``, possibly in another
+ # context; suspend now so the span isn't current before then and the
+ # context it detaches is the one ``create()`` ran in.
+ invocation.suspend()
return RawResponseStreamProxy(
result,
lambda stream: wrapper_cls(stream, invocation, capture_content),
- finalize=invocation.stop,
+ finalize=functools.partial(_stop_inside, invocation),
)
return wrapper_cls(result, invocation, capture_content)
diff --git a/instrumentation/opentelemetry-instrumentation-genai-openai/src/opentelemetry/instrumentation/genai/openai/chat_wrappers.py b/instrumentation/opentelemetry-instrumentation-genai-openai/src/opentelemetry/instrumentation/genai/openai/chat_wrappers.py
index f2d456923..ce44e00df 100644
--- a/instrumentation/opentelemetry-instrumentation-genai-openai/src/opentelemetry/instrumentation/genai/openai/chat_wrappers.py
+++ b/instrumentation/opentelemetry-instrumentation-genai-openai/src/opentelemetry/instrumentation/genai/openai/chat_wrappers.py
@@ -212,6 +212,7 @@ def __init__(
) -> None:
super().__init__(stream, invocation=invocation)
self._self_invocation = invocation
+ invocation.suspend()
self._self_choice_buffers = []
self._self_capture_content = capture_content
self._self_response_id = None
@@ -234,6 +235,7 @@ def __init__(
) -> None:
super().__init__(stream, invocation=invocation)
self._self_invocation = invocation
+ invocation.suspend()
self._self_choice_buffers = []
self._self_capture_content = capture_content
self._self_response_id = None
diff --git a/instrumentation/opentelemetry-instrumentation-genai-openai/src/opentelemetry/instrumentation/genai/openai/response_wrappers.py b/instrumentation/opentelemetry-instrumentation-genai-openai/src/opentelemetry/instrumentation/genai/openai/response_wrappers.py
index 9a28930e9..5df18a9a9 100644
--- a/instrumentation/opentelemetry-instrumentation-genai-openai/src/opentelemetry/instrumentation/genai/openai/response_wrappers.py
+++ b/instrumentation/opentelemetry-instrumentation-genai-openai/src/opentelemetry/instrumentation/genai/openai/response_wrappers.py
@@ -7,6 +7,7 @@
import logging
from collections.abc import Callable
+from contextlib import AbstractContextManager
from contextvars import ContextVar
from types import TracebackType
from typing import TYPE_CHECKING, Generic, TypeVar, cast
@@ -123,6 +124,8 @@ class _ResponseStreamMixin(Generic[TextFormatT]):
_self_invocation: _ResponseInvocation
_self_capture_content: bool
_self_response_telemetry_finalized: bool
+ # provided by the stream wrapper base class the mixin is combined with
+ _execution_context: Callable[[], AbstractContextManager[None]]
def __init__(
self,
@@ -132,6 +135,9 @@ def __init__(
self._self_invocation = invocation
self._self_capture_content = capture_content
self._self_response_telemetry_finalized = False
+ # The stream returns to the caller undrained: leave the caller's
+ # context as it was and make the span current only while reading.
+ invocation.suspend()
def _stop(
self, result: ParsedResponse[TextFormatT] | Response | None
@@ -199,7 +205,13 @@ def response(self):
response = _get_stream_response(self.stream)
if response is None:
return None
- return finalize_on_close(response, lambda: self._stop(None))
+ # A close through the HTTP response is stream cleanup too: run it and
+ # the finalizer in the same context as the wrapper's own close.
+ return finalize_on_close(
+ response,
+ lambda: self._stop(None),
+ execution_context=self._execution_context,
+ )
def process_event(self, event: ResponseStreamEvent[TextFormatT]) -> None:
# raw-response stream can be parsed into a caller-defined event type.
@@ -412,7 +424,12 @@ def response(self):
response = _get_stream_response(self.stream)
if response is None:
return None
- return finalize_on_aclose(response, lambda: self._stop(None))
+ # See _ResponseStreamMixin.response.
+ return finalize_on_aclose(
+ response,
+ lambda: self._stop(None),
+ execution_context=self._execution_context,
+ )
class AsyncFetchResponseStreamWrapper(
diff --git a/instrumentation/opentelemetry-instrumentation-genai-openai/tests/requirements.oldest.txt b/instrumentation/opentelemetry-instrumentation-genai-openai/tests/requirements.oldest.txt
index 1d7d06929..cdc6d82c3 100644
--- a/instrumentation/opentelemetry-instrumentation-genai-openai/tests/requirements.oldest.txt
+++ b/instrumentation/opentelemetry-instrumentation-genai-openai/tests/requirements.oldest.txt
@@ -30,3 +30,6 @@ httpx==0.27.2
Deprecated==1.2.14
importlib-metadata==6.11.0
packaging==24.0
+
+# The stream execution hook is not released yet.
+-e util/opentelemetry-util-genai
diff --git a/instrumentation/opentelemetry-instrumentation-genai-openai/tests/test_chat_stream_context.py b/instrumentation/opentelemetry-instrumentation-genai-openai/tests/test_chat_stream_context.py
new file mode 100644
index 000000000..dcdcae129
--- /dev/null
+++ b/instrumentation/opentelemetry-instrumentation-genai-openai/tests/test_chat_stream_context.py
@@ -0,0 +1,91 @@
+# Copyright The OpenTelemetry Authors
+# SPDX-License-Identifier: Apache-2.0
+
+from __future__ import annotations
+
+import asyncio
+from collections.abc import AsyncIterator, Iterator
+
+import pytest
+from openai.types.chat import ChatCompletionChunk
+
+from opentelemetry import context, trace
+from opentelemetry.instrumentation.genai.openai.chat_wrappers import (
+ AsyncChatStreamWrapper,
+ ChatStreamWrapper,
+)
+from opentelemetry.util.genai.handler import TelemetryHandler
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("asynchronous", [False, True])
+@pytest.mark.parametrize("ending", ["drain", "close", "error"])
+async def test_chat_stream_activates_only_during_reads(
+ tracer_provider, span_exporter, caplog, asynchronous: bool, ending: str
+) -> None:
+ before = context.get_current()
+ invocation = TelemetryHandler(
+ tracer_provider=tracer_provider,
+ ).inference("openai", request_model="test-model")
+ error = ConnectionError("stream failed")
+ chunk = ChatCompletionChunk(
+ id="test",
+ created=1,
+ model="test-model",
+ object="chat.completion.chunk",
+ choices=[],
+ )
+
+ def produce() -> Iterator[ChatCompletionChunk]:
+ assert (
+ trace.get_current_span().get_span_context()
+ == trace.get_current_span(invocation.context).get_span_context()
+ )
+ yield chunk
+ assert (
+ trace.get_current_span().get_span_context()
+ == trace.get_current_span(invocation.context).get_span_context()
+ )
+ if ending == "error":
+ raise error
+
+ async def aproduce() -> AsyncIterator[ChatCompletionChunk]:
+ for item in produce():
+ await asyncio.sleep(0)
+ yield item
+
+ if asynchronous:
+ stream = AsyncChatStreamWrapper(aproduce(), invocation, False)
+ else:
+ stream = ChatStreamWrapper(produce(), invocation, False)
+ assert context.get_current() is before
+ if asynchronous:
+ assert await asyncio.create_task(anext(stream)) is chunk
+ else:
+ assert next(stream) is chunk
+ assert context.get_current() is before
+ assert span_exporter.get_finished_spans() == ()
+
+ if ending == "close":
+ if asynchronous:
+ await stream.aclose()
+ else:
+ stream.close()
+ elif ending == "error":
+ with pytest.raises(ConnectionError) as raised:
+ if asynchronous:
+ await anext(stream)
+ else:
+ next(stream)
+ assert raised.value is error
+ elif asynchronous:
+ assert [item async for item in stream] == []
+ else:
+ assert list(stream) == []
+ assert context.get_current() is before
+ assert len(span_exporter.get_finished_spans()) == 1
+ assert not [
+ r
+ for r in caplog.records
+ if r.name == "opentelemetry.context" and r.levelno >= 40
+ ]
diff --git a/instrumentation/opentelemetry-instrumentation-genai-openai/tests/test_raw_response_stream_suspend.py b/instrumentation/opentelemetry-instrumentation-genai-openai/tests/test_raw_response_stream_suspend.py
new file mode 100644
index 000000000..4997a39b9
--- /dev/null
+++ b/instrumentation/opentelemetry-instrumentation-genai-openai/tests/test_raw_response_stream_suspend.py
@@ -0,0 +1,120 @@
+# Copyright The OpenTelemetry Authors
+# SPDX-License-Identifier: Apache-2.0
+
+import asyncio
+
+from openai import AsyncStream, Stream
+
+from opentelemetry import context, trace
+from opentelemetry.instrumentation.genai.openai._raw_response import (
+ wrap_stream_result,
+)
+from opentelemetry.instrumentation.genai.openai.chat_wrappers import (
+ AsyncChatStreamWrapper,
+ ChatStreamWrapper,
+)
+from opentelemetry.util.genai.handler import TelemetryHandler
+
+from .test_raw_response_proxy import _AsyncRawResponse, _RawResponse
+
+
+class _EmptyStream(Stream):
+ def __init__(self):
+ pass
+
+ def __iter__(self):
+ return iter(())
+
+ def close(self):
+ pass
+
+
+class _EmptyAsyncStream(AsyncStream):
+ def __init__(self):
+ pass
+
+ def __aiter__(self):
+ return self
+
+ async def __anext__(self):
+ raise StopAsyncIteration
+
+ async def close(self):
+ pass
+
+
+def _start(tracer_provider):
+ return TelemetryHandler(tracer_provider=tracer_provider).inference(
+ "openai", request_model="m"
+ )
+
+
+def test_raw_response_span_not_current_after_create(tracer_provider):
+ before = context.get_current()
+ raw = wrap_stream_result(
+ ChatStreamWrapper,
+ _RawResponse(_EmptyStream()),
+ _start(tracer_provider),
+ False,
+ )
+ assert context.get_current() is before
+ raw.http_response.close()
+
+
+def test_unparsed_raw_response_stops_inside_its_span(tracer_provider):
+ invocation = _start(tracer_provider)
+ stop = invocation.stop
+ current_at_stop = []
+
+ def recording_stop():
+ current_at_stop.append(trace.get_current_span())
+ stop()
+
+ invocation.stop = recording_stop
+ raw = wrap_stream_result(
+ ChatStreamWrapper, _RawResponse(_EmptyStream()), invocation, False
+ )
+ before = context.get_current()
+ with tracer_provider.get_tracer("t").start_as_current_span("caller"):
+ raw.http_response.close()
+
+ assert current_at_stop == [invocation.span]
+ assert context.get_current() is before
+
+
+def test_raw_response_parse_under_caller_span_then_drain(
+ tracer_provider, span_exporter
+):
+ before = context.get_current()
+ raw = wrap_stream_result(
+ ChatStreamWrapper,
+ _RawResponse(_EmptyStream()),
+ _start(tracer_provider),
+ False,
+ )
+ with tracer_provider.get_tracer("t").start_as_current_span("caller"):
+ stream = raw.parse()
+ assert list(stream) == []
+ assert len(span_exporter.get_finished_spans()) == 2
+ assert context.get_current() is before
+
+
+def test_async_raw_response_parse_under_caller_span_then_drain(
+ tracer_provider, span_exporter
+):
+ async def run():
+ before = context.get_current()
+ raw = wrap_stream_result(
+ AsyncChatStreamWrapper,
+ _AsyncRawResponse(_EmptyAsyncStream()),
+ _start(tracer_provider),
+ False,
+ )
+ assert context.get_current() is before
+ with tracer_provider.get_tracer("t").start_as_current_span("caller"):
+ stream = await raw.parse()
+ assert [chunk async for chunk in stream] == []
+ assert len(span_exporter.get_finished_spans()) == 2
+ assert context.get_current() is before
+
+ asyncio.run(run())
diff --git a/instrumentation/opentelemetry-instrumentation-genai-openai/tests/test_response_wrappers.py b/instrumentation/opentelemetry-instrumentation-genai-openai/tests/test_response_wrappers.py
index e42d6875f..3aa1fb129 100644
--- a/instrumentation/opentelemetry-instrumentation-genai-openai/tests/test_response_wrappers.py
+++ b/instrumentation/opentelemetry-instrumentation-genai-openai/tests/test_response_wrappers.py
@@ -1,6 +1,7 @@
# Copyright The OpenTelemetry Authors
# SPDX-License-Identifier: Apache-2.0
+from contextlib import nullcontext
from types import SimpleNamespace
import pytest
@@ -78,13 +79,22 @@ def _noop_on_stream_chunk(chunk_at):
del chunk_at
+def _fake_invocation(**overrides):
+ """A stand-in for an invocation with the stream wrapper's context hooks."""
+ fields = {
+ "request_model": None,
+ "stop": _noop_stop,
+ "fail": _noop_fail,
+ "_on_stream_chunk": _noop_on_stream_chunk,
+ "suspend": _noop_stop,
+ "activate": nullcontext,
+ }
+ fields.update(overrides)
+ return SimpleNamespace(**fields)
+
+
def _make_wrapper(manager):
- invocation = SimpleNamespace(
- request_model=None,
- stop=_noop_stop,
- _on_stream_chunk=_noop_on_stream_chunk,
- fail=_noop_fail,
- )
+ invocation = _fake_invocation()
return ResponseStreamManagerWrapper(
manager=manager,
invocation_factory=lambda: invocation,
@@ -94,12 +104,7 @@ def _make_wrapper(manager):
def _make_stream_wrapper(stream, invocation=None):
if invocation is None:
- invocation = SimpleNamespace(
- request_model=None,
- stop=_noop_stop,
- fail=_noop_fail,
- _on_stream_chunk=_noop_on_stream_chunk,
- )
+ invocation = _fake_invocation()
return ResponseStreamWrapper(
stream=stream,
invocation=invocation,
@@ -108,12 +113,7 @@ def _make_stream_wrapper(stream, invocation=None):
def _make_async_manager_wrapper(manager):
- invocation = SimpleNamespace(
- request_model=None,
- stop=_noop_stop,
- _on_stream_chunk=_noop_on_stream_chunk,
- fail=_noop_fail,
- )
+ invocation = _fake_invocation()
return AsyncResponseStreamManagerWrapper(
manager=manager,
invocation_factory=lambda: invocation,
@@ -123,12 +123,7 @@ def _make_async_manager_wrapper(manager):
def _make_async_stream_wrapper(stream, invocation=None):
if invocation is None:
- invocation = SimpleNamespace(
- request_model=None,
- stop=_noop_stop,
- fail=_noop_fail,
- _on_stream_chunk=_noop_on_stream_chunk,
- )
+ invocation = _fake_invocation()
return AsyncResponseStreamWrapper(
stream=stream,
invocation=invocation,
@@ -205,12 +200,7 @@ def test_manager_enter_failure_fails_invocation_and_reraises():
error = RuntimeError("enter failure")
manager = _FakeManager(stream=SimpleNamespace(), enter_error=error)
failures = []
- invocation = SimpleNamespace(
- request_model=None,
- stop=_noop_stop,
- _on_stream_chunk=_noop_on_stream_chunk,
- fail=failures.append,
- )
+ invocation = _fake_invocation(fail=failures.append)
wrapper = ResponseStreamManagerWrapper(
manager=manager,
invocation_factory=lambda: invocation,
@@ -272,12 +262,7 @@ async def test_async_manager_enter_failure_fails_invocation_and_reraises():
error = RuntimeError("enter failure")
manager = _FakeAsyncManager(stream=SimpleNamespace(), enter_error=error)
failures = []
- invocation = SimpleNamespace(
- request_model=None,
- stop=_noop_stop,
- _on_stream_chunk=_noop_on_stream_chunk,
- fail=failures.append,
- )
+ invocation = _fake_invocation(fail=failures.append)
wrapper = AsyncResponseStreamManagerWrapper(
manager=manager,
invocation_factory=lambda: invocation,
@@ -310,9 +295,7 @@ async def test_async_stream_wrapper_exit_fails_and_closes_on_exception():
stream = _FakeAsyncStream()
stopped = []
failures = []
- invocation = SimpleNamespace(
- request_model=None, stop=_noop_stop, fail=failures.append
- )
+ invocation = _fake_invocation(fail=failures.append)
wrapper = _make_async_stream_wrapper(stream, invocation=invocation)
wrapper._stop = stopped.append
@@ -387,9 +370,7 @@ async def test_async_stream_wrapper_fails_and_reraises_stream_errors():
error = ValueError("boom")
stream = _FakeAsyncStream(error=error)
failures = []
- invocation = SimpleNamespace(
- request_model=None, stop=_noop_stop, fail=failures.append
- )
+ invocation = _fake_invocation(fail=failures.append)
wrapper = _make_async_stream_wrapper(stream, invocation=invocation)
with pytest.raises(ValueError, match="boom"):
@@ -458,9 +439,7 @@ def _make_response(**overrides):
def _capturing_invocation():
calls = {"stop": 0, "fail": []}
- invocation = SimpleNamespace(
- request_model=None, response_model_name=None, attributes={}
- )
+ invocation = _fake_invocation(response_model_name=None, attributes={})
invocation.stop = lambda: calls.__setitem__("stop", calls["stop"] + 1)
invocation.fail = lambda error: calls["fail"].append(error)
return invocation, calls
@@ -562,7 +541,7 @@ def _make_fetch_stream_wrapper(stream, invocation):
def _capturing_fetch_invocation():
calls = {"stop": 0, "fail": []}
- invocation = SimpleNamespace(
+ invocation = _fake_invocation(
response_model_name=None,
response_status=None,
finish_reasons=None,
diff --git a/instrumentation/opentelemetry-instrumentation-genai-openai/tests/test_responses_stream_context.py b/instrumentation/opentelemetry-instrumentation-genai-openai/tests/test_responses_stream_context.py
new file mode 100644
index 000000000..0bdac9fba
--- /dev/null
+++ b/instrumentation/opentelemetry-instrumentation-genai-openai/tests/test_responses_stream_context.py
@@ -0,0 +1,246 @@
+# Copyright The OpenTelemetry Authors
+# SPDX-License-Identifier: Apache-2.0
+
+from __future__ import annotations
+
+import asyncio
+from collections.abc import AsyncIterator, Iterator
+from types import SimpleNamespace
+from typing import Any
+
+import pytest
+
+from opentelemetry import context, trace
+from opentelemetry.instrumentation.genai.openai.response_wrappers import (
+ AsyncResponseStreamManagerWrapper,
+ AsyncResponseStreamWrapper,
+ ResponseStreamManagerWrapper,
+ ResponseStreamWrapper,
+)
+from opentelemetry.util.genai.handler import TelemetryHandler
+
+
+def _assert_no_detach_errors(caplog) -> None:
+ assert not [
+ r
+ for r in caplog.records
+ if r.name == "opentelemetry.context" and r.levelno >= 40
+ ]
+
+
+class _Manager:
+ def __init__(self, stream: Any) -> None:
+ self._stream = stream
+
+ def __enter__(self) -> Any:
+ return self._stream
+
+ def __exit__(self, *exc: Any) -> bool:
+ return False
+
+ async def __aenter__(self) -> Any:
+ return self._stream
+
+ async def __aexit__(self, *exc: Any) -> bool:
+ return False
+
+
+def _events(
+ invocation: Any, event: Any, error: BaseException | None
+) -> Iterator[Any]:
+ assert (
+ trace.get_current_span().get_span_context()
+ == trace.get_current_span(invocation.context).get_span_context()
+ )
+ yield event
+ assert (
+ trace.get_current_span().get_span_context()
+ == trace.get_current_span(invocation.context).get_span_context()
+ )
+ if error is not None:
+ raise error
+
+
+async def _aevents(
+ invocation: Any, event: Any, error: BaseException | None
+) -> AsyncIterator[Any]:
+ for item in _events(invocation, event, error):
+ await asyncio.sleep(0)
+ yield item
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("asynchronous", [False, True])
+@pytest.mark.parametrize("ending", ["drain", "close", "error"])
+async def test_responses_stream_activates_only_during_reads(
+ tracer_provider, span_exporter, caplog, asynchronous: bool, ending: str
+) -> None:
+ before = context.get_current()
+ invocation = TelemetryHandler(tracer_provider=tracer_provider).inference(
+ "openai", request_model="test-model"
+ )
+ error = ConnectionError("stream failed") if ending == "error" else None
+ event = SimpleNamespace(type="response.output_text.delta")
+
+ if asynchronous:
+ stream = AsyncResponseStreamWrapper(
+ _aevents(invocation, event, error), invocation, False
+ )
+ else:
+ stream = ResponseStreamWrapper(
+ _events(invocation, event, error), invocation, False
+ )
+ assert context.get_current() is before
+ if asynchronous:
+ assert await asyncio.create_task(anext(stream)) is event
+ else:
+ assert next(stream) is event
+ assert context.get_current() is before
+ assert span_exporter.get_finished_spans() == ()
+
+ async def drain() -> list[Any]:
+ return [item async for item in stream]
+
+ if ending == "close":
+ if asynchronous:
+ await stream.aclose()
+ else:
+ stream.close()
+ elif ending == "error":
+ with pytest.raises(ConnectionError) as raised:
+ if asynchronous:
+ await asyncio.create_task(anext(stream))
+ else:
+ next(stream)
+ assert raised.value is error
+ elif asynchronous:
+ assert await asyncio.create_task(drain()) == []
+ else:
+ assert list(stream) == []
+ assert context.get_current() is before
+ assert len(span_exporter.get_finished_spans()) == 1
+ _assert_no_detach_errors(caplog)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("asynchronous", [False, True])
+async def test_responses_stream_manager_activates_only_during_reads(
+ tracer_provider, span_exporter, caplog, asynchronous: bool
+) -> None:
+ before = context.get_current()
+ handler = TelemetryHandler(tracer_provider=tracer_provider)
+ invocations: list[Any] = []
+ event = SimpleNamespace(type="response.output_text.delta")
+
+ def factory() -> Any:
+ invocations.append(
+ handler.inference("openai", request_model="test-model")
+ )
+ return invocations[-1]
+
+ def produce() -> Any:
+ return (
+ _aevents(invocations[-1], event, None)
+ if asynchronous
+ else _events(invocations[-1], event, None)
+ )
+
+ class Manager(_Manager):
+ def __init__(self) -> None:
+ super().__init__(None)
+
+ def __enter__(self) -> Any:
+ return produce()
+
+ async def __aenter__(self) -> Any:
+ return produce()
+
+ if asynchronous:
+ manager = AsyncResponseStreamManagerWrapper(Manager(), factory, False)
+ async with manager as stream:
+ assert context.get_current() is before
+ assert await asyncio.create_task(anext(stream)) is event
+ assert context.get_current() is before
+ assert span_exporter.get_finished_spans() == ()
+ else:
+ manager = ResponseStreamManagerWrapper(Manager(), factory, False)
+ with manager as stream:
+ assert context.get_current() is before
+ assert next(stream) is event
+ assert context.get_current() is before
+ assert span_exporter.get_finished_spans() == ()
+ assert context.get_current() is before
+ assert len(span_exporter.get_finished_spans()) == 1
+ _assert_no_detach_errors(caplog)
+
+
+class _RecordingHook:
+ def __init__(self) -> None:
+ self.seen: list[Any] = []
+
+ def on_completion(self, **kwargs: Any) -> None:
+ self.seen.append(trace.get_current_span().get_span_context())
+
+
+class _Response:
+ def __init__(self) -> None:
+ self.closed_in: list[Any] = []
+
+ def close(self) -> None:
+ self.closed_in.append(trace.get_current_span().get_span_context())
+
+ async def aclose(self) -> None:
+ self.close()
+
+
+class _Stream:
+ def __init__(self, response: _Response) -> None:
+ self._response = response
+
+ def __iter__(self) -> Iterator[Any]:
+ return iter(())
+
+ def __aiter__(self) -> _Stream:
+ return self
+
+ async def __anext__(self) -> Any:
+ raise StopAsyncIteration
+
+ def close(self) -> None:
+ pass
+
+ async def aclose(self) -> None:
+ pass
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("asynchronous", [False, True])
+async def test_responses_stream_response_close_activates_invocation(
+ tracer_provider, span_exporter, caplog, asynchronous: bool
+) -> None:
+ before = context.get_current()
+ hook = _RecordingHook()
+ invocation = TelemetryHandler(
+ tracer_provider=tracer_provider, completion_hook=hook
+ ).inference("openai", request_model="test-model")
+ expected = trace.get_current_span(invocation.context).get_span_context()
+ response = _Response()
+
+ if asynchronous:
+ stream = AsyncResponseStreamWrapper(
+ _Stream(response), invocation, False
+ )
+ else:
+ stream = ResponseStreamWrapper(_Stream(response), invocation, False)
+ assert context.get_current() is before
+
+ if asynchronous:
+ await stream.response.aclose()
+ else:
+ stream.response.close()
+
+ assert context.get_current() is before
+ assert response.closed_in == [expected]
+ assert hook.seen == [expected]
+ assert len(span_exporter.get_finished_spans()) == 1
+ _assert_no_detach_errors(caplog)
diff --git a/util/opentelemetry-util-genai/.changelog/817.added b/util/opentelemetry-util-genai/.changelog/817.added
new file mode 100644
index 000000000..2e4f3bc1d
--- /dev/null
+++ b/util/opentelemetry-util-genai/.changelog/817.added
@@ -0,0 +1 @@
+Stream wrappers make a suspended invocation current only while a chunk is read or the stream is cleaned up.
diff --git a/util/opentelemetry-util-genai/AGENTS.md b/util/opentelemetry-util-genai/AGENTS.md
index f336c224c..091248c4b 100644
--- a/util/opentelemetry-util-genai/AGENTS.md
+++ b/util/opentelemetry-util-genai/AGENTS.md
@@ -70,13 +70,19 @@ is usually hardcoded in specific invocation and does not need to be passed.
A streamed response only finishes once the caller has drained the stream, so the invocation must
stay open until then. Do **not** call `invocation.stop()` when the SDK returns the stream — the
-span would close before any chunks arrive.
+span would close before any chunks arrive. The invocation's span must also not stay current in
+the caller's context while the stream is unconsumed: call `invocation.suspend()` before returning
+the stream, and re-activate the invocation only while a chunk is being read.
Instrument streams by subclassing `SyncStreamWrapper` / `AsyncStreamWrapper` from
`opentelemetry.util.genai.stream` (the public, supported helpers). The base class proxies the
underlying SDK stream, drives iteration, and finalizes telemetry exactly once on success, error,
-or `close()`. Subclasses pass the SDK stream to `super().__init__(stream)` and implement three
-hooks:
+or `close()`. Subclasses pass the SDK stream and the invocation to
+`super().__init__(stream, invocation)`, which also marks the invocation as a streamed request
+(`gen_ai.request.stream`), call `invocation.suspend()` before returning the stream, and implement
+three hooks. The base `_execution_context()` makes a suspended invocation current again for each
+read and cleanup operation and restores the context before the chunk is returned, so
+invocation-backed wrappers don't override it.
- `_process_chunk(chunk)` — accumulate per-chunk state (e.g. response model, finish reasons,
token usage, streamed content) onto the invocation.
@@ -87,8 +93,8 @@ hooks:
```python
class MyStreamWrapper(SyncStreamWrapper[Chunk]):
def __init__(self, stream, invocation, capture_content):
- super().__init__(stream)
- self._self_invocation = invocation
+ super().__init__(stream, invocation)
+ invocation.suspend()
...
def _process_chunk(self, chunk): ... # accumulate state
diff --git a/util/opentelemetry-util-genai/README.rst b/util/opentelemetry-util-genai/README.rst
index 729621246..6e61193f1 100644
--- a/util/opentelemetry-util-genai/README.rst
+++ b/util/opentelemetry-util-genai/README.rst
@@ -35,6 +35,22 @@ to manage context:
ambient context is used.
+Stream execution context
+------------------------
+
+``SyncStreamWrapper`` and ``AsyncStreamWrapper`` scope each stream read and
+cleanup operation with ``_execution_context()``. Context is restored before
+returning a chunk to the consumer. Generator ``send``/``throw`` and
+``asend``/``athrow`` also use this scope when the underlying stream supports them.
+
+For invocation-backed streams, pass the invocation to the base constructor
+(``super().__init__(stream, invocation)``) and call ``invocation.suspend()``
+before returning the stream. The base wrapper then makes the suspended invocation
+current during reads, cleanup and finalization of an abandoned stream, so no
+``_execution_context()`` override is needed.
+Keep finalization in the existing ``_on_stream_end`` and ``_on_stream_error`` hooks.
+
+
Modalities
----------
@@ -209,4 +225,4 @@ References
----------
* `OpenTelemetry Project `_
-* `OpenTelemetry GenAI semantic conventions `_
\ No newline at end of file
+* `OpenTelemetry GenAI semantic conventions `_
diff --git a/util/opentelemetry-util-genai/src/opentelemetry/util/genai/_invocation.py b/util/opentelemetry-util-genai/src/opentelemetry/util/genai/_invocation.py
index 12645c105..17ed37e98 100644
--- a/util/opentelemetry-util-genai/src/opentelemetry/util/genai/_invocation.py
+++ b/util/opentelemetry-util-genai/src/opentelemetry/util/genai/_invocation.py
@@ -154,6 +154,7 @@ def __init__(
context=ctx,
)
self._span_context: Context = ctx
+ self._attach_to_context = _attach_to_context
self._context_token: ContextToken | None = (
attach(self._span_context) if _attach_to_context else None
)
@@ -161,6 +162,8 @@ def __init__(
ctx = get_current() if context is None else context
self.span = get_current_span(context=ctx)
self._span_context = ctx
+ # Never attached, so never treated as suspended by the stream wrappers.
+ self._attach_to_context = False
self._context_token = None
self._monotonic_start_s: float = timeit.default_timer()
# Streaming state, set when the invocation is handed to a stream
@@ -193,6 +196,19 @@ def context(self) -> Context:
"""The OpenTelemetry Context containing this invocation's span."""
return self._span_context
+ @property
+ def _suspended(self) -> bool:
+ """True when this invocation attached itself to the context and ``suspend`` has since run.
+
+ Describes the original attachment only: it stays true during a
+ temporary ``activate`` and after the invocation has finished. The
+ stream wrappers activate such an invocation around each read. An
+ invocation that is still attached is left alone: its original
+ attachment is outstanding, and activating it around the read that
+ finishes it would reinstate the ended span after ``stop`` detached it.
+ """
+ return self._attach_to_context and self._context_token is None
+
def suspend(self) -> None:
"""Restore the context that was current before this invocation started.
diff --git a/util/opentelemetry-util-genai/src/opentelemetry/util/genai/_tool_invocation.py b/util/opentelemetry-util-genai/src/opentelemetry/util/genai/_tool_invocation.py
index a903d0143..5d08eff49 100644
--- a/util/opentelemetry-util-genai/src/opentelemetry/util/genai/_tool_invocation.py
+++ b/util/opentelemetry-util-genai/src/opentelemetry/util/genai/_tool_invocation.py
@@ -161,3 +161,6 @@ def _record_metrics(self) -> None:
attributes=self._get_metric_attributes(),
context=self._span_context,
)
+
+ def _on_stream_chunk(self, chunk_at: float) -> None:
+ """A streamed tool result only scopes context; the conventions define no chunk timing metrics for tools."""
diff --git a/util/opentelemetry-util-genai/src/opentelemetry/util/genai/stream.py b/util/opentelemetry-util-genai/src/opentelemetry/util/genai/stream.py
index 7ac7530a0..36c4130ed 100644
--- a/util/opentelemetry-util-genai/src/opentelemetry/util/genai/stream.py
+++ b/util/opentelemetry-util-genai/src/opentelemetry/util/genai/stream.py
@@ -7,7 +7,9 @@
import logging
import timeit
from abc import ABCMeta, abstractmethod
-from collections.abc import AsyncIterable, Callable, Iterable
+from collections.abc import AsyncIterable, Awaitable, Callable, Iterable
+from contextlib import AbstractContextManager, nullcontext
+from functools import wraps
from types import TracebackType
from typing import (
TYPE_CHECKING,
@@ -50,6 +52,8 @@ def __init__(self, wrapped: object) -> None: ...
class _StreamTimingInvocation(Protocol):
def _on_stream_chunk(self, chunk_at: float) -> None: ...
+ def activate(self) -> AbstractContextManager[None]: ...
+
class _StreamingInvocation(_StreamTimingInvocation, Protocol):
def fail(self, error: Error | BaseException) -> None: ...
@@ -91,6 +95,20 @@ class _StreamTelemetry(Generic[ChunkT], metaclass=ABCMeta):
_self_finalized: bool
+ def _execution_context(self) -> AbstractContextManager[None]:
+ """Scope stream reads and cleanup, restoring context before returning a chunk.
+
+ By default a suspended invocation is made current again for each read:
+ an instrumentation that hands the stream back to the caller calls
+ ``invocation.suspend()`` and gets the span restored around tool or
+ provider code without overriding this hook. Invocations that are still
+ attached, or that never attached, leave the context alone.
+ """
+ invocation = getattr(self, "_self_invocation", None)
+ if invocation is None or not getattr(invocation, "_suspended", False):
+ return nullcontext()
+ return invocation.activate()
+
def _finalize_success(self) -> None:
if self._self_finalized:
return
@@ -116,7 +134,8 @@ def __del__(self) -> None:
if getattr(self, "_self_finalized", True):
return
try:
- self._finalize_failure(AbandonedStreamError())
+ with self._execution_context():
+ self._finalize_failure(AbandonedStreamError())
except BaseException: # pylint: disable=broad-exception-caught
_logger.debug(
"GenAI stream finalization error for abandoned stream",
@@ -146,9 +165,11 @@ class SyncStreamWrapper(
Subclass this when wrapping a provider SDK stream that is consumed with
normal iteration. The subclass should pass the SDK stream to
- ``super().__init__(stream)`` and implement the three telemetry hooks:
- ``_process_chunk`` for per-chunk state, ``_on_stream_end`` for successful
- finalization, and ``_on_stream_error`` for failure finalization.
+ ``super().__init__(stream, invocation)`` and implement the three telemetry
+ hooks: ``_process_chunk`` for per-chunk state, ``_on_stream_end`` for
+ successful finalization, and ``_on_stream_error`` for failure finalization.
+ ``_execution_context`` scopes each read and cleanup; by default it makes a
+ suspended invocation current again and can be overridden.
Users should consume subclasses as normal streams, for example with
``for chunk in wrapper`` or ``with wrapper``. The hook methods are called
@@ -194,27 +215,29 @@ def __exit__(
exc_val: BaseException | None,
exc_tb: TracebackType | None,
) -> Literal[False]:
- if exc_val is not None:
- self._finalize_failure(exc_val)
- try:
- self._self_stream.close()
- except Exception: # pylint: disable=broad-exception-caught
- _logger.debug(
- "GenAI stream close error after user exception",
- exc_info=True,
- )
+ with self._execution_context():
+ if exc_val is not None:
+ self._finalize_failure(exc_val)
+ try:
+ self._self_stream.close()
+ except Exception: # pylint: disable=broad-exception-caught
+ _logger.debug(
+ "GenAI stream close error after user exception",
+ exc_info=True,
+ )
+ return False
+
+ self.close()
return False
- self.close()
- return False
-
def close(self) -> None:
- try:
- self._self_stream.close()
- except BaseException as error:
- self._finalize_failure(error)
- raise
- self._finalize_success()
+ with self._execution_context():
+ try:
+ self._self_stream.close()
+ except BaseException as error:
+ self._finalize_failure(error)
+ raise
+ self._finalize_success()
def __iter__(self):
# Override ``ObjectProxy.__iter__`` so iteration drives ``__next__``
@@ -223,21 +246,40 @@ def __iter__(self):
return self
def __next__(self) -> ChunkT:
- try:
- chunk = next(self._self_iterator)
- except StopIteration:
- self._finalize_success()
- raise
- except BaseException as error:
- self._finalize_failure(error)
- raise
- invocation = self._self_invocation
- chunk_at = timeit.default_timer() if invocation is not None else None
- self._process_chunk(chunk)
- # Record after _process_chunk so response.model is on the metrics.
- if invocation is not None and chunk_at is not None:
- invocation._on_stream_chunk(chunk_at)
- return chunk
+ return self._advance(lambda: next(self._self_iterator))
+
+ if not TYPE_CHECKING:
+
+ def __getattr__(self, name: str) -> Any:
+ method = getattr(self.__wrapped__, name)
+ if name not in ("send", "throw"):
+ return method
+
+ @wraps(method)
+ def advance(*args: Any, **kwargs: Any) -> ChunkT:
+ return self._advance(lambda: method(*args, **kwargs))
+
+ return advance
+
+ def _advance(self, read: Callable[[], ChunkT]) -> ChunkT:
+ with self._execution_context():
+ try:
+ chunk = read()
+ except StopIteration:
+ self._finalize_success()
+ raise
+ except BaseException as error:
+ self._finalize_failure(error)
+ raise
+ invocation = self._self_invocation
+ chunk_at = (
+ timeit.default_timer() if invocation is not None else None
+ )
+ self._process_chunk(chunk)
+ # Record after _process_chunk so response.model is on the metrics.
+ if invocation is not None and chunk_at is not None:
+ invocation._on_stream_chunk(chunk_at)
+ return chunk
class AsyncStreamWrapper(
@@ -250,9 +292,11 @@ class AsyncStreamWrapper(
Subclass this when wrapping a provider SDK stream that is consumed with
async iteration. The subclass should pass the SDK stream to
- ``super().__init__(stream)`` and implement the three telemetry hooks:
- ``_process_chunk`` for per-chunk state, ``_on_stream_end`` for successful
- finalization, and ``_on_stream_error`` for failure finalization.
+ ``super().__init__(stream, invocation)`` and implement the three telemetry
+ hooks: ``_process_chunk`` for per-chunk state, ``_on_stream_end`` for
+ successful finalization, and ``_on_stream_error`` for failure finalization.
+ ``_execution_context`` scopes each read and cleanup; by default it makes a
+ suspended invocation current again and can be overridden.
Users should consume subclasses as normal async streams, for example with
``async for chunk in wrapper`` or ``async with wrapper``. The hook methods
@@ -296,20 +340,21 @@ async def __aexit__(
exc_val: BaseException | None,
exc_tb: TracebackType | None,
) -> Literal[False]:
- if exc_val is not None:
- self._finalize_failure(exc_val)
- try:
- await self._close_stream()
- except Exception: # pylint: disable=broad-exception-caught
- _logger.debug(
- "GenAI stream close error after user exception",
- exc_info=True,
- )
+ with self._execution_context():
+ if exc_val is not None:
+ self._finalize_failure(exc_val)
+ try:
+ await self._close_stream()
+ except Exception: # pylint: disable=broad-exception-caught
+ _logger.debug(
+ "GenAI stream close error after user exception",
+ exc_info=True,
+ )
+ return False
+
+ await self._close()
return False
- await self._close()
- return False
-
async def _close_stream(self) -> None:
"""Close the wrapped stream, whichever close method it exposes.
@@ -339,45 +384,48 @@ async def _close(self) -> None:
Reached through ``aclose`` or an async ``close`` on the wrapped stream;
see ``__getattr__``.
"""
- try:
- await self._close_stream()
- except BaseException as error:
- self._finalize_failure(error)
- _logger.debug(
- "GenAI stream close error during close",
- exc_info=True,
- )
- raise
- self._finalize_success()
+ with self._execution_context():
+ try:
+ await self._close_stream()
+ except BaseException as error:
+ self._finalize_failure(error)
+ _logger.debug(
+ "GenAI stream close error during close",
+ exc_info=True,
+ )
+ raise
+ self._finalize_success()
async def _await_close(self, close_awaitable: Any) -> Any:
- try:
- res = await close_awaitable
- except BaseException as error:
- self._finalize_failure(error)
- _logger.debug(
- "GenAI stream close error during close",
- exc_info=True,
- )
- raise
- self._finalize_success()
- return res
+ with self._execution_context():
+ try:
+ res = await close_awaitable
+ except BaseException as error:
+ self._finalize_failure(error)
+ _logger.debug(
+ "GenAI stream close error during close",
+ exc_info=True,
+ )
+ raise
+ self._finalize_success()
+ return res
def _sync_close(self) -> Any:
"""Close a stream exposing a synchronous ``close`` and finalize telemetry."""
- try:
- res = self._self_stream.close()
- except BaseException as error:
- self._finalize_failure(error)
- _logger.debug(
- "GenAI stream close error during close",
- exc_info=True,
- )
- raise
- if inspect.isawaitable(res):
- return self._await_close(res)
- self._finalize_success()
- return res
+ with self._execution_context():
+ try:
+ res = self._self_stream.close()
+ except BaseException as error:
+ self._finalize_failure(error)
+ _logger.debug(
+ "GenAI stream close error during close",
+ exc_info=True,
+ )
+ raise
+ if inspect.isawaitable(res):
+ return self._await_close(res)
+ self._finalize_success()
+ return res
if TYPE_CHECKING:
# Declared for type checkers only. Defining them for real would make
@@ -407,7 +455,15 @@ def __getattr__(self, name):
if inspect.iscoroutinefunction(getattr(wrapped, name)):
return self._close
return self._sync_close
- return getattr(wrapped, name)
+ method = getattr(wrapped, name)
+ if name in ("asend", "athrow"):
+
+ @wraps(method)
+ async def advance(*args: Any, **kwargs: Any) -> ChunkT:
+ return await self._advance(lambda: method(*args, **kwargs))
+
+ return advance
+ return method
def __aiter__(self):
# Override ``ObjectProxy.__aiter__`` so iteration drives ``__anext__``
@@ -416,22 +472,28 @@ def __aiter__(self):
return self
async def __anext__(self) -> ChunkT:
- try:
- chunk = await anext(self._self_aiter)
- except StopAsyncIteration:
- self._finalize_success()
- raise
- except BaseException as error:
- self._finalize_failure(error)
- raise
+ return await self._advance(lambda: anext(self._self_aiter))
- invocation = self._self_invocation
- chunk_at = timeit.default_timer() if invocation is not None else None
- self._process_chunk(chunk)
- # Record after _process_chunk so response.model is on the metrics.
- if invocation is not None and chunk_at is not None:
- invocation._on_stream_chunk(chunk_at)
- return chunk
+ async def _advance(self, read: Callable[[], Awaitable[ChunkT]]) -> ChunkT:
+ with self._execution_context():
+ try:
+ chunk = await read()
+ except StopAsyncIteration:
+ self._finalize_success()
+ raise
+ except BaseException as error:
+ self._finalize_failure(error)
+ raise
+
+ invocation = self._self_invocation
+ chunk_at = (
+ timeit.default_timer() if invocation is not None else None
+ )
+ self._process_chunk(chunk)
+ # Record after _process_chunk so response.model is on the metrics.
+ if invocation is not None and chunk_at is not None:
+ invocation._on_stream_chunk(chunk_at)
+ return chunk
class SyncToolStreamWrapper(SyncStreamWrapper[ChunkT]):
@@ -441,7 +503,9 @@ class SyncToolStreamWrapper(SyncStreamWrapper[ChunkT]):
drained. This wrapper restores the caller's context before returning, and
makes the tool span current only while tool code runs -- producing a chunk,
closing, or finalizing -- so caller work between chunks is not parented
- under the tool. Per-chunk content is accumulated and set on
+ under the tool. This applies to invocations created attached (the default);
+ a tool invocation created with ``_attach_to_context=False`` is left alone
+ and its span is never made current by the wrapper. Per-chunk content is accumulated and set on
``invocation.tool_result`` upon completion.
"""
@@ -450,33 +514,17 @@ def __init__(
stream: _SyncStream[ChunkT],
invocation: ToolInvocation,
) -> None:
- super().__init__(stream)
+ super().__init__(stream, invocation)
self._self_tool_invocation = invocation
invocation.suspend()
self._self_chunks: list[Any] = []
- def __next__(self) -> ChunkT:
- with self._self_tool_invocation.activate():
- return super().__next__()
-
- def close(self) -> None:
- with self._self_tool_invocation.activate():
- super().close()
-
- def __exit__(
- self,
- exc_type: type[BaseException] | None,
- exc_val: BaseException | None,
- exc_tb: TracebackType | None,
- ) -> Literal[False]:
- with self._self_tool_invocation.activate():
- return super().__exit__(exc_type, exc_val, exc_tb)
-
def __del__(self) -> None:
try:
- self._finalize_failure(
- GeneratorExit("Stream garbage collected before completion")
- )
+ with self._execution_context():
+ self._finalize_failure(
+ GeneratorExit("Stream garbage collected before completion")
+ )
except BaseException: # pylint: disable=broad-exception-caught
pass
@@ -510,36 +558,20 @@ def __init__(
stream: _AsyncStream[ChunkT],
invocation: ToolInvocation,
) -> None:
- super().__init__(stream)
+ super().__init__(stream, invocation)
self._self_tool_invocation = invocation
invocation.suspend()
self._self_chunks: list[Any] = []
def __del__(self) -> None:
try:
- self._finalize_failure(
- GeneratorExit("Stream garbage collected before completion")
- )
+ with self._execution_context():
+ self._finalize_failure(
+ GeneratorExit("Stream garbage collected before completion")
+ )
except BaseException: # pylint: disable=broad-exception-caught
pass
- async def __anext__(self) -> ChunkT:
- with self._self_tool_invocation.activate():
- return await super().__anext__()
-
- async def _close(self) -> None:
- with self._self_tool_invocation.activate():
- await super()._close()
-
- async def __aexit__(
- self,
- exc_type: type[BaseException] | None,
- exc_val: BaseException | None,
- exc_tb: TracebackType | None,
- ) -> Literal[False]:
- with self._self_tool_invocation.activate():
- return await super().__aexit__(exc_type, exc_val, exc_tb)
-
def _process_chunk(self, chunk: ChunkT) -> None:
if self._self_tool_invocation.should_capture_content:
self._self_chunks.append(chunk)
@@ -558,32 +590,46 @@ def _on_stream_error(self, error: BaseException) -> None:
self._self_tool_invocation.fail(error)
+_ExecutionContextFactory = Callable[[], AbstractContextManager[None]]
+
+
class _CloseFinalizingProxy(_ObjectProxy):
- def __init__(self, wrapped: object, finalize: Callable[[], None]) -> None:
+ def __init__(
+ self,
+ wrapped: object,
+ finalize: Callable[[], None],
+ execution_context: _ExecutionContextFactory | None,
+ ) -> None:
super().__init__(wrapped)
self._self_finalize = finalize
+ self._self_execution_context = execution_context
- def close(self) -> None:
- try:
- self.__wrapped__.close()
- finally:
- self._self_finalize()
+ def _scope(self) -> AbstractContextManager[None]:
+ if self._self_execution_context is None:
+ return nullcontext()
+ return self._self_execution_context()
+ def close(self) -> None:
+ with self._scope():
+ try:
+ self.__wrapped__.close()
+ finally:
+ self._self_finalize()
-class _AcloseFinalizingProxy(_ObjectProxy):
- def __init__(self, wrapped: object, finalize: Callable[[], None]) -> None:
- super().__init__(wrapped)
- self._self_finalize = finalize
+class _AcloseFinalizingProxy(_CloseFinalizingProxy):
async def aclose(self) -> None:
- try:
- await self.__wrapped__.aclose()
- finally:
- self._self_finalize()
+ with self._scope():
+ try:
+ await self.__wrapped__.aclose()
+ finally:
+ self._self_finalize()
def finalize_on_close(
- wrapped: WrappedT, finalize: Callable[[], None]
+ wrapped: WrappedT,
+ finalize: Callable[[], None],
+ execution_context: _ExecutionContextFactory | None = None,
) -> WrappedT:
"""Proxy ``wrapped`` so closing it also finalizes telemetry.
@@ -592,15 +638,26 @@ def finalize_on_close(
``stream.response`` -- where a ``close()`` means the caller is done and the
invocation should be finalized. Everything but ``close`` forwards
unchanged.
+
+ ``execution_context`` is entered around the close and the finalizer, the
+ same way the stream wrapper scopes its own ``close``; a stream wrapper
+ passes its ``_execution_context`` so this cleanup path finalizes in the
+ same context as the others.
"""
- return cast(WrappedT, _CloseFinalizingProxy(wrapped, finalize))
+ return cast(
+ WrappedT, _CloseFinalizingProxy(wrapped, finalize, execution_context)
+ )
def finalize_on_aclose(
- wrapped: WrappedT, finalize: Callable[[], None]
+ wrapped: WrappedT,
+ finalize: Callable[[], None],
+ execution_context: _ExecutionContextFactory | None = None,
) -> WrappedT:
"""Async counterpart of ``finalize_on_close``, hooking ``aclose``."""
- return cast(WrappedT, _AcloseFinalizingProxy(wrapped, finalize))
+ return cast(
+ WrappedT, _AcloseFinalizingProxy(wrapped, finalize, execution_context)
+ )
class SyncStreamManagerWrapper(
@@ -663,24 +720,34 @@ def __exit__(
) -> bool | None:
stream_wrapper = self._self_stream_wrapper
self._self_stream_wrapper = None
- try:
- suppressed = self.__wrapped__.__exit__(exc_type, exc_val, exc_tb)
- except BaseException as error:
- if stream_wrapper is not None:
- stream_wrapper.__exit__(
- type(error), error, error.__traceback__
+ # The SDK manager closes its stream on exit, so that runs in the
+ # stream wrapper's execution context like every other cleanup.
+ with (
+ nullcontext()
+ if stream_wrapper is None
+ else stream_wrapper._execution_context()
+ ):
+ try:
+ suppressed = self.__wrapped__.__exit__(
+ exc_type, exc_val, exc_tb
)
- elif self._self_invocation is not None:
- self._self_invocation.fail(error)
- raise
- if stream_wrapper is not None:
- if suppressed:
- # The manager swallowed the caller's exception, so the stream
- # ended successfully as far as telemetry is concerned.
- stream_wrapper.__exit__(None, None, None)
- else:
- stream_wrapper.__exit__(exc_type, exc_val, exc_tb)
- return suppressed
+ except BaseException as error:
+ if stream_wrapper is not None:
+ stream_wrapper.__exit__(
+ type(error), error, error.__traceback__
+ )
+ elif self._self_invocation is not None:
+ self._self_invocation.fail(error)
+ raise
+ if stream_wrapper is not None:
+ if suppressed:
+ # The manager swallowed the caller's exception, so the
+ # stream ended successfully as far as telemetry is
+ # concerned.
+ stream_wrapper.__exit__(None, None, None)
+ else:
+ stream_wrapper.__exit__(exc_type, exc_val, exc_tb)
+ return suppressed
class AsyncStreamManagerWrapper(
@@ -730,25 +797,30 @@ async def __aexit__(
) -> bool | None:
stream_wrapper = self._self_stream_wrapper
self._self_stream_wrapper = None
- try:
- suppressed = await self.__wrapped__.__aexit__(
- exc_type, exc_val, exc_tb
- )
- except BaseException as error:
- if stream_wrapper is not None:
- await stream_wrapper.__aexit__(
- type(error), error, error.__traceback__
+ # See SyncStreamManagerWrapper.__exit__.
+ with (
+ nullcontext()
+ if stream_wrapper is None
+ else stream_wrapper._execution_context()
+ ):
+ try:
+ suppressed = await self.__wrapped__.__aexit__(
+ exc_type, exc_val, exc_tb
)
- elif self._self_invocation is not None:
- self._self_invocation.fail(error)
- raise
- if stream_wrapper is not None:
- if suppressed:
- # See SyncStreamManagerWrapper.__exit__.
- await stream_wrapper.__aexit__(None, None, None)
- else:
- await stream_wrapper.__aexit__(exc_type, exc_val, exc_tb)
- return suppressed
+ except BaseException as error:
+ if stream_wrapper is not None:
+ await stream_wrapper.__aexit__(
+ type(error), error, error.__traceback__
+ )
+ elif self._self_invocation is not None:
+ self._self_invocation.fail(error)
+ raise
+ if stream_wrapper is not None:
+ if suppressed:
+ await stream_wrapper.__aexit__(None, None, None)
+ else:
+ await stream_wrapper.__aexit__(exc_type, exc_val, exc_tb)
+ return suppressed
__all__ = [
diff --git a/util/opentelemetry-util-genai/tests/test_stream.py b/util/opentelemetry-util-genai/tests/test_stream.py
index c83093510..d160cd385 100644
--- a/util/opentelemetry-util-genai/tests/test_stream.py
+++ b/util/opentelemetry-util-genai/tests/test_stream.py
@@ -7,17 +7,23 @@
import gc
import inspect
import timeit
+from collections.abc import AsyncGenerator, Generator, Iterator
+from contextlib import AbstractContextManager, contextmanager, nullcontext
+from contextvars import ContextVar
+from typing import Any
from unittest.mock import MagicMock, patch
import pytest
+from opentelemetry.context import attach, detach
from opentelemetry.sdk.metrics import MeterProvider
+from opentelemetry.sdk.metrics.export import InMemoryMetricReader
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
from opentelemetry.sdk.trace.export.in_memory_span_exporter import (
InMemorySpanExporter,
)
-from opentelemetry.trace import get_current_span
+from opentelemetry.trace import get_current_span, set_span_in_context
from opentelemetry.trace.status import StatusCode
from opentelemetry.util.genai._inference_invocation import (
SuppressedInferenceInvocation,
@@ -1006,6 +1012,77 @@ async def exercise():
asyncio.run(exercise())
+@pytest.mark.parametrize("close_error", [None, RuntimeError("close failure")])
+def test_finalize_on_close_runs_inside_execution_context(close_error):
+ active = ContextVar("active", default=False)
+ seen: list[bool] = []
+
+ @contextmanager
+ def scope() -> Iterator[None]:
+ token = active.set(True)
+ try:
+ yield
+ finally:
+ active.reset(token)
+
+ class Closable(_FakeClosable):
+ def close(self):
+ seen.append(active.get())
+ super().close()
+
+ proxy = finalize_on_close(
+ Closable(close_error=close_error),
+ lambda: seen.append(active.get()),
+ execution_context=scope,
+ )
+
+ with (
+ pytest.raises(RuntimeError, match="close failure")
+ if close_error
+ else nullcontext()
+ ):
+ proxy.close()
+
+ assert seen == [True, True]
+ assert not active.get()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("close_error", [None, RuntimeError("close failure")])
+async def test_finalize_on_aclose_runs_inside_execution_context(close_error):
+ active = ContextVar("active", default=False)
+ seen: list[bool] = []
+
+ @contextmanager
+ def scope() -> Iterator[None]:
+ token = active.set(True)
+ try:
+ yield
+ finally:
+ active.reset(token)
+
+ class Closable(_FakeClosable):
+ async def aclose(self):
+ seen.append(active.get())
+ await super().aclose()
+
+ proxy = finalize_on_aclose(
+ Closable(close_error=close_error),
+ lambda: seen.append(active.get()),
+ execution_context=scope,
+ )
+
+ with (
+ pytest.raises(RuntimeError, match="close failure")
+ if close_error
+ else nullcontext()
+ ):
+ await proxy.aclose()
+
+ assert seen == [True, True]
+ assert not active.get()
+
+
class _FakeInvocation:
def __init__(self):
self.stop_count = 0
@@ -1813,3 +1890,557 @@ async def exercise():
assert spans[0].attributes["error.type"] == "GeneratorExit"
asyncio.run(exercise())
+
+
+@pytest.mark.parametrize("failure", [False, True])
+def test_sync_execution_scope_includes_send_throw_and_close(
+ failure: bool,
+) -> None:
+ active = ContextVar("active", default=False)
+ closed: list[bool] = []
+ error = ConnectionError("stream failed")
+
+ @contextmanager
+ def scope() -> Iterator[None]:
+ token = active.set(True)
+ try:
+ yield
+ finally:
+ active.reset(token)
+
+ class ScopedWrapper(_TestSyncStreamWrapper):
+ def _execution_context(self) -> AbstractContextManager[None]:
+ return scope()
+
+ def _process_chunk(self, chunk: Any) -> None:
+ assert active.get()
+ super()._process_chunk(chunk)
+
+ def _on_stream_end(self) -> None:
+ assert active.get()
+ super()._on_stream_end()
+
+ def _on_stream_error(self, error: BaseException) -> None:
+ assert active.get()
+ super()._on_stream_error(error)
+
+ def produce() -> Generator[str, str, None]:
+ assert active.get()
+ try:
+ value = yield "first"
+ assert active.get()
+ assert value == "sent"
+ try:
+ yield "second"
+ except ConnectionError as caught:
+ assert caught is error
+ assert active.get()
+ if failure:
+ raise
+ yield "recovered"
+ finally:
+ assert active.get()
+ closed.append(True)
+
+ wrapper = ScopedWrapper(produce())
+ assert next(wrapper) == "first"
+ assert not active.get()
+ assert wrapper.send("sent") == "second"
+ assert not active.get()
+ if failure:
+ with pytest.raises(ConnectionError) as raised:
+ wrapper.throw(error)
+ assert raised.value is error
+ assert wrapper._self_failures == [error]
+ else:
+ assert wrapper.throw(error) == "recovered"
+ assert not active.get()
+ wrapper.close()
+ assert not active.get()
+ assert closed == [True]
+ assert wrapper._self_stop_count == (0 if failure else 1)
+ assert not hasattr(ScopedWrapper(_FakeSyncStream()), "send")
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("failure", [False, True])
+async def test_async_execution_scope_includes_asend_athrow_and_close(
+ failure: bool,
+) -> None:
+ active = ContextVar("active", default=False)
+ closed: list[bool] = []
+ error = ConnectionError("stream failed")
+
+ @contextmanager
+ def scope() -> Iterator[None]:
+ token = active.set(True)
+ try:
+ yield
+ finally:
+ active.reset(token)
+
+ class ScopedWrapper(_TestAsyncStreamWrapper):
+ def _execution_context(self) -> AbstractContextManager[None]:
+ return scope()
+
+ def _process_chunk(self, chunk: Any) -> None:
+ assert active.get()
+ super()._process_chunk(chunk)
+
+ def _on_stream_end(self) -> None:
+ assert active.get()
+ super()._on_stream_end()
+
+ def _on_stream_error(self, error: BaseException) -> None:
+ assert active.get()
+ super()._on_stream_error(error)
+
+ async def produce() -> AsyncGenerator[str, str]:
+ assert active.get()
+ try:
+ value = yield "first"
+ assert active.get()
+ assert value == "sent"
+ try:
+ yield "second"
+ except ConnectionError as caught:
+ assert caught is error
+ assert active.get()
+ if failure:
+ raise
+ yield "recovered"
+ finally:
+ assert active.get()
+ closed.append(True)
+
+ wrapper = ScopedWrapper(produce())
+ assert await anext(wrapper) == "first"
+ assert not active.get()
+ assert await wrapper.asend("sent") == "second"
+ assert not active.get()
+ if failure:
+ with pytest.raises(ConnectionError) as raised:
+ await wrapper.athrow(error)
+ assert raised.value is error
+ assert wrapper._self_failures == [error]
+ else:
+ assert await wrapper.athrow(error) == "recovered"
+ assert not active.get()
+ await wrapper.aclose()
+ assert not active.get()
+ assert closed == [True]
+ assert wrapper._self_stop_count == (0 if failure else 1)
+ assert not hasattr(ScopedWrapper(_FakeAsyncStream()), "asend")
+
+
+@pytest.mark.parametrize("failure", [False, True])
+def test_sync_manager_exit_runs_inside_stream_execution_scope(
+ failure: bool,
+) -> None:
+ active = ContextVar("active", default=False)
+ exits: list[bool] = []
+ error = RuntimeError("exit failed")
+
+ @contextmanager
+ def scope() -> Iterator[None]:
+ token = active.set(True)
+ try:
+ yield
+ finally:
+ active.reset(token)
+
+ class ScopedWrapper(_TestSyncStreamWrapper):
+ def _execution_context(self) -> AbstractContextManager[None]:
+ return scope()
+
+ class ScopedManagerWrapper(SyncStreamManagerWrapper):
+ def _wrap_stream(self, stream, invocation):
+ return ScopedWrapper(stream, invocation=invocation)
+
+ class Manager(_FakeSyncManager):
+ def __exit__(self, exc_type, exc_val, exc_tb):
+ # The SDK closes its stream here, so this is stream cleanup.
+ exits.append(active.get())
+ return super().__exit__(exc_type, exc_val, exc_tb)
+
+ stream = _FakeSyncStream(chunks=["a"])
+ manager = Manager(stream, exit_error=error if failure else None)
+ wrapper = ScopedManagerWrapper(manager, _FakeInvocation)
+ invocation = None
+
+ with pytest.raises(RuntimeError) if failure else nullcontext():
+ with wrapper as stream_wrapper:
+ invocation = stream_wrapper._self_invocation
+ assert not active.get()
+ assert next(stream_wrapper) == "a"
+ assert not active.get()
+
+ assert exits == [True]
+ assert not active.get()
+ assert invocation is not None
+ if failure:
+ assert invocation.failures == [error]
+ assert invocation.stop_count == 0
+ else:
+ assert invocation.failures == []
+ assert invocation.stop_count == 1
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("failure", [False, True])
+async def test_async_manager_exit_runs_inside_stream_execution_scope(
+ failure: bool,
+) -> None:
+ active = ContextVar("active", default=False)
+ exits: list[bool] = []
+ error = RuntimeError("exit failed")
+
+ @contextmanager
+ def scope() -> Iterator[None]:
+ token = active.set(True)
+ try:
+ yield
+ finally:
+ active.reset(token)
+
+ class ScopedWrapper(_TestAsyncStreamWrapper):
+ def _execution_context(self) -> AbstractContextManager[None]:
+ return scope()
+
+ class ScopedManagerWrapper(AsyncStreamManagerWrapper):
+ def _wrap_stream(self, stream, invocation):
+ return ScopedWrapper(stream, invocation=invocation)
+
+ class Manager(_FakeAsyncManager):
+ async def __aexit__(self, exc_type, exc_val, exc_tb):
+ exits.append(active.get())
+ return await super().__aexit__(exc_type, exc_val, exc_tb)
+
+ stream = _FakeAsyncStream(chunks=["a"])
+ manager = Manager(stream, exit_error=error if failure else None)
+ wrapper = ScopedManagerWrapper(manager, _FakeInvocation)
+ invocation = None
+
+ with pytest.raises(RuntimeError) if failure else nullcontext():
+ async with wrapper as stream_wrapper:
+ invocation = stream_wrapper._self_invocation
+ assert not active.get()
+ assert await anext(stream_wrapper) == "a"
+ assert not active.get()
+
+ assert exits == [True]
+ assert not active.get()
+ assert invocation is not None
+ if failure:
+ assert invocation.failures == [error]
+ assert invocation.stop_count == 0
+ else:
+ assert invocation.failures == []
+ assert invocation.stop_count == 1
+
+
+class _ActivatingInvocation:
+ """Invocation double that records whether reads ran inside activate()."""
+
+ def __init__(self, suspended: bool) -> None:
+ self._suspended = suspended
+ self.active = False
+ self.activations = 0
+
+ def _on_stream_chunk(self, chunk_at: float) -> None:
+ pass
+
+ def stop(self) -> None:
+ pass
+
+ def fail(self, error: Any) -> None:
+ pass
+
+ @contextmanager
+ def activate(self) -> Iterator[None]:
+ self.activations += 1
+ self.active = True
+ try:
+ yield
+ finally:
+ self.active = False
+
+
+@pytest.mark.parametrize("attach", [True, False])
+def test_default_execution_context_follows_suspended(attach: bool):
+ invocation = _ActivatingInvocation(attach)
+ seen: list[bool] = []
+
+ class Wrapper(_TestSyncStreamWrapper):
+ def _process_chunk(self, chunk: Any) -> None:
+ seen.append(invocation.active)
+ super()._process_chunk(chunk)
+
+ wrapper = Wrapper(iter(["a", "b"]), invocation)
+ assert list(wrapper) == ["a", "b"]
+ assert seen == [attach, attach]
+ # one activation per read: two chunks and the read that ends the stream
+ assert invocation.activations == (3 if attach else 0)
+ assert invocation.active is False
+
+
+def test_default_execution_context_without_suspended_flag_is_a_noop():
+ class Bare:
+ def _on_stream_chunk(self, chunk_at: float) -> None:
+ pass
+
+ def stop(self) -> None:
+ pass
+
+ def fail(self, error: Any) -> None:
+ pass
+
+ wrapper = _TestSyncStreamWrapper(iter(["a"]), Bare())
+ assert list(wrapper) == ["a"]
+
+
+def _suspended_tool_invocation():
+ """A real invocation started under a caller span, then suspended as a stream wrapper expects."""
+ tracer_provider = TracerProvider()
+ tracer_provider.add_span_processor(
+ SimpleSpanProcessor(InMemorySpanExporter())
+ )
+ invocation = ToolInvocation(
+ tracer=tracer_provider.get_tracer(__name__),
+ instruments=_Instruments(MeterProvider().get_meter(__name__)),
+ logger=MagicMock(),
+ completion_hook=MagicMock(spec=CompletionHook),
+ name="abandoned_tool",
+ )
+ invocation.suspend()
+ return invocation
+
+
+@pytest.mark.parametrize(
+ "wrapper_cls", [_TestSyncStreamWrapper, SyncToolStreamWrapper]
+)
+def test_abandoned_sync_stream_finalizes_inside_execution_context(wrapper_cls):
+ invocation = _suspended_tool_invocation()
+ spans_at_finalize = []
+
+ class _RecordingWrapper(wrapper_cls):
+ def _on_stream_error(self, error):
+ spans_at_finalize.append(get_current_span())
+ super()._on_stream_error(error)
+
+ wrapper = _RecordingWrapper(iter(["a", "b"]), invocation)
+ with pytest.raises(RuntimeError):
+ for _ in wrapper:
+ raise RuntimeError("caller error")
+
+ del wrapper
+ gc.collect()
+
+ assert spans_at_finalize == [invocation.span]
+ assert get_current_span() is not invocation.span
+
+
+@pytest.mark.parametrize(
+ "wrapper_cls", [_TestAsyncStreamWrapper, AsyncToolStreamWrapper]
+)
+def test_abandoned_async_stream_finalizes_inside_execution_context(
+ wrapper_cls,
+):
+ invocation = _suspended_tool_invocation()
+ spans_at_finalize = []
+
+ class _RecordingWrapper(wrapper_cls):
+ def _on_stream_error(self, error):
+ spans_at_finalize.append(get_current_span())
+ super()._on_stream_error(error)
+
+ async def chunks():
+ yield "a"
+ yield "b"
+
+ async def exercise():
+ wrapper = _RecordingWrapper(chunks(), invocation)
+ with pytest.raises(RuntimeError):
+ async for _ in wrapper:
+ raise RuntimeError("caller error")
+
+ asyncio.run(exercise())
+ gc.collect()
+
+ assert spans_at_finalize == [invocation.span]
+ assert get_current_span() is not invocation.span
+
+
+def test_default_execution_context_leaves_suppressed_inference_invocation_alone():
+ """A nested inference invocation never attached itself, so reads don't activate it."""
+ from opentelemetry.util.genai._inference_invocation import ( # pylint: disable=import-outside-toplevel
+ SuppressedInferenceInvocation,
+ )
+ from opentelemetry.util.genai.handler import ( # pylint: disable=import-outside-toplevel
+ TelemetryHandler,
+ )
+
+ handler = TelemetryHandler(tracer_provider=TracerProvider())
+ outer = handler.inference("openai", request_model="model")
+ try:
+ inner = handler.inference("openai", request_model="model")
+ assert isinstance(inner, SuppressedInferenceInvocation)
+ inner.suspend()
+ assert not inner._suspended
+
+ wrapper = _TestSyncStreamWrapper(iter(["a"]), inner)
+ assert list(wrapper) == ["a"]
+ assert wrapper._self_finalized
+ assert get_current_span() is outer.span
+ finally:
+ outer.stop()
+
+
+def _parent_and_attached_invocation():
+ """A caller span made current, then a real invocation started under it and left attached."""
+ tracer_provider = TracerProvider()
+ tracer_provider.add_span_processor(
+ SimpleSpanProcessor(InMemorySpanExporter())
+ )
+ parent = tracer_provider.get_tracer(__name__).start_span("caller")
+ token = attach(set_span_in_context(parent))
+ invocation = ToolInvocation(
+ tracer=tracer_provider.get_tracer(__name__),
+ instruments=_Instruments(MeterProvider().get_meter(__name__)),
+ logger=MagicMock(),
+ completion_hook=MagicMock(spec=CompletionHook),
+ name="attached_tool",
+ )
+ assert get_current_span() is invocation.span
+ return parent, token, invocation
+
+
+@pytest.mark.parametrize("ending", ["exhaust", "error", "close"])
+def test_default_execution_context_leaves_unsuspended_invocation_alone(ending):
+ """An instrumentation that never calls suspend() must still end with the parent current."""
+ parent, token, invocation = _parent_and_attached_invocation()
+ try:
+
+ def chunks():
+ yield "a"
+ if ending == "error":
+ raise RuntimeError("boom")
+ yield "b"
+
+ wrapper = _TestSyncStreamWrapper(chunks(), invocation)
+ if ending == "exhaust":
+ assert list(wrapper) == ["a", "b"]
+ elif ending == "error":
+ with pytest.raises(RuntimeError):
+ list(wrapper)
+ else:
+ assert next(iter(wrapper)) == "a"
+ wrapper.close()
+ assert wrapper._self_finalized
+ assert get_current_span() is parent
+ finally:
+ detach(token)
+ parent.end()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("ending", ["exhaust", "error", "close"])
+async def test_default_execution_context_leaves_unsuspended_invocation_alone_async(
+ ending,
+):
+ parent, token, invocation = _parent_and_attached_invocation()
+ try:
+
+ async def chunks():
+ yield "a"
+ if ending == "error":
+ raise RuntimeError("boom")
+ yield "b"
+
+ wrapper = _TestAsyncStreamWrapper(chunks(), invocation)
+ if ending == "exhaust":
+ assert [c async for c in wrapper] == ["a", "b"]
+ elif ending == "error":
+ with pytest.raises(RuntimeError):
+ _ = [c async for c in wrapper]
+ else:
+ assert await wrapper.__anext__() == "a"
+ await wrapper.aclose()
+ assert wrapper._self_finalized
+ assert get_current_span() is parent
+ finally:
+ detach(token)
+ parent.end()
+
+
+def test_default_execution_context_activates_suspended_invocation():
+ parent, token, invocation = _parent_and_attached_invocation()
+ try:
+ invocation.suspend()
+ assert get_current_span() is parent
+ seen: list[Any] = []
+
+ class Wrapper(_TestSyncStreamWrapper):
+ def _process_chunk(self, chunk: Any) -> None:
+ seen.append(get_current_span())
+ super()._process_chunk(chunk)
+
+ wrapper = Wrapper(iter(["a", "b"]), invocation)
+ for _ in wrapper:
+ assert get_current_span() is parent
+ assert seen == [invocation.span, invocation.span]
+ assert get_current_span() is parent
+ finally:
+ detach(token)
+ parent.end()
+
+
+def test_tool_stream_wrapper_records_no_chunk_timing_metrics():
+ reader = InMemoryMetricReader()
+ invocation = ToolInvocation(
+ tracer=TracerProvider().get_tracer(__name__),
+ instruments=_Instruments(
+ MeterProvider(metric_readers=[reader]).get_meter(__name__)
+ ),
+ logger=MagicMock(),
+ completion_hook=MagicMock(spec=CompletionHook),
+ name="metrics_tool",
+ )
+ assert list(SyncToolStreamWrapper(iter(["a", "b"]), invocation)) == [
+ "a",
+ "b",
+ ]
+
+ names: set[str] = set()
+ data = reader.get_metrics_data()
+ for resource_metrics in data.resource_metrics if data else []:
+ for scope_metrics in resource_metrics.scope_metrics:
+ names.update(m.name for m in scope_metrics.metrics)
+ assert names == {"gen_ai.execute_tool.duration"}
+
+
+def test_tool_stream_wrapper_leaves_detached_tool_invocation_alone():
+ """A tool invocation created with _attach_to_context=False is never made current by the wrapper."""
+ caller_span = get_current_span()
+ span_exporter = InMemorySpanExporter()
+ tracer_provider = TracerProvider()
+ tracer_provider.add_span_processor(SimpleSpanProcessor(span_exporter))
+ invocation = ToolInvocation(
+ tracer=tracer_provider.get_tracer(__name__),
+ instruments=_Instruments(MeterProvider().get_meter(__name__)),
+ logger=MagicMock(),
+ completion_hook=MagicMock(spec=CompletionHook),
+ name="detached_tool",
+ _attach_to_context=False,
+ )
+ assert get_current_span() is caller_span
+ inside: list[Any] = []
+
+ def tool_body():
+ inside.append(get_current_span())
+ yield "a"
+
+ assert list(SyncToolStreamWrapper(tool_body(), invocation)) == ["a"]
+ assert inside == [caller_span]
+ assert get_current_span() is caller_span
+ assert len(span_exporter.get_finished_spans()) == 1