From 4849501ccfc0a21703d93d2d1ee71220c72eb742 Mon Sep 17 00:00:00 2001 From: William FH <13333726+hinthornw@users.noreply.github.com> Date: Tue, 9 Jul 2024 20:26:45 -0700 Subject: [PATCH] Copy core context in RunnableCallable (#973) --- libs/langgraph/langgraph/utils.py | 12 ++++++-- libs/langgraph/tests/test_utils.py | 49 ++++++++++++++++++++++++++++++ 2 files changed, 59 insertions(+), 2 deletions(-) diff --git a/libs/langgraph/langgraph/utils.py b/libs/langgraph/langgraph/utils.py index 7d94b81e9..621690e2d 100644 --- a/libs/langgraph/langgraph/utils.py +++ b/libs/langgraph/langgraph/utils.py @@ -22,6 +22,14 @@ from langchain_core.runnables.graph import Edge, Graph, Node, is_uuid from langchain_core.runnables.utils import accepts_config from typing_extensions import TypeGuard +try: + from langchain_core.runnables.config import _set_config_context +except ImportError: + # For forwards compatibility + def _set_config_context(context: RunnableConfig) -> None: # type: ignore + """Set the context for the current thread.""" + var_child_runnable_config.set(context) + # Before Python 3.11 native StrEnum is not available class StrEnum(str, enum.Enum): @@ -89,7 +97,7 @@ class RunnableCallable(Runnable): else: config = merge_configs(self.config, config) context = copy_context() - context.run(var_child_runnable_config.set, config) + context.run(_set_config_context, config) if accepts_config(self.func): kwargs["config"] = config ret = context.run(self.func, input, **kwargs) @@ -110,7 +118,7 @@ class RunnableCallable(Runnable): else: config = merge_configs(self.config, config) context = copy_context() - context.run(var_child_runnable_config.set, config) + context.run(_set_config_context, config) if accepts_config(self.afunc): kwargs["config"] = config if sys.version_info >= (3, 11): diff --git a/libs/langgraph/tests/test_utils.py b/libs/langgraph/tests/test_utils.py index ca16e4c4c..1e3bf8ea1 100644 --- a/libs/langgraph/tests/test_utils.py +++ b/libs/langgraph/tests/test_utils.py @@ -1,5 +1,14 @@ import functools +import sys +import uuid +from typing import TypedDict +from unittest.mock import patch +import langsmith +import pytest + +from langgraph.graph import END, StateGraph +from langgraph.graph.graph import CompiledGraph from langgraph.utils import is_async_callable, is_async_generator @@ -70,3 +79,43 @@ def test_is_generator() -> None: assert not is_async_generator(sync_runnable) wrapped_sync_runnable = functools.wraps(sync_runnable)(sync_runnable) assert not is_async_generator(wrapped_sync_runnable) + + +@pytest.fixture +def rt_graph() -> CompiledGraph: + class State(TypedDict): + foo: int + node_run_id: int + + def node(_: State): + from langsmith import get_current_run_tree # type: ignore + + return {"node_run_id": get_current_run_tree().id} # type: ignore + + graph = StateGraph(State) + graph.add_node(node) + graph.set_entry_point("node") + graph.add_edge("node", END) + return graph.compile() + + +def test_runnable_callable_tracing_nested(rt_graph: CompiledGraph) -> None: + with patch("langsmith.client.Client", spec=langsmith.Client) as mock_client: + with patch("langchain_core.tracers.langchain.get_client") as mock_get_client: + mock_get_client.return_value = mock_client + with langsmith.tracing_context(enabled=True): + res = rt_graph.invoke({"foo": 1}) + assert isinstance(res["node_run_id"], uuid.UUID) + + +@pytest.mark.skipif( + sys.version_info < (3, 11), + reason="Python 3.11+ is required for async contextvars support", +) +async def test_runnable_callable_tracing_nested_async(rt_graph: CompiledGraph) -> None: + with patch("langsmith.client.Client", spec=langsmith.Client) as mock_client: + with patch("langchain_core.tracers.langchain.get_client") as mock_get_client: + mock_get_client.return_value = mock_client + with langsmith.tracing_context(enabled=True): + res = await rt_graph.ainvoke({"foo": 1}) + assert isinstance(res["node_run_id"], uuid.UUID)