From 76ebdcadc3b0ceaff0955a3201930579b023ce6b Mon Sep 17 00:00:00 2001 From: Liudmila Molkova Date: Tue, 29 Sep 2026 18:11:12 -0700 Subject: [PATCH 1/4] Prototype LangChain execution-scoped context propagation (LangChain part) Keep callback-managed spans detached and activate their context around framework execution boundaries, including custom tools and retrievers, streaming, and graph persistence. Add inference and HTTP correlation tests covering concurrency, cancellation, and context restoration. Assisted-by: GPT-6 --- .../pyproject.toml | 2 +- .../genai/langchain/__init__.py | 62 +- .../genai/langchain/_execution_context.py | 229 +++++++ .../genai/langchain/_run_context.py | 226 ++++++ .../genai/langchain/callback_handler.py | 14 +- .../tests/requirements.latest.txt | 5 +- .../tests/requirements.oldest.txt | 8 + .../tests/test_agent_classification_corpus.py | 2 +- .../tests/test_callback_handler.py | 60 +- .../tests/test_custom_execution_context.py | 214 ++++++ .../tests/test_execution_context.py | 648 ++++++++++++++++++ .../tests/test_persistence_context.py | 371 ++++++++++ .../tests/test_stream_context.py | 390 +++++++++++ 13 files changed, 2138 insertions(+), 93 deletions(-) create mode 100644 instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/_execution_context.py create mode 100644 instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/_run_context.py create mode 100644 instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_custom_execution_context.py create mode 100644 instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_execution_context.py create mode 100644 instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_persistence_context.py create mode 100644 instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_stream_context.py diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/pyproject.toml b/instrumentation/opentelemetry-instrumentation-genai-langchain/pyproject.toml index 9955207e7..5312f9539 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/pyproject.toml +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/pyproject.toml @@ -26,7 +26,7 @@ classifiers = [ ] dependencies = [ "opentelemetry-instrumentation >= 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-langchain/src/opentelemetry/instrumentation/genai/langchain/__init__.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/__init__.py index 218969449..32a206958 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 @@ -127,6 +131,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 +143,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 +155,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..740e4b103 --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/_execution_context.py @@ -0,0 +1,229 @@ +# Copyright The OpenTelemetry Authors +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from collections.abc import Callable, Iterator +from contextlib import contextmanager +from contextvars import Context as PythonContext +from functools import wraps +from importlib import import_module +from inspect import iscoroutinefunction, signature +from typing import Any +from uuid import UUID + +from wrapt import wrap_function_wrapper + +from opentelemetry.context import attach, detach +from opentelemetry.instrumentation.genai.langchain._run_context import ( + _astart_run, + _start_run, + _wrap_run, + _wrap_stream, +) +from opentelemetry.instrumentation.genai.langchain.invocation_manager import ( + _InvocationManager, +) +from opentelemetry.instrumentation.utils import unwrap + +__all__ = ["_ExecutionContext"] + +_METHODS = ( + ( + "langchain_core.language_models.chat_models", + "BaseChatModel", + "_generate_with_cache", + ), + ( + "langchain_core.language_models.chat_models", + "BaseChatModel", + "_agenerate_with_cache", + ), + ("langchain_core.runnables.base", "RunnableLambda", "_invoke"), + ("langchain_core.runnables.base", "RunnableLambda", "_ainvoke"), +) + + +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: + yield + return + token = attach(context) + try: + yield + finally: + detach(token) + + 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: + with self._activate(run_id_for(instance, args, kwargs)): + return await wrapped(*args, **kwargs) + + return asynchronous if iscoroutinefunction(original) else sync + + @contextmanager + def _graph_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: + # LangGraph executes nodes in this copied context. Restore the token + # in that same context, never in a callback or the consuming task. + token = context.run(attach, parent) if parent is not None else None + try: + yield context + finally: + if token is not None: + context.run(detach, token) + + def instrument(self) -> None: + for module_name, class_name, method in _METHODS: + cls = getattr(import_module(module_name), class_name) + wrap_function_wrapper( + cls, method, self._wrap(getattr(cls, method)) + ) + self._patched.append((cls, method)) + + 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, + ), + ): + cls = getattr(import_module(module_name), class_name) + for method, asynchronous in zip(methods, (False, True)): + wrap_function_wrapper( + cls, + method, + _wrap_run( + getattr(cls, method), + self._invocations, + on_error, + asynchronous, + ), + ) + self._patched.append((cls, method)) + + for class_name, wrapper in ( + ("CallbackManager", _start_run), + ("AsyncCallbackManager", _astart_run), + ): + cls = getattr( + import_module("langchain_core.callbacks.manager"), class_name + ) + for method in ( + "on_chat_model_start", + "on_chain_start", + "on_tool_start", + "on_retriever_start", + ): + wrap_function_wrapper(cls, method, wrapper) + self._patched.append((cls, method)) + + 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"), + ), + ("langgraph.pregel", "Pregel", ("stream", "astream")), + ): + try: + cls = getattr(import_module(module_name), class_name) + except ImportError: + continue + for method, asynchronous in zip(methods, (False, True)): + wrap_function_wrapper( + cls, + method, + _wrap_stream( + getattr(cls, method), self._invocations, asynchronous + ), + ) + self._patched.append((cls, method)) + + try: + module = import_module("langgraph._internal._runnable") + except ImportError: + return + if hasattr(module, "set_config_context"): + wrap_function_wrapper( + module, "set_config_context", self._graph_context + ) + self._patched.append((module, "set_config_context")) + + 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..b4df8cc64 --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/src/opentelemetry/instrumentation/genai/langchain/_run_context.py @@ -0,0 +1,226 @@ +# Copyright The OpenTelemetry Authors +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import logging +from collections.abc import Callable, Iterator +from contextlib import AbstractContextManager, ExitStack, contextmanager +from contextvars import ContextVar +from dataclasses import dataclass +from functools import wraps +from inspect import signature +from typing import Any +from uuid import UUID, uuid4 + +from langchain_core.callbacks.manager import BaseRunManager + +from opentelemetry.context import Context, attach, detach +from opentelemetry.instrumentation.genai.langchain.invocation_manager import ( + _InvocationManager, +) +from opentelemetry.util.genai.stream import ( + AsyncStreamWrapper, + SyncStreamWrapper, +) + +__all__ = ["_astart_run", "_start_run", "_wrap_run", "_wrap_stream"] + +_logger = logging.getLogger(__name__) + + +class _RunScope: + def __init__( + self, + run_id: UUID, + invocations: _InvocationManager, + on_error: Callable[..., None] | None = None, + ) -> None: + 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) + if self.context is not None: + stack.callback(detach, attach(self.context)) + yield + + def finish(self, error: BaseException | None = None) -> None: + 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. + read.scope.context = 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 _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: + bound = parameters.bind(instance, *args, **kwargs) + config = dict(bound.arguments.get("config") or {}) + run_id = config.get("run_id") or uuid4() + config["run_id"] = run_id + bound.arguments["config"] = config + stream = wrapped(*bound.args[1:], **bound.kwargs) + scope = _RunScope(run_id, invocations) + cls = _AsyncContextStream if asynchronous else _SyncContextStream + return cls(stream, scope) + + return wrapper + + +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 uuid4() + 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 66319ed52..080ee10f2 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 @@ -163,12 +163,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 @@ -218,7 +216,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) @@ -252,7 +250,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) @@ -278,7 +276,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( @@ -439,7 +437,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 @@ -726,7 +724,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 @@ -790,7 +788,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/tests/requirements.latest.txt b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.latest.txt index c47957ff3..0437fbeab 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.110 -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..9cce2035e 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.oldest.txt +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.oldest.txt @@ -25,3 +25,11 @@ 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 0c686fe6a..59043953d 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, ) @@ -1601,7 +1601,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, ) @@ -1667,7 +1667,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): @@ -1684,7 +1684,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): @@ -1705,7 +1705,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): @@ -1723,7 +1723,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): @@ -1763,7 +1763,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, ) @@ -3584,32 +3584,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 @@ -3628,11 +3603,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, @@ -3651,7 +3626,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([]) @@ -3661,7 +3635,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..545910a81 --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_custom_execution_context.py @@ -0,0 +1,214 @@ +# Copyright The OpenTelemetry Authors +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +from collections.abc import Awaitable, Callable +from contextlib import nullcontext +from typing import Any + +import pytest +from langchain_core.documents import Document +from langchain_core.retrievers import BaseRetriever +from langchain_core.tools import BaseTool + +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"] + + +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() 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..c517847dd --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_execution_context.py @@ -0,0 +1,648 @@ +# Copyright The OpenTelemetry Authors +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +import json +from collections.abc import AsyncIterator, Iterator +from contextlib import ExitStack +from contextvars import copy_context +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, RunnableLambda +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 +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() + + 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=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 + + +@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") + helper = pytest.importorskip("langgraph._internal._runnable") + if not hasattr(helper, "set_config_context"): + 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, "_transform_stream_with_config"), + (Runnable, "_atransform_stream_with_config"), + (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 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..407acdd74 --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_persistence_context.py @@ -0,0 +1,371 @@ +# Copyright The OpenTelemetry Authors +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +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.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"] + + +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", ["sync", "async", "exit"]) +async def test_checkpoint_and_store_context( + clients, + span_exporter, + asynchronous: bool, + streaming: bool, + durability: str, +) -> None: + saver = _saver(clients) + graph = _graph(saver, _store(clients), asynchronous) + before = context.get_current() + config = {"configurable": {"thread_id": "one"}} + for text in ("first", "second"): + if streaming: + if asynchronous: + async for _ in graph.astream( + {"text": text}, config, durability=durability + ): + assert context.get_current() is before + else: + for _ in graph.stream( + {"text": text}, config, durability=durability + ): + assert context.get_current() is before + elif asynchronous: + assert await graph.ainvoke( + {"text": text}, config, durability=durability + ) == {"text": text} + else: + assert graph.invoke( + {"text": text}, config, durability=durability + ) == {"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..0930fd101 --- /dev/null +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_stream_context.py @@ -0,0 +1,390 @@ +# 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 uuid4 + +import pytest +from langchain_core.callbacks.manager import CallbackManager +from langchain_core.runnables import RunnableGenerator, RunnableLambda +from langchain_openai import ChatOpenAI + +from opentelemetry import context, trace +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"] + + +def _model(clients: Any) -> ChatOpenAI: + return ChatOpenAI( + model="test-model", + api_key="test", + http_client=clients.http, + http_async_client=clients.ahttp, + ) + + +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() + + +@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 +@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() From 72b58396daec99a84f508d9ccb65daacdb0e4e25 Mon Sep 17 00:00:00 2001 From: Sangkyoon Nam Date: Sun, 11 Oct 2026 15:06:11 +0900 Subject: [PATCH 2/4] fix(langchain): activate context at execution boundaries instead of in callbacks --- README.md | 2 +- .../README.rst | 14 +- .../pyproject.toml | 2 +- .../genai/langchain/__init__.py | 2 + .../genai/langchain/_execution_context.py | 246 ++++++++++++++---- .../genai/langchain/_run_context.py | 140 ++++++++-- .../genai/langchain/package.py | 2 +- .../tests/requirements.oldest.txt | 8 +- uv.lock | 2 +- 9 files changed, 345 insertions(+), 73 deletions(-) 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/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 5312f9539..c9ff6829d 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/pyproject.toml +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/pyproject.toml @@ -31,7 +31,7 @@ dependencies = [ [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 32a206958..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 @@ -124,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 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 index 740e4b103..efbd65c48 100644 --- 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 @@ -3,21 +3,25 @@ from __future__ import annotations -from collections.abc import Callable, Iterator +import logging +from collections.abc import Callable, Iterator, Sequence from contextlib import contextmanager from contextvars import Context as PythonContext -from functools import wraps +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 +from typing import Any, cast from uuid import UUID from wrapt import wrap_function_wrapper -from opentelemetry.context import attach, detach +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, ) @@ -28,6 +32,8 @@ __all__ = ["_ExecutionContext"] +_logger = logging.getLogger(__name__) + _METHODS = ( ( "langchain_core.language_models.chat_models", @@ -39,8 +45,6 @@ "BaseChatModel", "_agenerate_with_cache", ), - ("langchain_core.runnables.base", "RunnableLambda", "_invoke"), - ("langchain_core.runnables.base", "RunnableLambda", "_ainvoke"), ) @@ -59,7 +63,7 @@ def __init__( @contextmanager def _activate(self, run_id: UUID | None) -> Iterator[None]: context = self._invocations.get_parent_context(run_id) - if context is None: + if context is None or get_current() is context: yield return token = attach(context) @@ -68,6 +72,45 @@ def _activate(self, run_id: UUID | None) -> Iterator[None]: 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) @@ -96,13 +139,23 @@ async def asynchronous( args: tuple[Any, ...], kwargs: dict[str, Any], ) -> Any: - with self._activate(run_id_for(instance, args, kwargs)): - return await wrapped(*args, **kwargs) + 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 _graph_context( + def _config_context( self, wrapped: Callable[..., Any], instance: Any, @@ -115,22 +168,67 @@ def _graph_context( getattr(callbacks, "parent_run_id", None) ) with wrapped(*args, **kwargs) as context: - # LangGraph executes nodes in this copied context. Restore the token - # in that same context, never in a callback or the consuming task. - token = context.run(attach, parent) if parent is not None else None + # 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: - cls = getattr(import_module(module_name), class_name) - wrap_function_wrapper( - cls, method, self._wrap(getattr(cls, method)) - ) - self._patched.append((cls, method)) + self._patch(module_name, class_name, method, self._wrap) for module_name, class_name, methods, on_error in ( ( @@ -146,35 +244,72 @@ def instrument(self) -> None: self._on_retriever_error, ), ): - cls = getattr(import_module(module_name), class_name) for method, asynchronous in zip(methods, (False, True)): - wrap_function_wrapper( - cls, + 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, - _wrap_run( - getattr(cls, method), - self._invocations, - on_error, - asynchronous, + partial( + _wrap_call, + invocations=self._invocations, + asynchronous=asynchronous, ), ) - self._patched.append((cls, method)) + + # 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), ): - cls = getattr( - import_module("langchain_core.callbacks.manager"), class_name - ) for method in ( "on_chat_model_start", "on_chain_start", "on_tool_start", "on_retriever_start", ): - wrap_function_wrapper(cls, method, wrapper) - self._patched.append((cls, method)) + self._patch( + "langchain_core.callbacks.manager", + class_name, + method, + lambda _original, wrapper=wrapper: wrapper, + ) for module_name, class_name, methods in ( ( @@ -197,31 +332,44 @@ def instrument(self) -> None: "RunnableLambda", ("stream", "astream"), ), + ( + "langchain_core.runnables.branch", + "RunnableBranch", + ("stream", "astream"), + ), ("langgraph.pregel", "Pregel", ("stream", "astream")), ): - try: - cls = getattr(import_module(module_name), class_name) - except ImportError: - continue for method, asynchronous in zip(methods, (False, True)): - wrap_function_wrapper( - cls, + self._patch( + module_name, + class_name, method, - _wrap_stream( - getattr(cls, method), self._invocations, asynchronous + partial( + _wrap_stream, + invocations=self._invocations, + asynchronous=asynchronous, ), ) - self._patched.append((cls, method)) - try: - module = import_module("langgraph._internal._runnable") - except ImportError: - return - if hasattr(module, "set_config_context"): - wrap_function_wrapper( - module, "set_config_context", self._graph_context + # 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, ) - self._patched.append((module, "set_config_context")) def uninstrument(self) -> None: for owner, name in reversed(self._patched): 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 index b4df8cc64..474a476fe 100644 --- 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 @@ -4,18 +4,31 @@ from __future__ import annotations import logging -from collections.abc import Callable, Iterator +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 signature +from inspect import BoundArguments, signature from typing import Any -from uuid import UUID, uuid4 +from uuid import UUID from langchain_core.callbacks.manager import BaseRunManager -from opentelemetry.context import Context, attach, detach +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, ) @@ -24,7 +37,14 @@ SyncStreamWrapper, ) -__all__ = ["_astart_run", "_start_run", "_wrap_run", "_wrap_stream"] +__all__ = [ + "_RunScope", + "_astart_run", + "_start_run", + "_wrap_call", + "_wrap_run", + "_wrap_stream", +] _logger = logging.getLogger(__name__) @@ -32,10 +52,11 @@ class _RunScope: def __init__( self, - run_id: UUID, + 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 @@ -46,11 +67,24 @@ def activate(self) -> Iterator[None]: with ExitStack() as stack: token = _active_run.set(_RunActivation(self, stack)) stack.callback(_active_run.reset, token) - if self.context is not None: + # 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: @@ -88,8 +122,10 @@ def _started(result: BaseRunManager | list[BaseRunManager]) -> None: 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 - read.stack.callback(detach, attach(context)) + if get_current() is not context: + read.stack.callback(detach, attach(context)) def _start_run( @@ -154,6 +190,35 @@ 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, @@ -168,19 +233,60 @@ def wrapper( args: tuple[Any, ...], kwargs: dict[str, Any], ) -> Any: + scope = _RunScope(None, invocations) bound = parameters.bind(instance, *args, **kwargs) - config = dict(bound.arguments.get("config") or {}) - run_id = config.get("run_id") or uuid4() - config["run_id"] = run_id - bound.arguments["config"] = config - stream = wrapped(*bound.args[1:], **bound.kwargs) - scope = _RunScope(run_id, invocations) - cls = _AsyncContextStream if asynchronous else _SyncContextStream - return cls(stream, scope) + + 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, @@ -188,7 +294,7 @@ def _wrap_run( asynchronous: bool, ) -> Callable[..., Any]: def scope_for(kwargs: dict[str, Any]) -> _RunScope: - run_id = kwargs.get("run_id") or uuid4() + run_id = kwargs.get("run_id") or _new_run_id() kwargs["run_id"] = run_id return _RunScope(run_id, invocations, on_error) 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.oldest.txt b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.oldest.txt index 9cce2035e..536933f56 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.oldest.txt +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/requirements.oldest.txt @@ -20,6 +20,12 @@ # 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 @@ -27,7 +33,7 @@ langchain-anthropic==0.3.0 boto3==1.37.0 opentelemetry-instrumentation-httpx >= 0.64b0, <1 --e instrumentation/opentelemetry-instrumentation-genai-openai[instruments] +opentelemetry-instrumentation-genai-openai[instruments]==1.2b0 # The oldest OpenAI client still passes the removed httpx proxies argument. httpx==0.27.2 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" }, ] From ba1598bfc9b45af2165f9aa68cfd64bee26eb7f5 Mon Sep 17 00:00:00 2001 From: Sangkyoon Nam Date: Sun, 11 Oct 2026 15:06:11 +0900 Subject: [PATCH 3/4] test(langchain): cover execution-scoped context propagation --- .../tests/test_custom_execution_context.py | 278 +++++++++++++++- .../tests/test_execution_context.py | 315 +++++++++++++++++- .../tests/test_langgraph_interrupt.py | 12 +- .../tests/test_persistence_context.py | 41 ++- .../tests/test_stream_context.py | 260 ++++++++++++++- 5 files changed, 859 insertions(+), 47 deletions(-) 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 index 545910a81..7a17025cc 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_custom_execution_context.py +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_custom_execution_context.py @@ -4,14 +4,24 @@ 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 @@ -20,6 +30,7 @@ _HTTP_SCOPE, _LC_SCOPE, _OPENAI_SCOPE, + _assert_cancellation, _child, _spans, clients, @@ -29,6 +40,13 @@ __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, @@ -186,8 +204,8 @@ 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",) + _assert_cancellation(error, errors[0]) + assert errors[0].args == ("cancelled by test",) raise finally: assert context.get_current() is before @@ -212,3 +230,259 @@ async def request() -> None: == "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 index c517847dd..681f3d777 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_execution_context.py +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_execution_context.py @@ -5,9 +5,14 @@ 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 @@ -18,14 +23,18 @@ 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, RunnableLambda +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 +from opentelemetry.instrumentation.genai.langchain import ( + LangChainInstrumentor, + agent_context, +) from opentelemetry.instrumentation.genai.langchain.callback_handler import ( OpenTelemetryLangChainCallbackHandler, ) @@ -97,6 +106,65 @@ async def __aiter__(self) -> AsyncIterator[bytes]: 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", @@ -125,7 +193,11 @@ def respond(request: httpx.Request) -> httpx.Response: return httpx.Response( 200, headers={"content-type": "text/event-stream"}, - stream=ResponseStream(), + stream=( + ResponsesApiStream() + if request.url.path.endswith("/responses") + else ResponseStream() + ), ) return httpx.Response( 200, @@ -195,6 +267,22 @@ def _child(child: ReadableSpan, parent: ReadableSpan) -> None: ) +def _assert_cancellation( + error: BaseException, original: BaseException +) -> None: + # langchain-core awaits a tool or stream body in a task of its own. On + # Python 3.10 a task re-raises its cancellation as a new CancelledError + # chained to the original; 3.11 re-raises the original (bpo-45390). + if sys.version_info >= (3, 11): + assert error is original + return + assert isinstance(error, asyncio.CancelledError) + cause: BaseException | None = error + while cause is not None and cause is not original: + cause = cause.__context__ + assert cause is original + + @pytest.mark.asyncio @pytest.mark.parametrize("kind", ["structured", "simple"]) @pytest.mark.parametrize("mode", ["sync", "async", "executor"]) @@ -270,10 +358,11 @@ async def test_chat_model_correlates_sdk_and_http( assert result.content == "answer" assert context.get_current() is before (model_span,) = _spans(span_exporter, _LC_SCOPE) - (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + # The OpenAI inference runs inside the LangChain chat span, so util-genai + # suppresses it and the request parents to the chat span. + assert _spans(span_exporter, _OPENAI_SCOPE) == [] (http,) = _spans(span_exporter, _HTTP_SCOPE) - _child(inference, model_span) - _child(http, inference) + _child(http, model_span) @pytest.mark.asyncio @@ -341,6 +430,13 @@ 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( @@ -349,8 +445,13 @@ async def test_graph_node_and_nested_tool_correlate_http( asynchronous: bool, ) -> None: graph_module = pytest.importorskip("langgraph.graph") - helper = pytest.importorskip("langgraph._internal._runnable") - if not hasattr(helper, "set_config_context"): + 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( @@ -459,7 +560,10 @@ async def afail(text: str) -> str: tool.invoke("hello") else: await tool.ainvoke("hello") - assert raised.value is error + if "cancel" in mode: + _assert_cancellation(raised.value, error) + else: + assert raised.value is error assert context.get_current() is before (http,) = _spans(span_exporter, _HTTP_SCOPE) (tool_span,) = _spans(span_exporter, _LC_SCOPE) @@ -521,8 +625,8 @@ async def request() -> None: "hello", config={"callbacks": [OtherHandler()]} ) except asyncio.CancelledError as error: - assert error is errors[0] - assert error.args == ("cancelled by test",) + _assert_cancellation(error, errors[0]) + assert errors[0].args == ("cancelled by test",) raise finally: assert context.get_current() is before @@ -624,8 +728,17 @@ def test_uninstrument_restores_execution_methods( (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: @@ -646,3 +759,183 @@ def test_uninstrument_restores_execution_methods( 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_langgraph_interrupt.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_langgraph_interrupt.py index b3cb7feed..dec44d7bb 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_langgraph_interrupt.py +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_langgraph_interrupt.py @@ -42,10 +42,12 @@ def test_invoke_interrupt_and_resume_do_not_log_callback_warnings( config: Any = {"configurable": {"thread_id": "sync"}} with assert_no_warnings(caplog, _CALLBACK_LOGGER): - interrupted = graph.invoke({"answer": ""}, config) + graph.invoke({"answer": ""}, config) + # LangGraph < 0.4 leaves __interrupt__ out of the returned state, + # so check the pending node instead. + assert graph.get_state(config).next == ("ask",) resumed = graph.invoke(Command(resume="yes"), config) - assert "__interrupt__" in interrupted assert resumed == {"answer": "yes"} @@ -61,8 +63,10 @@ async def test_ainvoke_interrupt_and_resume_do_not_log_callback_warnings( config: Any = {"configurable": {"thread_id": "async"}} with assert_no_warnings(caplog, _CALLBACK_LOGGER): - interrupted = await graph.ainvoke({"answer": ""}, config) + await graph.ainvoke({"answer": ""}, config) + # LangGraph < 0.4 leaves __interrupt__ out of the returned state, + # so check the pending node instead. + assert (await graph.aget_state(config)).next == ("ask",) resumed = await graph.ainvoke(Command(resume="yes"), config) - assert "__interrupt__" in interrupted assert resumed == {"answer": "yes"} diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_persistence_context.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_persistence_context.py index 407acdd74..c894147f6 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_persistence_context.py +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_persistence_context.py @@ -4,6 +4,7 @@ from __future__ import annotations import asyncio +from inspect import signature from typing import Any, TypedDict import pytest @@ -12,6 +13,7 @@ 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 @@ -29,6 +31,12 @@ __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 @@ -205,38 +213,43 @@ def _assert_persistence_parents( @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", [False, True]) @pytest.mark.parametrize("streaming", [False, True]) -@pytest.mark.parametrize("durability", ["sync", "async", "exit"]) +@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, + 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, durability=durability - ): + async for _ in graph.astream({"text": text}, config, **kwargs): assert context.get_current() is before else: - for _ in graph.stream( - {"text": text}, config, durability=durability - ): + for _ in graph.stream({"text": text}, config, **kwargs): assert context.get_current() is before elif asynchronous: - assert await graph.ainvoke( - {"text": text}, config, durability=durability - ) == {"text": text} + assert await graph.ainvoke({"text": text}, config, **kwargs) == { + "text": text + } else: - assert graph.invoke( - {"text": text}, config, durability=durability - ) == {"text": text} + assert graph.invoke({"text": text}, config, **kwargs) == { + "text": text + } assert context.get_current() is before assert ( InMemorySaver.get_tuple(saver, config).checkpoint[ diff --git a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_stream_context.py b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_stream_context.py index 0930fd101..d65fad9a9 100644 --- a/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_stream_context.py +++ b/instrumentation/opentelemetry-instrumentation-genai-langchain/tests/test_stream_context.py @@ -4,17 +4,24 @@ from __future__ import annotations import asyncio +import sys from collections.abc import AsyncIterator, Iterator from types import AsyncGeneratorType, GeneratorType from typing import Any -from uuid import uuid4 +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 +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, ) @@ -25,6 +32,7 @@ _HTTP_SCOPE, _LC_SCOPE, _OPENAI_SCOPE, + _assert_cancellation, _child, _spans, clients, @@ -34,12 +42,16 @@ __all__ = ["clients", "no_detach_errors"] -def _model(clients: Any) -> ChatOpenAI: +_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, ) @@ -90,11 +102,103 @@ async def test_chat_stream_is_lazy_and_restores_consumer_context( assert context.get_current() is before assert config == original_config (model_span,) = _spans(span_exporter, _LC_SCOPE) - (inference,) = _spans(span_exporter, _OPENAI_SCOPE) + assert _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) + _child(http, model_span) + 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) + assert _spans(span_exporter, _OPENAI_SCOPE) == [] + http, consumer_http = _spans(span_exporter, _HTTP_SCOPE) + assert model_span.parent == root.get_span_context() + _child(http, model_span) assert consumer_http.parent == root.get_span_context() _assert_no_runs() @@ -184,15 +288,17 @@ async def test_interleaved_chat_streams_do_not_share_parents( 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: + assert _spans(span_exporter, _OPENAI_SCOPE) == [] + requests = _spans(span_exporter, _HTTP_SCOPE) + assert len(requests) == 2 + assert {span.parent.span_id for span in requests} == { + span.context.span_id for span in model_spans + } + for request in requests: _child( - inference, + request, next( - span - for span in model_spans - if span.context == inference.parent + span for span in model_spans if span.context == request.parent ), ) _assert_no_runs() @@ -283,7 +389,7 @@ async def consume() -> None: try: await anext(stream) except asyncio.CancelledError as error: - assert error is errors[0] + _assert_cancellation(error, errors[0]) raise finally: assert context.get_current() is before @@ -335,8 +441,40 @@ async def consume() -> None: 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 _spans(span_exporter, _OPENAI_SCOPE) == [] + 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) + assert _spans(span_exporter, _OPENAI_SCOPE) == [] assert ( model_span.attributes[error_attributes.ERROR_TYPE] == "asyncio.exceptions.CancelledError" @@ -344,6 +482,63 @@ async def consume() -> None: _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() + kwargs: dict[str, Any] = {} + if sys.version_info < (3, 12) and ( + "stream_chunk_timeout" in ChatOpenAI.model_fields + ): + # Before 3.12, asyncio.wait_for reads each chunk in a task of its + # own and re-raises the interrupt when asyncio.run cancels it, so + # the loop is left again before this caller's finally runs. + kwargs["stream_chunk_timeout"] = None + model = _model(clients, streaming=True, **kwargs) + 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) + assert _spans(span_exporter, _OPENAI_SCOPE) == [] + 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) + assert _spans(span_exporter, _OPENAI_SCOPE) == [] + 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"]) @@ -388,3 +583,36 @@ async def atransform(inputs: AsyncIterator[str]) -> AsyncIterator[str]: (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() From 608bc6cc940f4cdec195c382c00135cb3767b47d Mon Sep 17 00:00:00 2001 From: Sangkyoon Nam Date: Sun, 11 Oct 2026 15:06:11 +0900 Subject: [PATCH 4/4] chore(langchain): add changelog fragments for #818 --- .../.changelog/818.changed | 1 + .../.changelog/818.fixed | 1 + 2 files changed, 2 insertions(+) create mode 100644 instrumentation/opentelemetry-instrumentation-genai-langchain/.changelog/818.changed create mode 100644 instrumentation/opentelemetry-instrumentation-genai-langchain/.changelog/818.fixed 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.