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