diff --git a/.github/instructions/instrumentation.instructions.md b/.github/instructions/instrumentation.instructions.md index 98c6664c6..4bac8f0b6 100644 --- a/.github/instructions/instrumentation.instructions.md +++ b/.github/instructions/instrumentation.instructions.md @@ -55,7 +55,9 @@ 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 `invocation.suspend()` before the stream is returned and + return `invocation.activate()` from `_execution_context()`. - 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..88a23c326 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -229,14 +229,20 @@ 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 +or `close()`. Subclasses pass the SDK stream to `super().__init__(stream)` and implement four hooks: +- `_execution_context()` — return a fresh context manager for each stream read and cleanup + operation; for invocation-backed streams return `invocation.activate()`. Context is restored + before the chunk is returned to the consumer, so the token is created and released in the + same frame. - `_process_chunk(chunk)` — accumulate per-chunk state (e.g. response model, finish reasons, token usage, streamed content) onto the invocation. - `_on_stream_end()` — finalize on success; set the accumulated response attributes and call @@ -248,8 +254,12 @@ class MyStreamWrapper(SyncStreamWrapper[Chunk]): def __init__(self, stream, invocation, capture_content): super().__init__(stream) self._self_invocation = invocation + invocation.suspend() ... + def _execution_context(self): + return self._self_invocation.activate() + def _process_chunk(self, chunk): ... # accumulate state def _on_stream_end(self): self._self_invocation.stop() diff --git a/README.md b/README.md index d5fb69933..b37f79349 100644 --- a/README.md +++ b/README.md @@ -17,7 +17,7 @@ All instrumentations use [opentelemetry-util-genai](./util/opentelemetry-util-ge | [opentelemetry-instrumentation-genai-anthropic](./instrumentation/opentelemetry-instrumentation-genai-anthropic) | anthropic >= 0.51.0, < 2 | [1.2b0](https://pypi.org/project/opentelemetry-instrumentation-genai-anthropic/) | | [opentelemetry-instrumentation-genai-bedrock](./instrumentation/opentelemetry-instrumentation-genai-bedrock) | boto3 >= 1.40.46, < 2 | [1.2b0](https://pypi.org/project/opentelemetry-instrumentation-genai-bedrock/) | | [opentelemetry-instrumentation-genai-dspy](./instrumentation/opentelemetry-instrumentation-genai-dspy) | dspy >= 3.3.0, < 4 | [1.2b0](https://pypi.org/project/opentelemetry-instrumentation-genai-dspy/) | -| [opentelemetry-instrumentation-genai-langchain](./instrumentation/opentelemetry-instrumentation-genai-langchain) | langchain >= 0.3.21, < 2 | [1.2b0](https://pypi.org/project/opentelemetry-instrumentation-genai-langchain/) | +| [opentelemetry-instrumentation-genai-langchain](./instrumentation/opentelemetry-instrumentation-genai-langchain) | langchain >= 0.3.22, < 2 | [1.2b0](https://pypi.org/project/opentelemetry-instrumentation-genai-langchain/) | | [opentelemetry-instrumentation-genai-llama-index](./instrumentation/opentelemetry-instrumentation-genai-llama-index) | llama-index-core >= 0.14.19, < 1, llama-index-instrumentation >= 0.4.3, < 1, llama-index-workflows >= 2.17.1, != 2.24.0, < 3 | [1.2b0](https://pypi.org/project/opentelemetry-instrumentation-genai-llama-index/) | | [opentelemetry-instrumentation-genai-openai](./instrumentation/opentelemetry-instrumentation-genai-openai) | openai >= 1.26.0, < 4 | [1.2b0](https://pypi.org/project/opentelemetry-instrumentation-genai-openai/) | | [opentelemetry-instrumentation-genai-openai-agents](./instrumentation/opentelemetry-instrumentation-genai-openai-agents) | openai-agents >= 0.3.3, < 1 | [1.2b0](https://pypi.org/project/opentelemetry-instrumentation-genai-openai-agents/) | diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/.changelog/818.changed b/instrumentation/opentelemetry-instrumentation-genai-langchain/.changelog/818.changed new file mode 100644 index 000000000..25f50afac --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/.changelog/818.changed @@ -0,0 +1 @@ +Require `langchain >= 0.3.22`: the execution boundaries rely on `set_config_context`, which langchain-core added in 0.3.46 (langchain 0.3.22 is the first release whose langchain-core floor includes it). LangGraph, when installed, needs 0.3.18 or newer for node boundaries. diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/.changelog/818.fixed b/instrumentation/opentelemetry-instrumentation-genai-langchain/.changelog/818.fixed new file mode 100644 index 000000000..bb5baf602 --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/.changelog/818.fixed @@ -0,0 +1 @@ +Activate the span context around LangChain and LangGraph execution boundaries (chat models, tools, retrievers, runnables, streams, graph nodes) instead of attaching it from callbacks, so nested SDK and HTTP spans correlate across asyncio tasks and threads without `Failed to detach context` errors. A chat model call cancelled or interrupted during `ainvoke` now ends its span with that error instead of leaking it. diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/README.rst b/instrumentation/opentelemetry-instrumentation-genai-langchain/README.rst index 57685a6c2..b30397c4a 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/README.rst +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/README.rst @@ -29,6 +29,10 @@ Installation pip install opentelemetry-instrumentation-genai-langchain +LangGraph is optional. Graph node boundaries are instrumented from LangGraph +0.3.18 on; an older LangGraph still runs, with a warning that context is not +propagated across its nodes. + See the `examples `_ directory for runnable ``workflow``, ``agent``, ``tools``, and ``zero-code`` scenarios. @@ -141,8 +145,14 @@ programmatically, which takes precedence over the environment variable:: Known Limitations ----------------- -Context propagation to nested calls (such as auto-instrumented HTTP clients -or database queries within tools) is not supported when using LangChain async API. +Context is propagated to nested calls (such as auto-instrumented HTTP clients or +database queries) where LangChain hands control to user code: chat models, tools, +retrievers, runnables built on ``_call_with_config``, streams, the steps of sequence, +parallel, branch and fallback runnables, single-input ``batch``/``abatch`` calls, and +LangGraph nodes. A ``Runnable`` whose ``invoke`` starts no run of its own is not +correlated when ``batch``/``abatch`` runs it for several inputs at once: the default +implementation dispatches each input from a thread pool or task with no frame of its +own to carry that input's parent. References ---------- diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/pyproject.toml b/instrumentation/opentelemetry-instrumentation-genai-langchain/pyproject.toml index 9955207e7..c9ff6829d 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/pyproject.toml +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/pyproject.toml @@ -26,12 +26,12 @@ classifiers = [ ] dependencies = [ "opentelemetry-instrumentation >= 0.64b0, <1", - "opentelemetry-util-genai >= 1.2b0, <2", + "opentelemetry-util-genai >= 1.3b0.dev, <2", ] [project.optional-dependencies] instruments = [ - "langchain >= 0.3.21, < 2", + "langchain >= 0.3.22, < 2", ] [project.entry-points.opentelemetry_instrumentor] diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/__init__.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/__init__.py index 218969449..78b38a868 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/__init__.py +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/__init__.py @@ -29,9 +29,11 @@ from typing import Any from langchain_core.callbacks import BaseCallbackManager -from langchain_core.callbacks.manager import AsyncCallbackManager from wrapt import wrap_function_wrapper +from opentelemetry.instrumentation.genai.langchain._execution_context import ( + _ExecutionContext, +) from opentelemetry.instrumentation.genai.langchain.agent_context import ( wrap_astream, wrap_stream, @@ -57,6 +59,8 @@ class LangChainInstrumentor(BaseInstrumentor): to capture LLM telemetry. """ + _execution_context: _ExecutionContext | None = None + def __init__( self, ): @@ -83,22 +87,22 @@ def _instrument(self, **kwargs: Any): instrumentation_scope_version=__version__, ) invocation_manager = _InvocationManager() - sync_handler = OpenTelemetryLangChainCallbackHandler( - telemetry_handler=telemetry_handler, - _attach_to_context=True, - invocation_manager=invocation_manager, - ) - async_handler = OpenTelemetryLangChainCallbackHandler( + handler = OpenTelemetryLangChainCallbackHandler( telemetry_handler=telemetry_handler, - _attach_to_context=False, invocation_manager=invocation_manager, ) wrap_function_wrapper( "langchain_core.callbacks", "BaseCallbackManager.__init__", - _BaseCallbackManagerInitWrapper(sync_handler, async_handler), + _BaseCallbackManagerInitWrapper(handler), ) + self._execution_context = _ExecutionContext( + invocation_manager, + handler.on_tool_error, + handler.on_retriever_error, + ) + self._execution_context.instrument() self._instrument_agent_entry_points() @staticmethod @@ -120,6 +124,8 @@ def _uninstrument(self, **kwargs: Any): Cleanup instrumentation (unwrap). """ unwrap("langchain_core.callbacks.base.BaseCallbackManager", "__init__") + # The agent entry points wrap last, over the execution boundary, so + # they come off first. try: import langgraph.pregel @@ -127,6 +133,9 @@ def _uninstrument(self, **kwargs: Any): unwrap(langgraph.pregel.Pregel, method) except (ImportError, AttributeError): pass + if self._execution_context is not None: + self._execution_context.uninstrument() + self._execution_context = None class _BaseCallbackManagerInitWrapper: @@ -136,11 +145,9 @@ class _BaseCallbackManagerInitWrapper: def __init__( self, - sync_handler: OpenTelemetryLangChainCallbackHandler, - async_handler: OpenTelemetryLangChainCallbackHandler, + handler: OpenTelemetryLangChainCallbackHandler, ): - self._sync_handler = sync_handler - self._async_handler = async_handler + self._handler = handler def __call__( self, @@ -150,29 +157,8 @@ def __call__( kwargs: dict[str, Any], ): wrapped(*args, **kwargs) - target_handler = ( - self._async_handler - if isinstance(instance, AsyncCallbackManager) - else self._sync_handler - ) - other_handler = ( - self._sync_handler - if isinstance(instance, AsyncCallbackManager) - else self._async_handler - ) - if other_handler in instance.handlers: - instance.handlers = [ - target_handler if h is other_handler else h - for h in instance.handlers - ] - if other_handler in instance.inheritable_handlers: - instance.inheritable_handlers = [ - target_handler if h is other_handler else h - for h in instance.inheritable_handlers - ] - if target_handler not in instance.inheritable_handlers: - for handler in instance.inheritable_handlers: - if isinstance(handler, OpenTelemetryLangChainCallbackHandler): - break - else: - instance.add_handler(target_handler, inherit=True) + if not any( + isinstance(handler, OpenTelemetryLangChainCallbackHandler) + for handler in instance.inheritable_handlers + ): + instance.add_handler(self._handler, inherit=True) diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/_execution_context.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/_execution_context.py new file mode 100644 index 000000000..efbd65c48 --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/_execution_context.py @@ -0,0 +1,377 @@ +# Copyright The OpenTelemetry Authors +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import logging +from collections.abc import Callable, Iterator, Sequence +from contextlib import contextmanager +from contextvars import Context as PythonContext +from functools import partial, wraps +from importlib import import_module +from importlib.metadata import PackageNotFoundError, distribution +from inspect import iscoroutinefunction, signature +from typing import Any, cast +from uuid import UUID + +from wrapt import wrap_function_wrapper + +from opentelemetry.context import attach, detach, get_current +from opentelemetry.instrumentation.genai.langchain._run_context import ( + _astart_run, + _RunScope, + _start_run, + _wrap_call, + _wrap_run, + _wrap_stream, +) +from opentelemetry.instrumentation.genai.langchain.invocation_manager import ( + _InvocationManager, +) +from opentelemetry.instrumentation.utils import unwrap + +__all__ = ["_ExecutionContext"] + +_logger = logging.getLogger(__name__) + +_METHODS = ( + ( + "langchain_core.language_models.chat_models", + "BaseChatModel", + "_generate_with_cache", + ), + ( + "langchain_core.language_models.chat_models", + "BaseChatModel", + "_agenerate_with_cache", + ), +) + + +class _ExecutionContext: + def __init__( + self, + invocations: _InvocationManager, + on_tool_error: Callable[..., None], + on_retriever_error: Callable[..., None], + ) -> None: + self._invocations = invocations + self._on_tool_error = on_tool_error + self._on_retriever_error = on_retriever_error + self._patched: list[tuple[Any, str]] = [] + + @contextmanager + def _activate(self, run_id: UUID | None) -> Iterator[None]: + context = self._invocations.get_parent_context(run_id) + if context is None or get_current() is context: + yield + return + token = attach(context) + try: + yield + finally: + detach(token) + + def _wrap_batch(self, original: Callable[..., Any]) -> Callable[..., Any]: + parameters = signature(original) + + def parent_run_id_for( + instance: Any, args: tuple[Any, ...], kwargs: dict[str, Any] + ) -> UUID | None: + # One input runs in this frame; several run in the executor or + # as tasks, each in a copied context that this frame never sees. + bound = parameters.bind(instance, *args, **kwargs).arguments + if len(bound["inputs"]) != 1: + return None + config = bound.get("config") + if isinstance(config, Sequence): + config = cast(Any, config[0]) if config else None + callbacks = cast(Any, config or {}).get("callbacks") + return getattr(callbacks, "parent_run_id", None) + + @wraps(original) + def sync( + wrapped: Callable[..., Any], + instance: Any, + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> Any: + with self._activate(parent_run_id_for(instance, args, kwargs)): + return wrapped(*args, **kwargs) + + @wraps(original) + async def asynchronous( + wrapped: Callable[..., Any], + instance: Any, + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> Any: + with self._activate(parent_run_id_for(instance, args, kwargs)): + return await wrapped(*args, **kwargs) + + return asynchronous if iscoroutinefunction(original) else sync + + def _wrap(self, original: Callable[..., Any]) -> Callable[..., Any]: + parameters = signature(original) + + def run_id_for( + instance: Any, args: tuple[Any, ...], kwargs: dict[str, Any] + ) -> UUID | None: + run_manager = parameters.bind( + instance, *args, **kwargs + ).arguments.get("run_manager") + return run_manager.run_id if run_manager is not None else None + + @wraps(original) + def sync( + wrapped: Callable[..., Any], + instance: Any, + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> Any: + with self._activate(run_id_for(instance, args, kwargs)): + return wrapped(*args, **kwargs) + + @wraps(original) + async def asynchronous( + wrapped: Callable[..., Any], + instance: Any, + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> Any: + run_id = run_id_for(instance, args, kwargs) + with self._activate(run_id): + try: + return await wrapped(*args, **kwargs) + except BaseException as error: + # agenerate reports an Exception from its gather results + # and reaches on_llm_error; a cancellation or interrupt + # leaves the gather before that, with no handler around + # the await. + if not isinstance(error, Exception): + _RunScope(run_id, self._invocations).finish(error) + raise + + return asynchronous if iscoroutinefunction(original) else sync + + @contextmanager + def _config_context( + self, + wrapped: Callable[..., Any], + instance: Any, + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> Iterator[PythonContext]: + config = args[0] if args else kwargs["config"] + callbacks = config.get("callbacks") + parent = self._invocations.get_parent_context( + getattr(callbacks, "parent_run_id", None) + ) + with wrapped(*args, **kwargs) as context: + # LangChain runs each child (a sequence step, a parallel branch, a + # graph node) in this copied context. Restore the token in that + # same context, never in a callback or the consuming task. A child + # whose own boundary already attached the parent needs no second + # token. + token = ( + context.run(attach, parent) + if parent is not None + and context.run(get_current) is not parent + else None + ) + try: + yield context + finally: + if token is not None: + context.run(detach, token) + + def _patch( + self, + module_name: str, + class_name: str | None, + method: str, + wrapper_for: Callable[[Callable[..., Any]], Callable[..., Any]], + *, + fallback_modules: tuple[str, ...] = (), + ) -> None: + for name in (module_name, *fallback_modules): + try: + owner: Any = import_module(name) + if class_name is not None: + owner = getattr(owner, class_name) + original = getattr(owner, method) + except (ImportError, AttributeError): + continue + wrap_function_wrapper(owner, method, wrapper_for(original)) + self._patched.append((owner, method)) + return + target = ".".join(filter(None, (module_name, class_name, method))) + library = module_name.partition(".")[0] + try: + # The namespace alone proves nothing: langgraph-checkpoint fills + # it without the runtime. Look for the distribution instead. + distribution(library) + except PackageNotFoundError: + _logger.debug( + "Skipping execution boundary %s: %s is not installed", + target, + library, + ) + else: + # Installed but changed: spans under it will not correlate. + _logger.warning( + "Skipping execution boundary %s: not found in the " + "installed %s, so context is not propagated across it", + target, + library, + ) + + def instrument(self) -> None: + for module_name, class_name, method in _METHODS: + self._patch(module_name, class_name, method, self._wrap) + + for module_name, class_name, methods, on_error in ( + ( + "langchain_core.tools.base", + "BaseTool", + ("run", "arun"), + self._on_tool_error, + ), + ( + "langchain_core.retrievers", + "BaseRetriever", + ("invoke", "ainvoke"), + self._on_retriever_error, + ), + ): + for method, asynchronous in zip(methods, (False, True)): + self._patch( + module_name, + class_name, + method, + partial( + _wrap_run, + invocations=self._invocations, + on_error=on_error, + asynchronous=asynchronous, + ), + ) + + # These start the chain run themselves, so the run id is chosen up + # front for _started to match. + for module_name, class_name, methods in ( + ( + "langchain_core.runnables.base", + "Runnable", + ("_call_with_config", "_acall_with_config"), + ), + # RunnableBranch dispatches the chosen branch straight from its + # own run, so a branch without a boundary inherits it from here. + ( + "langchain_core.runnables.branch", + "RunnableBranch", + ("invoke", "ainvoke"), + ), + ): + for method, asynchronous in zip(methods, (False, True)): + self._patch( + module_name, + class_name, + method, + partial( + _wrap_call, + invocations=self._invocations, + asynchronous=asynchronous, + ), + ) + + # The default batch runs a single input in the caller's frame. + for method in ("batch", "abatch"): + self._patch( + "langchain_core.runnables.base", + "Runnable", + method, + self._wrap_batch, + ) + + for class_name, wrapper in ( + ("CallbackManager", _start_run), + ("AsyncCallbackManager", _astart_run), + ): + for method in ( + "on_chat_model_start", + "on_chain_start", + "on_tool_start", + "on_retriever_start", + ): + self._patch( + "langchain_core.callbacks.manager", + class_name, + method, + lambda _original, wrapper=wrapper: wrapper, + ) + + for module_name, class_name, methods in ( + ( + "langchain_core.language_models.chat_models", + "BaseChatModel", + ("stream", "astream"), + ), + ( + "langchain_core.runnables.base", + "Runnable", + ( + "_transform_stream_with_config", + "_atransform_stream_with_config", + ), + ), + # RunnableLambda.astream does not forward aclose to its inner + # transform iterator, so its own boundary needs finalization too. + ( + "langchain_core.runnables.base", + "RunnableLambda", + ("stream", "astream"), + ), + ( + "langchain_core.runnables.branch", + "RunnableBranch", + ("stream", "astream"), + ), + ("langgraph.pregel", "Pregel", ("stream", "astream")), + ): + for method, asynchronous in zip(methods, (False, True)): + self._patch( + module_name, + class_name, + method, + partial( + _wrap_stream, + invocations=self._invocations, + asynchronous=asynchronous, + ), + ) + + # Composite runnables drive steps that may start no run of their own + # (a plain ``invoke`` override), so the child context is attached where + # LangChain and LangGraph enter the step's copied context. Each caller + # binds ``set_config_context`` by name at import, so the replacement + # goes into the calling modules rather than ``runnables.config``; tools + # bind it too, but ``BaseTool.run`` above already covers them. + for module_name, fallback_modules in ( + ("langchain_core.runnables.base", ()), + ("langchain_core.runnables.fallbacks", ()), + # LangGraph moved the helper into _internal in 0.6. + ("langgraph._internal._runnable", ("langgraph.utils.runnable",)), + ): + self._patch( + module_name, + None, + "set_config_context", + lambda _original: self._config_context, + fallback_modules=fallback_modules, + ) + + def uninstrument(self) -> None: + for owner, name in reversed(self._patched): + unwrap(owner, name) + self._patched.clear() diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/_run_context.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/_run_context.py new file mode 100644 index 000000000..474a476fe --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/_run_context.py @@ -0,0 +1,332 @@ +# Copyright The OpenTelemetry Authors +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import logging +from collections.abc import ( + AsyncGenerator, + AsyncIterator, + Callable, + Generator, + Iterator, +) +from contextlib import AbstractContextManager, ExitStack, contextmanager +from contextvars import ContextVar +from dataclasses import dataclass +from functools import wraps +from inspect import BoundArguments, signature +from typing import Any +from uuid import UUID + +from langchain_core.callbacks.manager import BaseRunManager + +try: + # LangChain mints time-ordered run ids; use the same generator so the ids + # tracers see for runs started here match the ones LangChain starts. + from langchain_core.utils.uuid import uuid7 as _new_run_id +except ImportError: # older langchain-core + from uuid import uuid4 as _new_run_id + +from opentelemetry.context import Context, attach, detach, get_current +from opentelemetry.instrumentation.genai.langchain.invocation_manager import ( + _InvocationManager, +) +from opentelemetry.util.genai.stream import ( + AsyncStreamWrapper, + SyncStreamWrapper, +) + +__all__ = [ + "_RunScope", + "_astart_run", + "_start_run", + "_wrap_call", + "_wrap_run", + "_wrap_stream", +] + +_logger = logging.getLogger(__name__) + + +class _RunScope: + def __init__( + self, + run_id: UUID | None, + invocations: _InvocationManager, + on_error: Callable[..., None] | None = None, + ) -> None: + # None until a deferred stream chooses its run id on the first read. + self.run_id = run_id + self.invocations = invocations + self.context: Context | None = None + self.on_error = on_error + + @contextmanager + def activate(self) -> Iterator[None]: + with ExitStack() as stack: + token = _active_run.set(_RunActivation(self, stack)) + stack.callback(_active_run.reset, token) + # A stream read inside an enclosing read of the same run tree + # already has this context current. + if self.context is not None and get_current() is not self.context: + stack.callback(detach, attach(self.context)) + yield + + def finish(self, error: BaseException | None = None) -> None: + # The callback handler ends the invocation from its end/error callback; + # its state stays registered while a child is live, so this may find + # an ended invocation, and finishing it again is a no-op in the util. + # This is the end when no callback fires: a tool or retriever cancelled + # by a BaseException the runner does not report, a chat model whose + # agenerate gather child is cancelled or interrupted, or a stream + # closed early by a runner that leaves its inner iterator open + # (RunnableLambda.astream). on_error routes a tool or retriever error + # through the handler's own error path so the attributes match. + if self.run_id is None: + return + invocation = self.invocations.get_invocation(self.run_id) + try: + if invocation is not None: + if error is None: + invocation.stop() + elif self.on_error is not None: + self.on_error(error, run_id=self.run_id) + else: + invocation.fail(error) + except Exception: + _logger.exception("Failed to finalize LangChain run") + finally: + self.invocations.delete_invocation_state(self.run_id) + + +@dataclass +class _RunActivation: + scope: _RunScope + stack: ExitStack + + +_active_run: ContextVar[_RunActivation | None] = ContextVar( + "otel_langchain_run_activation", default=None +) + + +def _started(result: BaseRunManager | list[BaseRunManager]) -> None: + read = _active_run.get() + if read is None: + return + managers = result if isinstance(result, list) else [result] + if not any(manager.run_id == read.scope.run_id for manager in managers): + return + context = read.scope.invocations.get_parent_context(read.scope.run_id) + if context is not None: + # The manager returns in the runner's context, after handlers finish. + # Its token belongs to this activation, never to a callback handler. + # An enclosing composite may have attached the same context already. + read.scope.context = context + if get_current() is not context: + read.stack.callback(detach, attach(context)) + + +def _start_run( + wrapped: Callable[..., Any], + instance: Any, + args: tuple[Any, ...], + kwargs: dict[str, Any], +) -> Any: + result = wrapped(*args, **kwargs) + _started(result) + return result + + +async def _astart_run( + wrapped: Callable[..., Any], + instance: Any, + args: tuple[Any, ...], + kwargs: dict[str, Any], +) -> Any: + result = await wrapped(*args, **kwargs) + _started(result) + return result + + +class _SyncContextStream(SyncStreamWrapper[Any]): + _self_scope: _RunScope + + def __init__(self, stream: Any, scope: _RunScope) -> None: + super().__init__(stream) + self._self_scope = scope + + def _execution_context(self) -> AbstractContextManager[None]: + return self._self_scope.activate() + + def _process_chunk(self, chunk: Any) -> None: + pass + + def _on_stream_end(self) -> None: + self._self_scope.finish() + + def _on_stream_error(self, error: BaseException) -> None: + self._self_scope.finish(error) + + +class _AsyncContextStream(AsyncStreamWrapper[Any]): + _self_scope: _RunScope + + def __init__(self, stream: Any, scope: _RunScope) -> None: + super().__init__(stream) + self._self_scope = scope + + def _execution_context(self) -> AbstractContextManager[None]: + return self._self_scope.activate() + + def _process_chunk(self, chunk: Any) -> None: + pass + + def _on_stream_end(self) -> None: + self._self_scope.finish() + + def _on_stream_error(self, error: BaseException) -> None: + self._self_scope.finish(error) + + +def _config_run(bound: BoundArguments, scope: _RunScope) -> BoundArguments: + # The run id is chosen here so _started can tell this frame's own run from + # one started under the same activation before reaching its own boundary + # (a chat model invoked from a lambda body): attaching that one here would + # leave its ended span current for the rest of this frame. + config = dict(bound.arguments.get("config") or {}) + scope.run_id = config.get("run_id") or _new_run_id() + config["run_id"] = scope.run_id + bound.arguments["config"] = config + return bound + + +def _deferred(start: Callable[[], Iterator[Any]]) -> Generator[Any, Any, None]: + yield from start() + + +async def _adeferred( + start: Callable[[], AsyncIterator[Any]], +) -> AsyncGenerator[Any, None]: + stream = start() + try: + async for chunk in stream: + yield chunk + finally: + close = getattr(stream, "aclose", None) + if close is not None: + await close() + + +def _wrap_stream( + original: Callable[..., Any], + invocations: _InvocationManager, + asynchronous: bool, +) -> Callable[..., Any]: + parameters = signature(original) + + @wraps(original) + def wrapper( + wrapped: Callable[..., Any], + instance: Any, + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> Any: + scope = _RunScope(None, invocations) + bound = parameters.bind(instance, *args, **kwargs) + + def start() -> Any: + # A generator reads its config on the first advancement; keep + # that so the caller may still edit the config until then. + _config_run(bound, scope) + return wrapped(*bound.args[1:], **bound.kwargs) + + stream: Any = _adeferred(start) if asynchronous else _deferred(start) + # Keep the SDK's name on the generator the caller sees. + stream.__qualname__ = original.__qualname__ + stream.__name__ = original.__name__ + if asynchronous: + return _AsyncContextStream(stream, scope) + return _SyncContextStream(stream, scope) + + return wrapper + + +def _wrap_call( + original: Callable[..., Any], + invocations: _InvocationManager, + asynchronous: bool, +) -> Callable[..., Any]: + parameters = signature(original) + + @wraps(original) + def sync( + wrapped: Callable[..., Any], + instance: Any, + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> Any: + scope = _RunScope(None, invocations) + bound = _config_run(parameters.bind(instance, *args, **kwargs), scope) + with scope.activate(): + return wrapped(*bound.args[1:], **bound.kwargs) + + @wraps(original) + async def asynchronous_call( + wrapped: Callable[..., Any], + instance: Any, + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> Any: + scope = _RunScope(None, invocations) + bound = _config_run(parameters.bind(instance, *args, **kwargs), scope) + with scope.activate(): + return await wrapped(*bound.args[1:], **bound.kwargs) + + return asynchronous_call if asynchronous else sync + + +def _wrap_run( + original: Callable[..., Any], + invocations: _InvocationManager, + on_error: Callable[..., None], + asynchronous: bool, +) -> Callable[..., Any]: + def scope_for(kwargs: dict[str, Any]) -> _RunScope: + run_id = kwargs.get("run_id") or _new_run_id() + kwargs["run_id"] = run_id + return _RunScope(run_id, invocations, on_error) + + @wraps(original) + def sync( + wrapped: Callable[..., Any], + instance: Any, + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> Any: + scope = scope_for(kwargs) + with scope.activate(): + try: + return wrapped(*args, **kwargs) + except BaseException as error: + # Some runners omit error callbacks for cancellation. + scope.finish(error) + raise + + @wraps(original) + async def asynchronous_call( + wrapped: Callable[..., Any], + instance: Any, + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> Any: + scope = scope_for(kwargs) + with scope.activate(): + try: + return await wrapped(*args, **kwargs) + except BaseException as error: + scope.finish(error) + raise + + return asynchronous_call if asynchronous else sync diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/callback_handler.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/callback_handler.py index e339712cd..a27cc6be9 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/callback_handler.py +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/callback_handler.py @@ -166,12 +166,10 @@ def __init__( self, telemetry_handler: TelemetryHandler, *, - _attach_to_context: bool = True, invocation_manager: _InvocationManager | None = None, ) -> None: super().__init__() self._telemetry_handler = telemetry_handler - self._attach_to_context = _attach_to_context self._invocation_manager = ( invocation_manager if invocation_manager is not None @@ -221,7 +219,7 @@ def on_chain_start( name=workflow_name_override or workflow_name, context=parent_context, conversation_id=conversation_id, - _attach_to_context=self._attach_to_context, + _attach_to_context=False, ) if capture_content: workflow.input_messages = make_input_message(inputs) @@ -255,7 +253,7 @@ def on_chain_start( agent_name=suggested_agent_name, context=parent_context, conversation_id=conversation_id, - _attach_to_context=self._attach_to_context, + _attach_to_context=False, ) if capture_content: agent.input_messages = make_input_message(inputs) @@ -281,7 +279,7 @@ def on_chain_start( agent_name=None, context=parent_context, conversation_id=conversation_id, - _attach_to_context=self._attach_to_context, + _attach_to_context=False, ) agent.input_messages = make_input_message(inputs) self._invocation_manager.add_invocation_state( @@ -442,7 +440,7 @@ def on_chat_model_start( request_model=request_model, context=parent_context, conversation_id=_conversation_id(metadata), - _attach_to_context=self._attach_to_context, + _attach_to_context=False, ) llm_invocation.input_messages = input_messages llm_invocation.top_p = top_p @@ -751,7 +749,7 @@ def on_tool_start( tool_type="function", agent_name=agent_name, context=parent_context, - _attach_to_context=self._attach_to_context, + _attach_to_context=False, ) tool_invocation.tool_description = description tool_invocation.arguments = arguments @@ -815,7 +813,7 @@ def on_retriever_start( provider=provider, request_model=request_model, context=parent_context, - _attach_to_context=self._attach_to_context, + _attach_to_context=False, ) retrieval.query_text = query self._invocation_manager.add_invocation_state( diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/package.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/package.py index 74028c734..35156b6f1 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/package.py +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/package.py @@ -2,4 +2,4 @@ # SPDX-License-Identifier: Apache-2.0 -_instruments = ("langchain >= 0.3.21, < 2",) +_instruments = ("langchain >= 0.3.22, < 2",) diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.latest.txt b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.latest.txt index dfacf6af1..714bff1a8 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.latest.txt +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.latest.txt @@ -22,4 +22,7 @@ langchain-anthropic==1.7.5 boto3==1.43.108 -e util/opentelemetry-util-genai --e instrumentation/opentelemetry-instrumentation-genai-langchain[instruments] \ No newline at end of file +-e instrumentation/opentelemetry-instrumentation-genai-langchain[instruments] + +opentelemetry-instrumentation-httpx >= 0.64b0, <1 +-e instrumentation/opentelemetry-instrumentation-genai-openai[instruments] diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.oldest.txt b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.oldest.txt index db47a8396..4ee613722 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.oldest.txt +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.oldest.txt @@ -20,8 +20,22 @@ # pyproject floors by UV_RESOLUTION=lowest-direct on the oldest tox factor, so they are NOT pinned # here to avoid drift between the declared bound and the tested version. +# langchain-core is transitive, so lowest-direct does not floor it: pin the floor that +# langchain 0.3.22 declares, which is the oldest core with the set_config_context boundary. +langchain-core==0.3.49 +# LangGraph is optional and has no bound in pyproject.toml: 0.3.18 is the oldest release with +# the node boundary (langgraph.utils.runnable.set_config_context) the instrumentation wraps. +langgraph==0.3.18 langchain-openai==0.2.0 langchain-aws==0.2.2 langchain-google-genai==2.0.0 langchain-anthropic==0.3.0 boto3==1.37.0 + +opentelemetry-instrumentation-httpx >= 0.64b0, <1 +-e instrumentation/opentelemetry-instrumentation-genai-openai[instruments] +# The oldest OpenAI client still passes the removed httpx proxies argument. +httpx==0.27.2 + +# The stream execution hook is in the unreleased workspace util package. +-e util/opentelemetry-util-genai diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_agent_classification_corpus.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_agent_classification_corpus.py index 158659d3d..d54305134 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_agent_classification_corpus.py +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_agent_classification_corpus.py @@ -397,7 +397,7 @@ def test_agent_named_runnable_is_an_agent() -> None: agent_name="SupportAgentRunner", context=None, conversation_id=None, - _attach_to_context=True, + _attach_to_context=False, ) diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_callback_handler.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_callback_handler.py index e6fd013f6..b79dcf18e 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_callback_handler.py +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_callback_handler.py @@ -154,7 +154,7 @@ def test_workflow_name_from_serialized(self): name="MyLangGraph", context=None, conversation_id=None, - _attach_to_context=True, + _attach_to_context=False, ) def test_workflow_name_overridden_by_metadata(self): @@ -173,7 +173,7 @@ def test_workflow_name_overridden_by_metadata(self): name="custom_workflow", context=None, conversation_id=None, - _attach_to_context=True, + _attach_to_context=False, ) def test_workflow_conversation_id_from_metadata(self): @@ -230,7 +230,7 @@ def test_child_agent_passes_parent_context_to_telemetry_handler(self): agent_name="math_agent", context=workflow_inv.context, conversation_id=None, - _attach_to_context=True, + _attach_to_context=False, ) @@ -256,7 +256,7 @@ def test_new_agent_span_created(self): agent_name="math_agent", context=None, conversation_id=None, - _attach_to_context=True, + _attach_to_context=False, ) assert ( handler._invocation_manager.get_agent_name(run_id) == "math_agent" @@ -286,7 +286,7 @@ def test_agent_name_heuristic_sets_standard_agent_attributes(self): agent_name="AgentExecutor", context=None, conversation_id="thread-abc", - _attach_to_context=True, + _attach_to_context=False, ) assert agent_inv.input_messages[0].parts[0].content == "Solve this" assert agent_inv.output_messages[0].parts[0].content == "Solved" @@ -577,7 +577,7 @@ def test_chat_model_passes_parent_context_to_telemetry_handler(self): request_model="gpt-4", context=parent_wf.context, conversation_id=None, - _attach_to_context=True, + _attach_to_context=False, ) @@ -765,7 +765,7 @@ def test_named_child_under_workflow_opens_agent_layer(self): agent_name="math_agent", context=workflow_inv.context, conversation_id=None, - _attach_to_context=True, + _attach_to_context=False, ) @@ -1479,7 +1479,7 @@ def test_tool_passes_parent_context_to_telemetry_handler(self): tool_type="function", agent_name=None, context=parent_wf.context, - _attach_to_context=True, + _attach_to_context=False, ) @@ -1545,7 +1545,7 @@ def test_provider_passed_from_metadata(self): provider="Chroma", request_model=None, context=None, - _attach_to_context=True, + _attach_to_context=False, ) def test_provider_none_when_metadata_absent(self): @@ -1562,7 +1562,7 @@ def test_provider_none_when_metadata_absent(self): provider=None, request_model=None, context=None, - _attach_to_context=True, + _attach_to_context=False, ) def test_request_model_passed_from_ls_embedding_model(self): @@ -1583,7 +1583,7 @@ def test_request_model_passed_from_ls_embedding_model(self): provider="Chroma", request_model="text-embedding-3-small", context=None, - _attach_to_context=True, + _attach_to_context=False, ) def test_request_model_none_when_ls_embedding_model_absent(self): @@ -1601,7 +1601,7 @@ def test_request_model_none_when_ls_embedding_model_absent(self): provider="Chroma", request_model=None, context=None, - _attach_to_context=True, + _attach_to_context=False, ) def test_registered_in_invocation_manager(self): @@ -1641,7 +1641,7 @@ def test_retriever_passes_parent_context_to_telemetry_handler(self): provider=None, request_model=None, context=parent_wf.context, - _attach_to_context=True, + _attach_to_context=False, ) @@ -3462,32 +3462,7 @@ def test_on_chat_model_start_preserves_message_name(): assert llm_inv.input_messages[0].name == "Alice" -def test_explicit_attach_to_context_false(): - telemetry = mock.MagicMock() - workflow_inv = mock.MagicMock(spec=WorkflowInvocation) - telemetry.workflow.return_value = workflow_inv - - handler = OpenTelemetryLangChainCallbackHandler( - telemetry, _attach_to_context=False - ) - run_id = _run_id() - - handler.on_chain_start( - serialized={"name": "LangGraph", "id": ["langgraph"]}, - inputs={}, - run_id=run_id, - parent_run_id=None, - ) - - telemetry.workflow.assert_called_once_with( - name="LangGraph", - context=None, - conversation_id=None, - _attach_to_context=False, - ) - - -def test_sync_defaults_attach_to_context_true(): +def test_handler_does_not_attach_context(): telemetry = mock.MagicMock() workflow_inv = mock.MagicMock(spec=WorkflowInvocation) telemetry.workflow.return_value = workflow_inv @@ -3506,11 +3481,11 @@ def test_sync_defaults_attach_to_context_true(): name="LangGraph", context=None, conversation_id=None, - _attach_to_context=True, + _attach_to_context=False, ) -def test_instrumentor_routes_sync_and_async_callback_managers(): +def test_instrumentor_shares_handler_across_callback_managers(): from langchain_core.callbacks.manager import ( AsyncCallbackManager, CallbackManager, @@ -3529,7 +3504,6 @@ def test_instrumentor_routes_sync_and_async_callback_managers(): if isinstance(h, OpenTelemetryLangChainCallbackHandler) ] assert len(sync_handlers) == 1 - assert sync_handlers[0]._attach_to_context is True assert sync_handlers[0].run_inline is False acm = AsyncCallbackManager([]) @@ -3539,7 +3513,7 @@ def test_instrumentor_routes_sync_and_async_callback_managers(): if isinstance(h, OpenTelemetryLangChainCallbackHandler) ] assert len(async_handlers) == 1 - assert async_handlers[0]._attach_to_context is False + assert sync_handlers[0] is async_handlers[0] assert async_handlers[0].run_inline is False finally: LangChainInstrumentor().uninstrument() diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_custom_execution_context.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_custom_execution_context.py new file mode 100644 index 000000000..abdd81f2b --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_custom_execution_context.py @@ -0,0 +1,487 @@ +# Copyright The OpenTelemetry Authors +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +import sys +from collections.abc import Awaitable, Callable +from contextlib import nullcontext +from importlib.metadata import version +from typing import Any + +import pytest +from langchain_core.documents import Document +from langchain_core.retrievers import BaseRetriever +from langchain_core.runnables import ( + Runnable, + RunnableBranch, + RunnableConfig, + RunnableLambda, + RunnableParallel, +) +from langchain_core.tools import BaseTool +from packaging.version import Version + +from opentelemetry import context +from opentelemetry.semconv.attributes import error_attributes + +from .test_execution_context import ( + _HTTP_SCOPE, + _LC_SCOPE, + _OPENAI_SCOPE, + _child, + _spans, + clients, + no_detach_errors, +) +from .test_stream_context import _assert_no_runs + +__all__ = ["clients", "no_detach_errors"] + +requires_async_step_context = pytest.mark.skipif( + sys.version_info < (3, 11) + and Version(version("langchain-core")) < Version("1.5.0"), + reason="langchain-core < 1.5 on Python < 3.11 doesn't run async steps in the copied context", +) +_async_step = pytest.param("async", marks=requires_async_step_context) + + +def _custom_operation( + kind: str, + call: Callable[[str], str], + acall: Callable[[str], Awaitable[str]] | None = None, +) -> BaseTool | BaseRetriever: + class CustomTool(BaseTool): + name: str = "custom" + description: str = "Custom tool execution." + + def _run(self, text: str) -> str: + return call(text) + + class AsyncCustomTool(CustomTool): + async def _arun(self, text: str) -> str: + assert acall is not None + return await acall(text) + + class CustomRetriever(BaseRetriever): + def _get_relevant_documents(self, query: str) -> list[Document]: + return [Document(page_content=call(query))] + + class AsyncCustomRetriever(CustomRetriever): + async def _aget_relevant_documents(self, query: str) -> list[Document]: + assert acall is not None + return [Document(page_content=await acall(query))] + + if kind == "tool": + return CustomTool() if acall is None else AsyncCustomTool() + return CustomRetriever() if acall is None else AsyncCustomRetriever() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["tool", "retriever"]) +@pytest.mark.parametrize("mode", ["sync", "async", "executor"]) +@pytest.mark.parametrize("parented", [False, True]) +async def test_custom_execution_correlates_inference_and_http( + clients: Any, + span_exporter: Any, + tracer_provider: Any, + kind: str, + mode: str, + parented: bool, +) -> None: + def call(text: str) -> str: + clients.http.get("https://example.test/custom") + return clients.infer(text) + + async def acall(text: str) -> str: + await clients.ahttp.get("https://example.test/custom") + return await clients.ainfer(text) + + operation = _custom_operation( + kind, call, acall if mode == "async" else None + ) + parent = ( + tracer_provider.get_tracer("test").start_as_current_span("request") + if parented + else nullcontext() + ) + with parent as root: + before = context.get_current() + result = ( + operation.invoke("hello") + if mode == "sync" + else await operation.ainvoke("hello") + ) + assert ( + result if kind == "tool" else result[0].page_content + ) == "answer" + assert context.get_current() is before + (execution,) = _spans(span_exporter, _LC_SCOPE) + assert execution.parent == (root.get_span_context() if root else None) + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + direct_http, provider_http = _spans(span_exporter, _HTTP_SCOPE) + _child(direct_http, execution) + _child(inference, execution) + _child(provider_http, inference) + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["tool", "retriever"]) +@pytest.mark.parametrize("mode", ["sync", "async", "executor", "sync_cancel"]) +async def test_custom_execution_error_restores_context( + clients: Any, + span_exporter: Any, + kind: str, + mode: str, +) -> None: + error = ( + asyncio.CancelledError("cancelled") + if mode == "sync_cancel" + else ConnectionError("failed") + ) + + def call(text: str) -> str: + clients.http.get("https://example.test/custom") + raise error + + async def acall(text: str) -> str: + await clients.ahttp.get("https://example.test/custom") + raise error + + operation = _custom_operation( + kind, call, acall if mode == "async" else None + ) + before = context.get_current() + with pytest.raises(type(error)) as raised: + if mode in ("sync", "sync_cancel"): + operation.invoke("hello") + else: + await operation.ainvoke("hello") + assert raised.value is error + assert context.get_current() is before + (execution,) = _spans(span_exporter, _LC_SCOPE) + (http,) = _spans(span_exporter, _HTTP_SCOPE) + _child(http, execution) + assert execution.attributes[error_attributes.ERROR_TYPE] == ( + "asyncio.exceptions.CancelledError" + if mode == "sync_cancel" + else "ConnectionError" + ) + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["tool", "retriever"]) +async def test_custom_execution_task_cancellation( + clients: Any, + span_exporter: Any, + tracer_provider: Any, + kind: str, +) -> None: + started = asyncio.Event() + errors: list[BaseException] = [] + + async def acall(text: str) -> str: + await clients.ahttp.get("https://example.test/custom") + started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError as error: + errors.append(error) + raise + return text + + operation = _custom_operation(kind, lambda text: text, acall) + with tracer_provider.get_tracer("test").start_as_current_span( + "request" + ) as root: + before = context.get_current() + + async def request() -> None: + try: + await operation.ainvoke("hello") + except asyncio.CancelledError as error: + assert error is errors[0] + assert error.args == ("cancelled by test",) + raise + finally: + assert context.get_current() is before + + task = asyncio.create_task(request()) + try: + await asyncio.wait_for(started.wait(), 5) + task.cancel("cancelled by test") + with pytest.raises(asyncio.CancelledError): + await task + finally: + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + assert context.get_current() is before + (execution,) = _spans(span_exporter, _LC_SCOPE) + (http,) = _spans(span_exporter, _HTTP_SCOPE) + assert execution.parent == root.get_span_context() + _child(http, execution) + assert ( + execution.attributes[error_attributes.ERROR_TYPE] + == "asyncio.exceptions.CancelledError" + ) + _assert_no_runs() + + +class _HttpRunnable(Runnable[str, str]): + """A Runnable built on the documented ``_call_with_config`` helpers.""" + + def __init__(self, clients: Any) -> None: + self._clients = clients + + def invoke( + self, input: str, config: RunnableConfig | None = None, **kwargs: Any + ) -> str: + return self._call_with_config(self._call, input, config, **kwargs) + + async def ainvoke( + self, input: str, config: RunnableConfig | None = None, **kwargs: Any + ) -> str: + return await self._acall_with_config( + self._acall, input, config, **kwargs + ) + + def _call(self, text: str) -> str: + self._clients.http.get("https://example.test/runnable") + return text + + async def _acall(self, text: str) -> str: + await self._clients.ahttp.get("https://example.test/runnable") + return text + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["sync", "async", "stream", "astream"]) +async def test_custom_runnable_in_sequence_correlates_http( + clients: Any, span_exporter: Any, mode: str +) -> None: + chain = _HttpRunnable(clients) | RunnableLambda(clients.infer) + before = context.get_current() + if mode == "sync": + result = chain.invoke("hello") + elif mode == "async": + result = await chain.ainvoke("hello") + elif mode == "stream": + result = "".join(chain.stream("hello")) + else: + result = "".join([chunk async for chunk in chain.astream("hello")]) + assert result == "answer" + assert context.get_current() is before + (workflow,) = _spans(span_exporter, _LC_SCOPE) + assert workflow.name == "invoke_workflow RunnableSequence" + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + runnable_http, provider_http = _spans(span_exporter, _HTTP_SCOPE) + _child(runnable_http, workflow) + _child(inference, workflow) + _child(provider_http, inference) + _assert_no_runs() + + +def _url(span: Any) -> str: + attributes = span.attributes + return str(attributes.get("url.full") or attributes.get("http.url")) + + +class _PlainRunnable(Runnable[str, str]): + """A Runnable whose ``invoke`` starts no run and uses no helper.""" + + def __init__(self, clients: Any, error: Exception | None = None) -> None: + self._clients = clients + self._error = error + + def invoke( + self, input: str, config: RunnableConfig | None = None, **kwargs: Any + ) -> str: + self._clients.http.get("https://example.test/plain") + if self._error is not None: + raise self._error + return input + + async def ainvoke( + self, input: str, config: RunnableConfig | None = None, **kwargs: Any + ) -> str: + await self._clients.ahttp.get("https://example.test/plain") + if self._error is not None: + raise self._error + return input + + +@pytest.mark.asyncio +@pytest.mark.parametrize("composite", ["sequence", "parallel", "nested"]) +@pytest.mark.parametrize("mode", ["sync", _async_step]) +async def test_plain_runnable_in_composite_correlates_http( + clients: Any, span_exporter: Any, mode: str, composite: str +) -> None: + plain = _PlainRunnable(clients) + infer = RunnableLambda(clients.infer) + if composite == "sequence": + chain: Runnable[str, Any] = plain | infer + elif composite == "parallel": + chain = RunnableParallel(plain=plain, infer=infer) + else: + chain = RunnableLambda(lambda text: text) | (plain | infer) + before = context.get_current() + result = ( + chain.invoke("hello") + if mode == "sync" + else await chain.ainvoke("hello") + ) + assert ( + result == {"plain": "hello", "infer": "answer"} + if composite == "parallel" + else result == "answer" + ) + assert context.get_current() is before + (workflow,) = _spans(span_exporter, _LC_SCOPE) + assert workflow.name.startswith("invoke_workflow Runnable") + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + http_spans = _spans(span_exporter, _HTTP_SCOPE) + (plain_http,) = [s for s in http_spans if _url(s).endswith("/plain")] + (provider_http,) = [s for s in http_spans if s is not plain_http] + _child(plain_http, workflow) + _child(inference, workflow) + _child(provider_http, inference) + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["sync", _async_step]) +async def test_plain_runnable_fallback_correlates_http( + clients: Any, span_exporter: Any, mode: str +) -> None: + chain = _PlainRunnable(clients, ConnectionError("down")).with_fallbacks( + [RunnableLambda(clients.infer)] + ) + before = context.get_current() + result = ( + chain.invoke("hello") + if mode == "sync" + else await chain.ainvoke("hello") + ) + assert result == "answer" + assert context.get_current() is before + (workflow,) = _spans(span_exporter, _LC_SCOPE) + assert workflow.name == "invoke_workflow RunnableWithFallbacks" + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + http_spans = _spans(span_exporter, _HTTP_SCOPE) + (plain_http,) = [s for s in http_spans if _url(s).endswith("/plain")] + (provider_http,) = [s for s in http_spans if s is not plain_http] + _child(plain_http, workflow) + _child(inference, workflow) + _child(provider_http, inference) + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["batch", "abatch"]) +async def test_sequence_batch_correlates_each_input( + clients: Any, span_exporter: Any, mode: str +) -> None: + chain = _HttpRunnable(clients) | RunnableLambda(clients.infer) + before = context.get_current() + if mode == "batch": + results = chain.batch(["one", "two"]) + else: + results = await chain.abatch(["one", "two"]) + assert results == ["answer", "answer"] + assert context.get_current() is before + workflows = _spans(span_exporter, _LC_SCOPE) + assert [w.name for w in workflows] == [ + "invoke_workflow RunnableSequence" + ] * 2 + inferences = _spans(span_exporter, _OPENAI_SCOPE) + http_spans = _spans(span_exporter, _HTTP_SCOPE) + runnable_http = [s for s in http_spans if _url(s).endswith("/runnable")] + provider_http = [s for s in http_spans if s not in runnable_http] + assert len(inferences) == len(runnable_http) == len(provider_http) == 2 + # Every root run parents exactly one runnable step and one inference. + roots = sorted(w.context.span_id for w in workflows) + assert sorted(s.parent.span_id for s in runnable_http) == roots + assert sorted(s.parent.span_id for s in inferences) == roots + for http in provider_http: + _child(http, next(i for i in inferences if i.context == http.parent)) + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["sync", "async", "stream", "astream"]) +async def test_plain_runnable_in_branch_correlates_http( + clients: Any, span_exporter: Any, mode: str +) -> None: + chain = RunnableBranch( + (lambda text: text == "hello", _PlainRunnable(clients)), + RunnableLambda(clients.infer), + ) + before = context.get_current() + if mode == "sync": + result = chain.invoke("hello") + elif mode == "async": + result = await chain.ainvoke("hello") + elif mode == "stream": + result = "".join(chain.stream("hello")) + else: + result = "".join([chunk async for chunk in chain.astream("hello")]) + assert result == "hello" + assert context.get_current() is before + (workflow,) = _spans(span_exporter, _LC_SCOPE) + assert workflow.name == "invoke_workflow RunnableBranch" + assert _spans(span_exporter, _OPENAI_SCOPE) == [] + (plain_http,) = _spans(span_exporter, _HTTP_SCOPE) + _child(plain_http, workflow) + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["batch", "abatch"]) +@pytest.mark.parametrize( + "inputs", + [ + ["one"], + pytest.param( + ["one", "two"], + marks=pytest.mark.xfail( + strict=True, + reason="Runnable.batch dispatches several inputs from the " + "executor with no frame to attach each input's parent in; " + "see the README known limitations", + ), + ), + ], +) +async def test_plain_runnable_batch_correlates_each_input( + clients: Any, span_exporter: Any, mode: str, inputs: list[str] +) -> None: + chain = _PlainRunnable(clients) | RunnableLambda(clients.infer) + before = context.get_current() + if mode == "batch": + results = chain.batch(inputs) + else: + results = await chain.abatch(inputs) + assert results == ["answer"] * len(inputs) + assert context.get_current() is before + workflows = _spans(span_exporter, _LC_SCOPE) + assert [w.name for w in workflows] == [ + "invoke_workflow RunnableSequence" + ] * len(inputs) + inferences = _spans(span_exporter, _OPENAI_SCOPE) + http_spans = _spans(span_exporter, _HTTP_SCOPE) + plain_http = [s for s in http_spans if _url(s).endswith("/plain")] + provider_http = [s for s in http_spans if s not in plain_http] + assert ( + len(inferences) == len(plain_http) == len(provider_http) == len(inputs) + ) + # Every root run parents exactly one plain step and one inference. + roots = sorted(w.context.span_id for w in workflows) + assert sorted(s.parent.span_id for s in plain_http if s.parent) == roots + assert sorted(s.parent.span_id for s in inferences) == roots + for http in provider_http: + _child(http, next(i for i in inferences if i.context == http.parent)) + _assert_no_runs() diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_execution_context.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_execution_context.py new file mode 100644 index 000000000..ea8fbcc37 --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_execution_context.py @@ -0,0 +1,921 @@ +# Copyright The OpenTelemetry Authors +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +import json +import logging +import sys +from collections.abc import AsyncIterator, Iterator +from contextlib import ExitStack +from contextvars import copy_context +from importlib import import_module, invalidate_caches +from importlib.util import find_spec +from pathlib import Path +from typing import Any, TypedDict + +import httpx +import pytest +import pytest_asyncio +from langchain_core.callbacks import BaseCallbackHandler +from langchain_core.callbacks.manager import CallbackManager +from langchain_core.documents import Document +from langchain_core.language_models.chat_models import BaseChatModel +from langchain_core.retrievers import BaseRetriever +from langchain_core.runnables import Runnable, RunnableBranch, RunnableLambda +from langchain_core.runnables import base as runnables_base +from langchain_core.tools import BaseTool, StructuredTool, Tool +from langchain_core.vectorstores import VectorStore +from langchain_openai import ChatOpenAI +from openai import AsyncOpenAI, OpenAI + +from opentelemetry import baggage, context, trace +from opentelemetry.instrumentation.genai.langchain import ( + LangChainInstrumentor, + agent_context, +) +from opentelemetry.instrumentation.genai.langchain.callback_handler import ( + OpenTelemetryLangChainCallbackHandler, +) +from opentelemetry.instrumentation.genai.openai import OpenAIInstrumentor +from opentelemetry.instrumentation.httpx import HTTPXClientInstrumentor +from opentelemetry.sdk.trace import ReadableSpan +from opentelemetry.semconv.attributes import error_attributes +from opentelemetry.test_util_genai.instrumentor import instrument +from opentelemetry.trace import StatusCode + +_LC_SCOPE = "opentelemetry.instrumentation.genai.langchain" +_OPENAI_SCOPE = "opentelemetry.instrumentation.genai.openai" +_HTTP_SCOPE = "opentelemetry.instrumentation.httpx" + + +@pytest.fixture(autouse=True) +def no_detach_errors(caplog) -> Iterator[None]: + yield + assert not [ + record + for phase in ("setup", "call", "teardown") + for record in caplog.get_records(phase) + if record.name == "opentelemetry.context" and record.levelno >= 40 + ] + + +class _Clients: + def __init__(self, http: httpx.Client, ahttp: httpx.AsyncClient) -> None: + self.http = http + self.ahttp = ahttp + self.requests: list[httpx.Request] = [] + self.stream_error: BaseException | None = None + self.stream_waiting: asyncio.Event | None = None + self.openai = OpenAI(api_key="test", http_client=http, max_retries=0) + self.aopenai = AsyncOpenAI( + api_key="test", http_client=ahttp, max_retries=0 + ) + + def infer(self, text: str) -> str: + result = self.openai.chat.completions.create( + model="test-model", messages=[{"role": "user", "content": text}] + ) + return result.choices[0].message.content + + async def ainfer(self, text: str) -> str: + result = await self.aopenai.chat.completions.create( + model="test-model", messages=[{"role": "user", "content": text}] + ) + return result.choices[0].message.content + + +@pytest_asyncio.fixture +async def clients( + tracer_provider, meter_provider, logger_provider +) -> AsyncIterator[_Clients]: + class ResponseStream(httpx.SyncByteStream, httpx.AsyncByteStream): + def __iter__(self) -> Iterator[bytes]: + yield sse_chunk("answer", None) + if client_set.stream_error is not None: + raise client_set.stream_error + yield sse_chunk("", "stop") + yield b"data: [DONE]\n\n" + + async def __aiter__(self) -> AsyncIterator[bytes]: + for chunk in self: + await asyncio.sleep(0) + yield chunk + if client_set.stream_waiting is not None: + client_set.stream_waiting.set() + await asyncio.Event().wait() + + class ResponsesApiStream(ResponseStream): + def __iter__(self) -> Iterator[bytes]: + response = { + "id": "resp-test", + "object": "response", + "created_at": 1, + "model": "test-model", + "status": "in_progress", + "output": [], + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + } + message = { + "type": "message", + "id": "msg-test", + "role": "assistant", + "status": "completed", + "content": [ + { + "type": "output_text", + "text": "answer", + "annotations": [], + } + ], + } + yield sse_event("response.created", response=response) + yield sse_event( + "response.output_text.delta", + item_id="msg-test", + output_index=0, + content_index=0, + delta="answer", + logprobs=[], + ) + if client_set.stream_error is not None: + raise client_set.stream_error + yield sse_event( + "response.completed", + response={ + **response, + "status": "completed", + "output": [message], + "usage": { + "input_tokens": 1, + "output_tokens": 1, + "total_tokens": 2, + "input_tokens_details": {"cached_tokens": 0}, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + }, + ) + + def sse_event(event_type: str, **payload: Any) -> bytes: + data = json.dumps( + {"type": event_type, "sequence_number": 0, **payload} + ) + return f"event: {event_type}\ndata: {data}\n\n".encode() + + def sse_chunk(content: str, finish_reason: str | None) -> bytes: + payload = { + "id": "chatcmpl-test", + "object": "chat.completion.chunk", + "created": 1, + "model": "test-model", + "choices": [ + { + "index": 0, + "delta": {"content": content}, + "finish_reason": finish_reason, + } + ], + } + return f"data: {json.dumps(payload)}\n\n".encode() + + def respond(request: httpx.Request) -> httpx.Response: + client_set.requests.append(request) + current = trace.get_current_span().get_span_context() + assert request.headers["traceparent"] == ( + f"00-{current.trace_id:032x}-{current.span_id:016x}-{current.trace_flags:02x}" + ) + if request.method == "POST" and json.loads(request.content).get( + "stream" + ): + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + stream=( + ResponsesApiStream() + if request.url.path.endswith("/responses") + else ResponseStream() + ), + ) + return httpx.Response( + 200, + json={ + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 1, + "model": "test-model", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "answer"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + }, + ) + + with ExitStack() as stack: + for instrumentor in (LangChainInstrumentor(), OpenAIInstrumentor()): + stack.enter_context( + instrument( + instrumentor, + tracer_provider=tracer_provider, + meter_provider=meter_provider, + logger_provider=logger_provider, + ) + ) + with httpx.Client(transport=httpx.MockTransport(respond)) as http: + async with httpx.AsyncClient( + transport=httpx.MockTransport(respond) + ) as ahttp: + for client in (http, ahttp): + HTTPXClientInstrumentor.instrument_client( + client, + tracer_provider=tracer_provider, + meter_provider=meter_provider, + ) + stack.callback( + HTTPXClientInstrumentor.uninstrument_client, client + ) + client_set = _Clients(http, ahttp) + yield client_set + + +def _spans(exporter: Any, scope: str) -> list[ReadableSpan]: + return [ + s + for s in exporter.get_finished_spans() + if s.instrumentation_scope.name == scope + ] + + +def _child(child: ReadableSpan, parent: ReadableSpan) -> None: + assert child.parent == parent.context + assert child.context.trace_id == parent.context.trace_id + assert ( + parent.start_time + <= child.start_time + <= child.end_time + <= parent.end_time + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["structured", "simple"]) +@pytest.mark.parametrize("mode", ["sync", "async", "executor"]) +async def test_tool_correlates_inference_and_http( + clients, + span_exporter, + tracer_provider, + kind: str, + mode: str, +) -> None: + def lookup(text: str) -> str: + assert baggage.get_baggage("request") == "test-request" + clients.http.get("https://example.test/lookup") + return clients.infer(text) + + async def alookup(text: str) -> str: + assert baggage.get_baggage("request") == "test-request" + await clients.ahttp.get("https://example.test/lookup") + await asyncio.sleep(0) + return await clients.ainfer(text) + + cls = StructuredTool if kind == "structured" else Tool + tool = cls.from_function( + func=lookup, + coroutine=alookup if mode == "async" else None, + name="lookup", + description="Look up a value.", + ) + token = context.attach(baggage.set_baggage("request", "test-request")) + try: + with tracer_provider.get_tracer("test").start_as_current_span( + "request" + ) as root: + before = context.get_current() + output = ( + tool.invoke("hello") + if mode == "sync" + else await tool.ainvoke("hello") + ) + assert output == "answer" + assert context.get_current() is before + assert trace.get_current_span() is root + finally: + context.detach(token) + + (tool_span,) = _spans(span_exporter, _LC_SCOPE) + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + http, model_http = _spans(span_exporter, _HTTP_SCOPE) + assert tool_span.name == "execute_tool lookup" + assert tool_span.parent == root.get_span_context() + _child(inference, tool_span) + _child(http, tool_span) + _child(model_http, inference) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_chat_model_correlates_sdk_and_http( + clients, + span_exporter, + asynchronous: bool, +) -> None: + model = ChatOpenAI( + model="test-model", + api_key="test", + http_client=clients.http, + http_async_client=clients.ahttp, + ) + before = context.get_current() + result = ( + await model.ainvoke("hello") if asynchronous else model.invoke("hello") + ) + assert result.content == "answer" + assert context.get_current() is before + (model_span,) = _spans(span_exporter, _LC_SCOPE) + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + (http,) = _spans(span_exporter, _HTTP_SCOPE) + _child(inference, model_span) + _child(http, inference) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["sync", "async", "executor"]) +async def test_runnable_correlates_inference( + clients, + span_exporter, + mode: str, +) -> None: + runnable = RunnableLambda( + clients.infer, afunc=clients.ainfer if mode == "async" else None + ) + before = context.get_current() + result = ( + runnable.invoke("hello") + if mode == "sync" + else await runnable.ainvoke("hello") + ) + assert result == "answer" + assert context.get_current() is before + (workflow,) = _spans(span_exporter, _LC_SCOPE) + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + _child(inference, workflow) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_vector_retriever_correlates_http( + clients, + span_exporter, + asynchronous: bool, +) -> None: + class Store(VectorStore): + @classmethod + def from_texts(cls, texts, embedding, metadatas=None, **kwargs): + raise NotImplementedError + + def similarity_search( + self, query: str, k: int = 4, **kwargs: Any + ) -> list[Document]: + clients.http.get("https://example.test/search") + return [Document(page_content="found")] + + async def asimilarity_search( + self, query: str, k: int = 4, **kwargs: Any + ) -> list[Document]: + await clients.ahttp.get("https://example.test/search") + return [Document(page_content="found")] + + retriever = Store().as_retriever() + before = context.get_current() + result = ( + await retriever.ainvoke("query") + if asynchronous + else retriever.invoke("query") + ) + assert result[0].page_content == "found" + assert context.get_current() is before + (retrieval,) = _spans(span_exporter, _LC_SCOPE) + (http,) = _spans(span_exporter, _HTTP_SCOPE) + _child(http, retrieval) + + +class _State(TypedDict): + text: str + + +def _optional_module(name: str) -> Any: + try: + return import_module(name) + except ImportError: + return None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_graph_node_and_nested_tool_correlate_http( + clients, + span_exporter, + asynchronous: bool, +) -> None: + graph_module = pytest.importorskip("langgraph.graph") + if not any( + hasattr(_optional_module(name), "set_config_context") + for name in ( + "langgraph._internal._runnable", + "langgraph.utils.runnable", + ) + ): + pytest.skip("LangGraph version has no scoped node execution helper") + + tool = StructuredTool.from_function( + func=clients.infer, + coroutine=clients.ainfer, + name="lookup", + description="Look up a value.", + ) + + def node(state: _State) -> _State: + clients.http.get("https://example.test/node") + return {"text": tool.invoke(state["text"])} + + async def anode(state: _State) -> _State: + await clients.ahttp.get("https://example.test/node") + return {"text": await tool.ainvoke(state["text"])} + + builder = graph_module.StateGraph(_State) + builder.add_node("lookup", anode if asynchronous else node) + builder.add_edge(graph_module.START, "lookup") + builder.add_edge("lookup", graph_module.END) + graph = builder.compile() + before = context.get_current() + result = ( + await graph.ainvoke({"text": "hello"}) + if asynchronous + else graph.invoke({"text": "hello"}) + ) + assert result == {"text": "answer"} + assert context.get_current() is before + tool_span, workflow = _spans(span_exporter, _LC_SCOPE) + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + node_http, inference_http = _spans(span_exporter, _HTTP_SCOPE) + _child(tool_span, workflow) + _child(node_http, workflow) + _child(inference, tool_span) + _child(inference_http, inference) + + +@pytest.mark.asyncio +async def test_concurrent_tools_do_not_share_context( + clients, span_exporter, tracer_provider +) -> None: + ready = asyncio.Event() + entered = 0 + + async def lookup(text: str) -> str: + nonlocal entered + entered += 1 + if entered == 2: + ready.set() + await asyncio.wait_for(ready.wait(), 5) + return await clients.ainfer(text) + + tool = StructuredTool.from_function( + coroutine=lookup, name="lookup", description="Look up a value." + ) + + async def request(name: str) -> None: + with tracer_provider.get_tracer("test").start_as_current_span( + name + ) as root: + assert await tool.ainvoke(name) == "answer" + assert trace.get_current_span() is root + + await asyncio.gather(request("one"), request("two")) + tools = _spans(span_exporter, _LC_SCOPE) + inferences = _spans(span_exporter, _OPENAI_SCOPE) + assert len(tools) == len(inferences) == 2 + assert tools[0].context.trace_id != tools[1].context.trace_id + for inference in inferences: + _child( + inference, + next(tool for tool in tools if tool.context == inference.parent), + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["sync", "async", "cancel", "sync_cancel"]) +@pytest.mark.parametrize("kind", ["structured", "simple"]) +async def test_execution_failure_restores_context( + clients, span_exporter, mode: str, kind: str +) -> None: + error = ( + asyncio.CancelledError() + if "cancel" in mode + else RuntimeError("tool failed") + ) + + def fail(text: str) -> str: + clients.http.get("https://example.test/fail") + raise error + + async def afail(text: str) -> str: + await clients.ahttp.get("https://example.test/fail") + await asyncio.sleep(0) + raise error + + cls = StructuredTool if kind == "structured" else Tool + tool = cls.from_function( + func=fail, coroutine=afail, name="fail", description="Fail after HTTP." + ) + before = context.get_current() + with pytest.raises(type(error)) as raised: + if mode in ("sync", "sync_cancel"): + tool.invoke("hello") + else: + await tool.ainvoke("hello") + assert raised.value is error + assert context.get_current() is before + (http,) = _spans(span_exporter, _HTTP_SCOPE) + (tool_span,) = _spans(span_exporter, _LC_SCOPE) + _child(http, tool_span) + assert tool_span.status.status_code == StatusCode.ERROR + assert tool_span.attributes[error_attributes.ERROR_TYPE] == ( + "asyncio.exceptions.CancelledError" + if "cancel" in mode + else "RuntimeError" + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["structured", "simple"]) +async def test_task_cancellation_finalizes_tool_once( + clients, + span_exporter, + tracer_provider, + kind: str, +) -> None: + started = asyncio.Event() + errors: list[BaseException] = [] + other_errors: list[BaseException] = [] + + class OtherHandler(BaseCallbackHandler): + def on_tool_error(self, error: BaseException, **kwargs: Any) -> None: + other_errors.append(error) + + async def lookup(text: str) -> str: + await clients.ahttp.get("https://example.test/lookup") + started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError as error: + errors.append(error) + raise + return text + + cls = StructuredTool if kind == "structured" else Tool + tool = cls.from_function( + func=None, + coroutine=lookup, + name="lookup", + description="Wait for cancellation.", + ) + handler = next( + h + for h in CallbackManager.configure().handlers + if isinstance(h, OpenTelemetryLangChainCallbackHandler) + ) + + async def request() -> None: + with tracer_provider.get_tracer("test").start_as_current_span( + "request" + ) as root: + before = context.get_current() + try: + await tool.ainvoke( + "hello", config={"callbacks": [OtherHandler()]} + ) + except asyncio.CancelledError as error: + assert error is errors[0] + assert error.args == ("cancelled by test",) + raise + finally: + assert context.get_current() is before + assert trace.get_current_span() is root + + task = asyncio.create_task(request()) + try: + await asyncio.wait_for(started.wait(), 5) + (run_id,) = handler._invocation_manager._invocations + task.cancel("cancelled by test") + with pytest.raises(asyncio.CancelledError): + await task + assert task.cancelled() + finally: + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + assert handler._invocation_manager._invocations == {} + assert other_errors == [] + handler.on_tool_error(errors[0], run_id=run_id) + (tool_span,) = _spans(span_exporter, _LC_SCOPE) + (http,) = _spans(span_exporter, _HTTP_SCOPE) + _child(http, tool_span) + assert tool_span.status.status_code == StatusCode.ERROR + assert ( + tool_span.attributes[error_attributes.ERROR_TYPE] + == "asyncio.exceptions.CancelledError" + ) + + +@pytest.mark.asyncio +async def test_handled_child_cancellation_does_not_fail_parent_tool( + clients, + span_exporter, +) -> None: + async def child(text: str) -> str: + raise asyncio.CancelledError() + + async def lookup(text: str) -> str: + before = context.get_current() + with pytest.raises(asyncio.CancelledError): + await RunnableLambda(child).ainvoke(text) + assert context.get_current() is before + return await clients.ainfer(text) + + tool = StructuredTool.from_function( + coroutine=lookup, + name="lookup", + description="Handle a cancelled child.", + ) + assert await tool.ainvoke("hello") == "answer" + (tool_span,) = [ + span + for span in _spans(span_exporter, _LC_SCOPE) + if span.name == "execute_tool lookup" + ] + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + _child(inference, tool_span) + assert tool_span.status.status_code != StatusCode.ERROR + assert error_attributes.ERROR_TYPE not in tool_span.attributes + + +def test_callbacks_in_different_contexts_never_attach( + clients, + tracer_provider, + span_exporter, +) -> None: + with tracer_provider.get_tracer("test").start_as_current_span( + "request" + ) as root: + start_context = copy_context() + finish_context = copy_context() + callbacks = CallbackManager.configure() + run = start_context.run( + callbacks.on_chain_start, + {"name": "LangGraph", "id": ["langgraph"]}, + {}, + ) + assert start_context.run(trace.get_current_span) is root + finish_context.run(run.on_chain_end, {}) + assert start_context.run(trace.get_current_span) is root + assert finish_context.run(trace.get_current_span) is root + assert trace.get_current_span() is root + + (workflow,) = _spans(span_exporter, _LC_SCOPE) + assert workflow.parent == root.get_span_context() + + +def test_uninstrument_restores_execution_methods( + tracer_provider, + meter_provider, + logger_provider, +) -> None: + targets = [ + (BaseTool, "run"), + (BaseTool, "arun"), + (BaseRetriever, "invoke"), + (BaseRetriever, "ainvoke"), + (BaseChatModel, "stream"), + (BaseChatModel, "astream"), + (Runnable, "_call_with_config"), + (Runnable, "_acall_with_config"), + (Runnable, "_transform_stream_with_config"), + (Runnable, "_atransform_stream_with_config"), + (Runnable, "batch"), + (Runnable, "abatch"), + (RunnableBranch, "invoke"), + (RunnableBranch, "ainvoke"), + (RunnableBranch, "stream"), + (RunnableBranch, "astream"), + (runnables_base, "set_config_context"), + (CallbackManager, "on_chain_start"), + ] + try: + from langgraph.pregel import Pregel + + targets.extend([(Pregel, "stream"), (Pregel, "astream")]) + except ImportError: + pass + originals = [getattr(owner, name) for owner, name in targets] + for _ in range(2): + with instrument( + LangChainInstrumentor(), + tracer_provider=tracer_provider, + meter_provider=meter_provider, + logger_provider=logger_provider, + ): + for (owner, name), original in zip(targets, originals): + assert getattr(owner, name) is not original + for (owner, name), original in zip(targets, originals): + assert getattr(owner, name) is original + + +def test_uninstrument_restores_the_twice_wrapped_graph_stream( + tracer_provider, + meter_provider, + logger_provider, +) -> None: + Pregel = pytest.importorskip("langgraph.pregel").Pregel + originals = {name: vars(Pregel)[name] for name in ("stream", "astream")} + with instrument( + LangChainInstrumentor(), + tracer_provider=tracer_provider, + meter_provider=meter_provider, + logger_provider=logger_provider, + ): + for name, original in originals.items(): + # The agent announcement wraps the execution boundary, which + # wraps the original. + outer = vars(Pregel)[name] + assert outer._self_wrapper is getattr( + agent_context, f"wrap_{name}" + ) + inner = outer.__wrapped__ + assert inner is not original + assert inner.__wrapped__ is original + for name, original in originals.items(): + assert vars(Pregel)[name] is original + + +def _uninstall_langgraph(monkeypatch, tmp_path, keep: str | None) -> None: + """Take LangGraph off sys.path and out of sys.modules. + + ``keep`` names a sibling distribution under the ``langgraph`` namespace + (``langgraph-checkpoint``) that stays installed without the runtime. + Every other package on the affected path entries stays reachable. + """ + for name in list(sys.modules): + if name == "langgraph" or name.startswith("langgraph."): + monkeypatch.delitem(sys.modules, name) + path: list[str] = [] + for entry in sys.path: + root = Path(entry or ".") + if not (root / "langgraph").is_dir(): + path.append(entry) + continue + shadow = tmp_path / f"site-packages-{len(path)}" + shadow.mkdir() + for child in root.iterdir(): + if ( + not child.name.startswith( + ("langgraph.", "langgraph_", "langgraph-") + ) + and child.name != "langgraph" + ): + (shadow / child.name).symlink_to(child) + if keep is not None: + (shadow / "langgraph").mkdir() + (shadow / "langgraph" / keep).symlink_to(root / "langgraph" / keep) + for info in root.glob(f"langgraph_{keep}-*.dist-info"): + (shadow / info.name).symlink_to(info) + path.append(str(shadow)) + monkeypatch.setattr(sys, "path", path) + monkeypatch.setattr(sys, "path_importer_cache", {}) + invalidate_caches() + + +@pytest.mark.parametrize( + ("module_name", "class_name", "method", "level"), + [ + ( + "langchain_core.language_models.chat_models", + "BaseChatModel", + "_agenerate_with_cache", + logging.WARNING, + ), + ( + "langchain_core.runnables.base", + "Runnable", + "_atransform_stream_with_config", + logging.WARNING, + ), + ( + "langchain_core.runnables.base", + None, + "set_config_context", + logging.WARNING, + ), + ( + "langgraph._internal._runnable", + None, + "set_config_context", + logging.WARNING, + ), + ("langgraph._internal._runnable", None, None, logging.WARNING), + ], +) +def test_instrument_skips_missing_execution_boundary( + monkeypatch, + caplog, + tracer_provider, + meter_provider, + logger_provider, + module_name: str, + class_name: str | None, + method: str | None, + level: int, +) -> None: + module = pytest.importorskip(module_name) + if method is None: + # Hide the module and everything imported under it. + for name in list(sys.modules): + if name == module_name or name.startswith(f"{module_name}."): + monkeypatch.setitem(sys.modules, name, None) + else: + owner = getattr(module, class_name) if class_name else module + monkeypatch.delattr(owner, method) + run = BaseTool.run + generate = BaseChatModel._generate_with_cache + with caplog.at_level(logging.DEBUG, logger=_LC_SCOPE): + with instrument( + LangChainInstrumentor(), + tracer_provider=tracer_provider, + meter_provider=meter_provider, + logger_provider=logger_provider, + ): + assert BaseTool.run is not run + assert BaseChatModel._generate_with_cache is not generate + assert BaseTool.run is run + assert BaseChatModel._generate_with_cache is generate + target = ".".join(filter(None, (module_name, class_name, method))) + prefix = f"Skipping execution boundary {target}{': ' if method else '.'}" + skipped = [ + record + for record in caplog.records + if record.name == f"{_LC_SCOPE}._execution_context" + and record.getMessage().startswith(prefix) + ] + assert skipped + assert {record.levelno for record in skipped} == {level} + + +# The optional library is absent, or only langgraph-checkpoint is installed: +# the namespace then resolves without the runtime. Nothing to warn about. +@pytest.mark.parametrize("installed", [None, "checkpoint"]) +def test_instrument_skips_absent_langgraph( + monkeypatch, + caplog, + tmp_path, + tracer_provider, + meter_provider, + logger_provider, + installed: str | None, +) -> None: + pytest.importorskip("langgraph") + _uninstall_langgraph(monkeypatch, tmp_path, installed) + assert (find_spec("langgraph") is not None) == (installed is not None) + with caplog.at_level(logging.DEBUG, logger=_LC_SCOPE): + with instrument( + LangChainInstrumentor(), + tracer_provider=tracer_provider, + meter_provider=meter_provider, + logger_provider=logger_provider, + ): + pass + records = [ + record + for record in caplog.records + if record.name == f"{_LC_SCOPE}._execution_context" + ] + assert {record.levelno for record in records} == {logging.DEBUG} + assert sorted(record.getMessage() for record in records) == [ + f"Skipping execution boundary {target}: langgraph is not installed" + for target in sorted( + ( + "langgraph._internal._runnable.set_config_context", + "langgraph.pregel.Pregel.astream", + "langgraph.pregel.Pregel.stream", + ) + ) + ] diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_persistence_context.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_persistence_context.py new file mode 100644 index 000000000..c894147f6 --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_persistence_context.py @@ -0,0 +1,384 @@ +# Copyright The OpenTelemetry Authors +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +from inspect import signature +from typing import Any, TypedDict + +import pytest + +pytest.importorskip("langgraph") +from langchain_core.tools import StructuredTool +from langgraph.checkpoint.memory import InMemorySaver +from langgraph.graph import END, START, StateGraph +from langgraph.pregel import Pregel +from langgraph.store.memory import InMemoryStore + +from opentelemetry import context +from opentelemetry.semconv.attributes import error_attributes + +from .test_execution_context import ( + _HTTP_SCOPE, + _LC_SCOPE, + _child, + _spans, + clients, + no_detach_errors, +) +from .test_stream_context import _assert_no_runs + +__all__ = ["clients", "no_detach_errors"] + +# LangGraph 0.6 replaced checkpoint_during with durability modes. +_durability = pytest.mark.skipif( + "durability" not in signature(Pregel.stream).parameters, + reason="durability modes need LangGraph >= 0.6", +) + + +class _State(TypedDict): + text: str + + +def _saver( + clients: Any, + failure: str | None = None, + error: BaseException | None = None, +) -> InMemorySaver: + class Saver(InMemorySaver): + def observe(self, op: str, config: dict) -> None: + clients.http.get( + f"https://example.test/checkpoint/{op}", + params={"thread": config["configurable"]["thread_id"]}, + ) + if op == failure: + raise error + + async def aobserve(self, op: str, config: dict) -> None: + await clients.ahttp.get( + f"https://example.test/checkpoint/{op}", + params={"thread": config["configurable"]["thread_id"]}, + ) + await asyncio.sleep(0) + if op == failure: + raise error + + def get_tuple(self, config: dict) -> Any: + self.observe("get", config) + return super().get_tuple(config) + + def put( + self, + config: dict, + checkpoint: Any, + metadata: Any, + new_versions: Any, + ) -> Any: + self.observe("put", config) + return super().put(config, checkpoint, metadata, new_versions) + + def put_writes( + self, config: dict, writes: Any, task_id: str, task_path: str = "" + ) -> None: + self.observe("writes", config) + super().put_writes(config, writes, task_id, task_path) + + async def aget_tuple(self, config: dict) -> Any: + await self.aobserve("get", config) + return InMemorySaver.get_tuple(self, config) + + async def aput( + self, + config: dict, + checkpoint: Any, + metadata: Any, + new_versions: Any, + ) -> Any: + await self.aobserve("put", config) + return InMemorySaver.put( + self, config, checkpoint, metadata, new_versions + ) + + async def aput_writes( + self, config: dict, writes: Any, task_id: str, task_path: str = "" + ) -> None: + await self.aobserve("writes", config) + InMemorySaver.put_writes(self, config, writes, task_id, task_path) + + return Saver() + + +def _store(clients: Any, error: BaseException | None = None) -> InMemoryStore: + class Store(InMemoryStore): + def batch(self, ops: Any) -> Any: + ops = list(ops) + for op in ops: + clients.http.get( + f"https://example.test/store/{type(op).__name__}" + ) + if error is not None: + raise error + return super().batch(ops) + + async def abatch(self, ops: Any) -> Any: + ops = list(ops) + for op in ops: + await clients.ahttp.get( + f"https://example.test/store/{type(op).__name__}" + ) + if error is not None: + raise error + return await super().abatch(ops) + + return Store() + + +def _graph( + saver: InMemorySaver, store: InMemoryStore, asynchronous: bool +) -> Any: + def lookup(text: str) -> str: + result = store.get(("memory",), text) + assert result is not None + return result.value["text"] + + async def alookup(text: str) -> str: + result = await store.aget(("memory",), text) + assert result is not None + return result.value["text"] + + tool = StructuredTool.from_function( + func=lookup, + coroutine=alookup, + name="lookup", + description="Read stored state.", + ) + + def node(state: _State) -> _State: + store.put(("memory",), state["text"], {"text": state["text"]}) + assert store.search(("memory",)) + return {"text": tool.invoke(state["text"])} + + async def anode(state: _State) -> _State: + await store.aput(("memory",), state["text"], {"text": state["text"]}) + assert await store.asearch(("memory",)) + return {"text": await tool.ainvoke(state["text"])} + + graph = StateGraph(_State) + graph.add_node("lookup", anode if asynchronous else node) + graph.add_edge(START, "lookup") + graph.add_edge("lookup", END) + return graph.compile(checkpointer=saver, store=store) + + +def _assert_persistence_parents( + clients: Any, span_exporter: Any, writes_expected: bool = True +) -> None: + lc_spans = _spans(span_exporter, _LC_SCOPE) + by_id = {span.context.span_id: span for span in lc_spans} + http_spans = { + span.context.span_id: span + for span in _spans(span_exporter, _HTTP_SCOPE) + } + operations: set[str] = set() + for request in clients.requests: + http = http_spans[ + int(request.headers["traceparent"].split("-")[2], 16) + ] + parent = by_id[http.parent.span_id] + if request.url.path.startswith("/checkpoint/"): + assert parent.name.startswith("invoke_workflow") + assert ( + parent.attributes["gen_ai.conversation.id"] + == request.url.params["thread"] + ) + elif request.url.path == "/store/GetOp": + assert parent.name == "execute_tool lookup" + else: + assert parent.name.startswith("invoke_workflow") + _child(http, parent) + operations.add(request.url.path) + assert { + "/checkpoint/get", + "/checkpoint/put", + "/store/GetOp", + "/store/PutOp", + "/store/SearchOp", + } <= operations + assert ("/checkpoint/writes" in operations) is writes_expected + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize( + "durability", + [ + # The default mode; older LangGraph checkpoints the same way. + None, + pytest.param("async", marks=_durability), + pytest.param("exit", marks=_durability), + ], +) +async def test_checkpoint_and_store_context( + clients, + span_exporter, + asynchronous: bool, + streaming: bool, + durability: str | None, +) -> None: + saver = _saver(clients) + graph = _graph(saver, _store(clients), asynchronous) + before = context.get_current() + config = {"configurable": {"thread_id": "one"}} + kwargs = {} if durability is None else {"durability": durability} + for text in ("first", "second"): + if streaming: + if asynchronous: + async for _ in graph.astream({"text": text}, config, **kwargs): + assert context.get_current() is before + else: + for _ in graph.stream({"text": text}, config, **kwargs): + assert context.get_current() is before + elif asynchronous: + assert await graph.ainvoke({"text": text}, config, **kwargs) == { + "text": text + } + else: + assert graph.invoke({"text": text}, config, **kwargs) == { + "text": text + } + assert context.get_current() is before + assert ( + InMemorySaver.get_tuple(saver, config).checkpoint[ + "channel_values" + ]["text"] + == text + ) + _assert_persistence_parents(clients, span_exporter, durability != "exit") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_shared_saver_keeps_concurrent_workflows_separate( + clients, span_exporter, asynchronous: bool +) -> None: + graph = _graph(_saver(clients), _store(clients), asynchronous) + + async def run(name: str) -> None: + config = {"configurable": {"thread_id": name}} + before = context.get_current() + if asynchronous: + assert await graph.ainvoke({"text": name}, config) == { + "text": name + } + else: + assert await asyncio.to_thread( + graph.invoke, {"text": name}, config + ) == {"text": name} + assert context.get_current() is before + + await asyncio.gather(run("one"), run("two")) + roots = [ + span + for span in _spans(span_exporter, _LC_SCOPE) + if span.parent is None + ] + assert len(roots) == 2 + assert roots[0].context.trace_id != roots[1].context.trace_id + _assert_persistence_parents(clients, span_exporter) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("operation", ["get", "put", "writes"]) +async def test_checkpoint_error_restores_context( + clients, span_exporter, asynchronous: bool, operation: str +) -> None: + error = ConnectionError("checkpoint failed") + graph = _graph( + _saver(clients, operation, error), _store(clients), asynchronous + ) + before = context.get_current() + with pytest.raises(ConnectionError) as raised: + if asynchronous: + await graph.ainvoke( + {"text": "one"}, {"configurable": {"thread_id": "one"}} + ) + else: + graph.invoke( + {"text": "one"}, {"configurable": {"thread_id": "one"}} + ) + assert raised.value is error + assert context.get_current() is before + (workflow,) = [ + span + for span in _spans(span_exporter, _LC_SCOPE) + if span.parent is None + ] + assert ( + workflow.attributes[error_attributes.ERROR_TYPE] == "ConnectionError" + ) + by_id = { + span.context.span_id: span for span in _spans(span_exporter, _LC_SCOPE) + } + for span in _spans(span_exporter, _HTTP_SCOPE): + _child(span, by_id[span.parent.span_id]) + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_store_error_restores_context( + clients, span_exporter, asynchronous: bool +) -> None: + error = ConnectionError("store failed") + graph = _graph(_saver(clients), _store(clients, error), asynchronous) + before = context.get_current() + config = {"configurable": {"thread_id": "one"}} + with pytest.raises(ConnectionError) as raised: + if asynchronous: + await graph.ainvoke({"text": "one"}, config) + else: + graph.invoke({"text": "one"}, config) + assert raised.value is error + assert context.get_current() is before + by_id = { + span.context.span_id: span for span in _spans(span_exporter, _LC_SCOPE) + } + for span in _spans(span_exporter, _HTTP_SCOPE): + _child(span, by_id[span.parent.span_id]) + (workflow,) = [span for span in by_id.values() if span.parent is None] + assert ( + workflow.attributes[error_attributes.ERROR_TYPE] == "ConnectionError" + ) + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_graph_stream_close_restores_context( + clients, span_exporter, asynchronous: bool +) -> None: + graph = _graph(_saver(clients), _store(clients), asynchronous) + config = {"configurable": {"thread_id": "one"}} + before = context.get_current() + if asynchronous: + stream = graph.astream({"text": "one"}, config) + assert await anext(stream) + assert context.get_current() is before + await stream.aclose() + else: + stream = graph.stream({"text": "one"}, config) + assert next(stream) + assert context.get_current() is before + stream.close() + assert context.get_current() is before + by_id = { + span.context.span_id: span for span in _spans(span_exporter, _LC_SCOPE) + } + for span in _spans(span_exporter, _HTTP_SCOPE): + _child(span, by_id[span.parent.span_id]) + _assert_no_runs() diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_stream_context.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_stream_context.py new file mode 100644 index 000000000..7cf6263f3 --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_stream_context.py @@ -0,0 +1,612 @@ +# 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 AsyncGeneratorType, GeneratorType +from typing import Any +from uuid import UUID, uuid4 + +import pytest +from langchain_core.callbacks import BaseCallbackHandler +from langchain_core.callbacks.manager import CallbackManager +from langchain_core.runnables import ( + RunnableGenerator, + RunnableLambda, + RunnableParallel, +) +from langchain_openai import ChatOpenAI + +from opentelemetry import context, trace +from opentelemetry.instrumentation.genai.langchain import _run_context +from opentelemetry.instrumentation.genai.langchain.callback_handler import ( + OpenTelemetryLangChainCallbackHandler, +) +from opentelemetry.semconv.attributes import error_attributes +from opentelemetry.trace import StatusCode + +from .test_execution_context import ( + _HTTP_SCOPE, + _LC_SCOPE, + _OPENAI_SCOPE, + _child, + _spans, + clients, + no_detach_errors, +) + +__all__ = ["clients", "no_detach_errors"] + + +_RESPONSES_API = "use_responses_api" in ChatOpenAI.model_fields + + +def _model(clients: Any, **kwargs: Any) -> ChatOpenAI: + return ChatOpenAI( + model="test-model", + api_key="test", + http_client=clients.http, + http_async_client=clients.ahttp, + **kwargs, + ) + + +def _assert_no_runs() -> None: + handler = next( + h + for h in CallbackManager.configure().handlers + if isinstance(h, OpenTelemetryLangChainCallbackHandler) + ) + assert handler._invocation_manager._invocations == {} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_chat_stream_is_lazy_and_restores_consumer_context( + clients, span_exporter, tracer_provider, asynchronous: bool +) -> None: + model = _model(clients) + config = {"run_id": uuid4(), "tags": ["test"]} + original_config = dict(config) + with tracer_provider.get_tracer("test").start_as_current_span( + "request" + ) as root: + before = context.get_current() + stream = ( + model.astream("hello", config=config) + if asynchronous + else model.stream("hello", config=config) + ) + assert isinstance( + stream, AsyncGeneratorType if asynchronous else GeneratorType + ) + assert clients.requests == [] + assert span_exporter.get_finished_spans() == () + if asynchronous: + first = await asyncio.create_task(anext(stream)) + else: + first = next(stream) + assert first.content == "answer" + assert context.get_current() is before + assert _spans(span_exporter, _LC_SCOPE) == [] + clients.http.get("https://example.test/consumer") + if asynchronous: + rest = [chunk async for chunk in stream] + else: + rest = list(stream) + assert "".join(chunk.content for chunk in rest) == "" + assert context.get_current() is before + assert config == original_config + (model_span,) = _spans(span_exporter, _LC_SCOPE) + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + http, consumer_http = _spans(span_exporter, _HTTP_SCOPE) + assert model_span.parent == root.get_span_context() + _child(inference, model_span) + _child(http, inference) + assert consumer_http.parent == root.get_span_context() + _assert_no_runs() + + +class _RunIdRecorder(BaseCallbackHandler): + def __init__(self) -> None: + self.run_ids: list[UUID] = [] + + def on_chat_model_start( + self, *args: Any, run_id: UUID, **kwargs: Any + ) -> None: + self.run_ids.append(run_id) + + def on_chain_start(self, *args: Any, run_id: UUID, **kwargs: Any) -> None: + self.run_ids.append(run_id) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("kind", ["model", "lambda"]) +async def test_stream_reads_config_on_first_advancement( + clients, span_exporter, asynchronous: bool, kind: str +) -> None: + runnable = ( + _model(clients) + if kind == "model" + else RunnableLambda(lambda text: [text], afunc=None) + ) + recorder = _RunIdRecorder() + config: dict[str, Any] = {"run_id": uuid4(), "callbacks": [recorder]} + stream = ( + runnable.astream("hello", config=config) + if asynchronous + else runnable.stream("hello", config=config) + ) + # Nothing ran yet, so a caller may still edit the config it passed. + assert clients.requests == [] + late_run_id = uuid4() + config["run_id"] = late_run_id + if asynchronous: + assert [chunk async for chunk in stream] + else: + assert list(stream) + assert recorder.run_ids == [late_run_id] + (span,) = _spans(span_exporter, _LC_SCOPE) + assert span.name.startswith( + "chat " if kind == "model" else "invoke_workflow " + ) + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.skipif( + not _RESPONSES_API, reason="langchain-openai predates the Responses API" +) +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_responses_stream_restores_consumer_context( + clients, span_exporter, tracer_provider, asynchronous: bool +) -> None: + model = _model(clients, use_responses_api=True) + with tracer_provider.get_tracer("test").start_as_current_span( + "request" + ) as root: + before = context.get_current() + stream = ( + model.astream("hello") if asynchronous else model.stream("hello") + ) + if asynchronous: + first = await asyncio.create_task(anext(stream)) + else: + first = next(stream) + assert context.get_current() is before + clients.http.get("https://example.test/consumer") + + async def drain() -> list[Any]: + return [chunk async for chunk in stream] + + # Later reads move to another task, so each read must attach and + # detach within its own frame. + rest = ( + await asyncio.create_task(drain()) + if asynchronous + else list(stream) + ) + assert "answer" in "".join( + str(chunk.content) for chunk in (first, *rest) + ) + assert context.get_current() is before + (model_span,) = _spans(span_exporter, _LC_SCOPE) + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + http, consumer_http = _spans(span_exporter, _HTTP_SCOPE) + assert model_span.parent == root.get_span_context() + _child(inference, model_span) + _child(http, inference) + assert consumer_http.parent == root.get_span_context() + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("ending", ["provider_error", "caller_error", "close"]) +async def test_chat_stream_failure_and_close( + clients, span_exporter, asynchronous: bool, ending: str +) -> None: + error = ConnectionError("stream failed") + if ending == "provider_error": + clients.stream_error = error + model = _model(clients) + stream = model.astream("hello") if asynchronous else model.stream("hello") + before = context.get_current() + if ending == "caller_error": + with pytest.raises(ConnectionError) as raised: + if asynchronous: + async with stream: + await anext(stream) + raise error + else: + with stream: + next(stream) + raise error + assert raised.value is error + else: + if asynchronous: + await anext(stream) + else: + next(stream) + assert context.get_current() is before + if ending == "close": + if asynchronous: + await stream.aclose() + else: + stream.close() + else: + with pytest.raises(ConnectionError) as raised: + if asynchronous: + await anext(stream) + else: + next(stream) + assert raised.value is error + assert context.get_current() is before + (model_span,) = _spans(span_exporter, _LC_SCOPE) + if ending != "close": + assert model_span.status.status_code == StatusCode.ERROR + assert ( + model_span.attributes[error_attributes.ERROR_TYPE] + == "ConnectionError" + ) + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_interleaved_chat_streams_do_not_share_parents( + clients, span_exporter, tracer_provider, asynchronous: bool +) -> None: + model = _model(clients) + one = model.astream("one") if asynchronous else model.stream("one") + two = model.astream("two") if asynchronous else model.stream("two") + with tracer_provider.get_tracer("test").start_as_current_span( + "one" + ) as first_parent: + if asynchronous: + await anext(one) + else: + next(one) + with tracer_provider.get_tracer("test").start_as_current_span( + "two" + ) as second_parent: + if asynchronous: + await anext(two) + assert [chunk async for chunk in one] + assert [chunk async for chunk in two] + else: + next(two) + assert list(one) + assert list(two) + assert trace.get_current_span() is second_parent + assert trace.get_current_span() is first_parent + model_spans = _spans(span_exporter, _LC_SCOPE) + assert {span.parent.span_id for span in model_spans} == { + first_parent.get_span_context().span_id, + second_parent.get_span_context().span_id, + } + inferences = _spans(span_exporter, _OPENAI_SCOPE) + assert len(inferences) == 2 + for inference in inferences: + _child( + inference, + next( + span + for span in model_spans + if span.context == inference.parent + ), + ) + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("kind", ["lambda", "generator"]) +async def test_runnable_stream_correlates_each_step( + clients, span_exporter, tracer_provider, asynchronous: bool, kind: str +) -> None: + def produce(text: str) -> Iterator[str]: + clients.http.get("https://example.test/produce/one") + yield text + clients.http.get("https://example.test/produce/two") + yield "done" + + async def aproduce(text: str) -> AsyncIterator[str]: + await clients.ahttp.get("https://example.test/produce/one") + yield text + await clients.ahttp.get("https://example.test/produce/two") + yield "done" + + def transform(inputs: Iterator[str]) -> Iterator[str]: + for text in inputs: + yield from produce(text) + + async def atransform(inputs: AsyncIterator[str]) -> AsyncIterator[str]: + async for text in inputs: + async for chunk in aproduce(text): + yield chunk + + if kind == "lambda": + runnable = RunnableLambda(produce, afunc=aproduce) + else: + runnable = RunnableGenerator(transform, atransform=atransform) + with tracer_provider.get_tracer("test").start_as_current_span( + "request" + ) as root: + before = context.get_current() + stream = ( + runnable.astream("hello") + if asynchronous + else runnable.stream("hello") + ) + assert clients.requests == [] + if asynchronous: + assert await anext(stream) == "hello" + else: + assert next(stream) == "hello" + assert context.get_current() is before + clients.http.get("https://example.test/consumer") + if asynchronous: + assert [chunk async for chunk in stream] == ["done"] + else: + assert list(stream) == ["done"] + assert context.get_current() is before + (workflow,) = _spans(span_exporter, _LC_SCOPE) + one, consumer, two = _spans(span_exporter, _HTTP_SCOPE) + _child(one, workflow) + _child(two, workflow) + assert consumer.parent == root.get_span_context() + _assert_no_runs() + + +@pytest.mark.asyncio +async def test_runnable_stream_task_cancellation( + clients, span_exporter +) -> None: + waiting = asyncio.Event() + errors: list[BaseException] = [] + + async def produce(text: str) -> AsyncIterator[str]: + yield text + await clients.ahttp.get("https://example.test/waiting") + waiting.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError as error: + errors.append(error) + raise + + stream = RunnableLambda(produce).astream("hello") + before = context.get_current() + assert await anext(stream) == "hello" + + async def consume() -> None: + try: + await anext(stream) + except asyncio.CancelledError as error: + assert error is errors[0] + raise + finally: + assert context.get_current() is before + + task = asyncio.create_task(consume()) + try: + await asyncio.wait_for(waiting.wait(), 5) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + finally: + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + assert context.get_current() is before + (workflow,) = _spans(span_exporter, _LC_SCOPE) + (http,) = _spans(span_exporter, _HTTP_SCOPE) + _child(http, workflow) + assert ( + workflow.attributes[error_attributes.ERROR_TYPE] + == "asyncio.exceptions.CancelledError" + ) + _assert_no_runs() + + +@pytest.mark.asyncio +async def test_chat_stream_task_cancellation(clients, span_exporter) -> None: + waiting = asyncio.Event() + clients.stream_waiting = waiting + stream = _model(clients).astream("hello") + before = context.get_current() + assert (await anext(stream)).content == "answer" + + async def consume() -> None: + try: + await anext(stream) + finally: + assert context.get_current() is before + + task = asyncio.create_task(consume()) + try: + await asyncio.wait_for(waiting.wait(), 5) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + finally: + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + assert context.get_current() is before + (model_span,) = _spans(span_exporter, _LC_SCOPE) + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + _child(inference, model_span) + assert ( + model_span.attributes[error_attributes.ERROR_TYPE] + == "asyncio.exceptions.CancelledError" + ) + _assert_no_runs() + + +@pytest.mark.asyncio +async def test_chat_ainvoke_task_cancellation(clients, span_exporter) -> None: + waiting = asyncio.Event() + clients.stream_waiting = waiting + model = _model(clients, streaming=True) + before = context.get_current() + + async def request() -> None: + try: + await model.ainvoke("hello") + finally: + assert context.get_current() is before + + task = asyncio.create_task(request()) + try: + await asyncio.wait_for(waiting.wait(), 5) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + finally: + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + assert context.get_current() is before + (model_span,) = _spans(span_exporter, _LC_SCOPE) + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + _child(inference, model_span) + assert ( + model_span.attributes[error_attributes.ERROR_TYPE] + == "asyncio.exceptions.CancelledError" + ) + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error", [KeyboardInterrupt, SystemExit]) +async def test_chat_ainvoke_interrupt( + clients, span_exporter, error: type[BaseException] +) -> None: + clients.stream_error = error() + model = _model(clients, streaming=True) + restored: list[bool] = [] + + async def request() -> None: + before = context.get_current() + try: + await model.ainvoke("hello") + finally: + restored.append(context.get_current() is before) + + # Raised in a gather child, either leaves the loop from the task step + # before the gather resolves, and asyncio.run then cancels the caller. + # Run it on a loop of its own so this test's loop is not the one left. + with pytest.raises(error): + await asyncio.to_thread(asyncio.run, request()) + assert restored == [True] + (model_span,) = _spans(span_exporter, _LC_SCOPE) + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + _child(inference, model_span) + assert model_span.status.status_code == StatusCode.ERROR + assert model_span.attributes[error_attributes.ERROR_TYPE] == error.__name__ + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error", [KeyboardInterrupt, SystemExit]) +async def test_chat_invoke_interrupt( + clients, span_exporter, error: type[BaseException] +) -> None: + # generate catches the BaseException around _generate_with_cache and + # reports it, so the callback handler ends this run. + clients.stream_error = error() + model = _model(clients, streaming=True) + before = context.get_current() + with pytest.raises(error): + model.invoke("hello") + assert context.get_current() is before + (model_span,) = _spans(span_exporter, _LC_SCOPE) + (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + _child(inference, model_span) + assert model_span.status.status_code == StatusCode.ERROR + assert model_span.attributes[error_attributes.ERROR_TYPE] == error.__name__ + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("kind", ["lambda", "generator"]) +async def test_runnable_stream_close( + clients, span_exporter, asynchronous: bool, kind: str +) -> None: + def produce(text: str) -> Iterator[str]: + clients.http.get("https://example.test/produce") + yield text + yield "done" + + async def aproduce(text: str) -> AsyncIterator[str]: + await clients.ahttp.get("https://example.test/produce") + yield text + yield "done" + + def transform(inputs: Iterator[str]) -> Iterator[str]: + for text in inputs: + yield from produce(text) + + async def atransform(inputs: AsyncIterator[str]) -> AsyncIterator[str]: + async for text in inputs: + async for chunk in aproduce(text): + yield chunk + + runnable = ( + RunnableLambda(produce, afunc=aproduce) + if kind == "lambda" + else RunnableGenerator(transform, atransform=atransform) + ) + before = context.get_current() + if asynchronous: + stream = runnable.astream("hello") + assert await anext(stream) == "hello" + await stream.aclose() + else: + stream = runnable.stream("hello") + assert next(stream) == "hello" + stream.close() + assert context.get_current() is before + (workflow,) = _spans(span_exporter, _LC_SCOPE) + (http,) = _spans(span_exporter, _HTTP_SCOPE) + _child(http, workflow) + _assert_no_runs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("composite", ["sequence", "parallel"]) +async def test_composite_stream_attaches_each_context_once( + clients, span_exporter, monkeypatch, asynchronous: bool, composite: str +) -> None: + attached: list[bool] = [] + original = _run_context.attach + + def attach(ctx: context.Context) -> object: + # Attaching the current context again is a redundant token. + attached.append(context.get_current() is ctx) + return original(ctx) + + monkeypatch.setattr(_run_context, "attach", attach) + steps = [RunnableLambda(lambda text: text) for _ in range(2)] + chain = ( + steps[0] | steps[1] + if composite == "sequence" + else RunnableParallel(one=steps[0], two=steps[1]) + ) + before = context.get_current() + if asynchronous: + chunks = [chunk async for chunk in chain.astream("hello")] + else: + chunks = list(chain.stream("hello")) + assert chunks + assert context.get_current() is before + assert attached + assert not any(attached) + _assert_no_runs() 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..30fcde7a4 --- /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: 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/chat_wrappers.py b/instrumentation/opentelemetry-instrumentation-genai-openai/src/opentelemetry/instrumentation/genai/openai/chat_wrappers.py index f2d456923..609dd9e1e 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 @@ -5,6 +5,7 @@ import json import logging +from contextlib import AbstractContextManager from openai import AsyncStream, Stream from openai.types.chat import ChatCompletionChunk @@ -43,6 +44,9 @@ class _ChatStreamMixin: _self_cached_prompt_tokens: int | None _self_reasoning_tokens: int | None + def _execution_context(self) -> AbstractContextManager[None]: + return self._self_invocation.activate() + def _set_response_model(self, chunk: ChatCompletionChunk) -> None: # Set eagerly so the per-chunk streaming timing metrics carry # gen_ai.response.model. @@ -212,6 +216,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 +239,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..5c15891cc 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 @@ -132,6 +133,12 @@ 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 _execution_context(self) -> AbstractContextManager[None]: + return self._self_invocation.activate() def _stop( self, result: ParsedResponse[TextFormatT] | Response | None @@ -199,7 +206,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 +425,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_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..adf005fbc --- /dev/null +++ b/util/opentelemetry-util-genai/.changelog/817.added @@ -0,0 +1 @@ +Add the `_execution_context()` hook on `SyncStreamWrapper`/`AsyncStreamWrapper` to scope each stream read and cleanup, including generator `send`/`throw` and `asend`/`athrow`. `finalize_on_close`/`finalize_on_aclose` accept an optional `execution_context` so a close through the HTTP response runs inside the same scope, and the stream manager wrappers run the wrapped manager's exit inside it. diff --git a/util/opentelemetry-util-genai/AGENTS.md b/util/opentelemetry-util-genai/AGENTS.md index f336c224c..b58f8959a 100644 --- a/util/opentelemetry-util-genai/AGENTS.md +++ b/util/opentelemetry-util-genai/AGENTS.md @@ -70,14 +70,20 @@ 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 +or `close()`. Subclasses pass the SDK stream to `super().__init__(stream)` and implement four hooks: +- `_execution_context()` — return a fresh context manager for each stream read and cleanup + operation; for invocation-backed streams return `invocation.activate()`. Context is restored + before the chunk is returned to the consumer, so the token is created and released in the + same frame. - `_process_chunk(chunk)` — accumulate per-chunk state (e.g. response model, finish reasons, token usage, streamed content) onto the invocation. - `_on_stream_end()` — finalize on success; set the accumulated response attributes and call @@ -89,8 +95,12 @@ class MyStreamWrapper(SyncStreamWrapper[Chunk]): def __init__(self, stream, invocation, capture_content): super().__init__(stream) self._self_invocation = invocation + invocation.suspend() ... + def _execution_context(self): + return self._self_invocation.activate() + def _process_chunk(self, chunk): ... # accumulate state def _on_stream_end(self): self._self_invocation.stop() diff --git a/util/opentelemetry-util-genai/README.rst b/util/opentelemetry-util-genai/README.rst index 729621246..e74da9411 100644 --- a/util/opentelemetry-util-genai/README.rst +++ b/util/opentelemetry-util-genai/README.rst @@ -35,6 +35,20 @@ to manage context: ambient context is used. +Stream execution context +------------------------ + +Subclasses of ``SyncStreamWrapper`` and ``AsyncStreamWrapper`` can override +``_execution_context()`` to return a fresh context manager for each stream read +and cleanup operation. 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, call ``invocation.suspend()`` before returning the +stream and return ``invocation.activate()`` from ``_execution_context()``. +Keep finalization in the existing ``_on_stream_end`` and ``_on_stream_error`` hooks. + + Modalities ---------- @@ -209,4 +223,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/stream.py b/util/opentelemetry-util-genai/src/opentelemetry/util/genai/stream.py index d93848151..cb81d8d5f 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, @@ -81,6 +83,10 @@ 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.""" + return nullcontext() + def _finalize_success(self) -> None: if self._self_finalized: return @@ -116,9 +122,10 @@ 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)`` and implement the four telemetry hooks: + ``_execution_context`` for the context active around each read and + cleanup, ``_process_chunk`` for per-chunk state, ``_on_stream_end`` for + successful finalization, and ``_on_stream_error`` for failure finalization. Users should consume subclasses as normal streams, for example with ``for chunk in wrapper`` or ``with wrapper``. The hook methods are called @@ -164,27 +171,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__`` @@ -193,21 +202,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( @@ -220,9 +248,10 @@ 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)`` and implement the four telemetry hooks: + ``_execution_context`` for the context active around each read and + cleanup, ``_process_chunk`` for per-chunk state, ``_on_stream_end`` for + successful finalization, and ``_on_stream_error`` for failure finalization. Users should consume subclasses as normal async streams, for example with ``async for chunk in wrapper`` or ``async with wrapper``. The hook methods @@ -266,20 +295,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. @@ -309,45 +339,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 @@ -377,7 +410,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__`` @@ -386,22 +427,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]): @@ -425,22 +472,8 @@ def __init__( 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 _execution_context(self) -> AbstractContextManager[None]: + return self._self_tool_invocation.activate() def __del__(self) -> None: try: @@ -493,22 +526,8 @@ def __del__(self) -> None: 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 _execution_context(self) -> AbstractContextManager[None]: + return self._self_tool_invocation.activate() def _process_chunk(self, chunk: ChunkT) -> None: if self._self_tool_invocation.should_capture_content: @@ -528,32 +547,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. @@ -562,15 +595,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( @@ -633,24 +677,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( @@ -700,25 +754,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 6feb527f6..ecaf635be 100644 --- a/util/opentelemetry-util-genai/tests/test_stream.py +++ b/util/opentelemetry-util-genai/tests/test_stream.py @@ -6,6 +6,10 @@ import asyncio 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 @@ -840,6 +844,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 @@ -1607,3 +1682,250 @@ 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 diff --git a/uv.lock b/uv.lock index 76e95839c..34bc0dc86 100644 --- a/uv.lock +++ b/uv.lock @@ -3885,7 +3885,7 @@ instruments = [ [package.metadata] requires-dist = [ - { name = "langchain", marker = "extra == 'instruments'", specifier = ">=0.3.21,<2" }, + { name = "langchain", marker = "extra == 'instruments'", specifier = ">=0.3.22,<2" }, { name = "opentelemetry-instrumentation", specifier = ">=0.64b0,<1" }, { name = "opentelemetry-util-genai", editable = "util/opentelemetry-util-genai" }, ]