From 5fb2560078fba8ec84c3b267793dac888d25edc1 Mon Sep 17 00:00:00 2001 From: Andrew Klatzke Date: Fri, 2 Oct 2026 12:59:26 -0800 Subject: [PATCH] fix(langchain-agents): accept sync tool handlers Graph handoff tools stay synchronous so routing can record the selected edge. Awaiting every tool result raised before that record happened. Co-authored-by: Cursor --- .../handler.py | 10 +- .../langchain-agents/tests/test_handler.py | 146 ++++++++++++++++++ 2 files changed, 154 insertions(+), 2 deletions(-) diff --git a/packages/langchain-agents/src/launchdarkly_ai_langchain_agents/handler.py b/packages/langchain-agents/src/launchdarkly_ai_langchain_agents/handler.py index bc72beb3..41c69ed6 100644 --- a/packages/langchain-agents/src/launchdarkly_ai_langchain_agents/handler.py +++ b/packages/langchain-agents/src/launchdarkly_ai_langchain_agents/handler.py @@ -6,6 +6,7 @@ from __future__ import annotations import asyncio +import inspect import json from collections.abc import AsyncGenerator from typing import Any @@ -60,8 +61,13 @@ async def _handler(_name: str = name, **kwargs: Any) -> str: fn = tool_handlers.get(_name) if not fn: raise ValueError(f'No handler registered for tool "{_name}"') - res = await fn(kwargs) - return str(res) + # Handlers may be sync or async. Graph ``__handoff_*`` tools stay sync so + # routing records the selected edge on the call itself; awaiting a plain + # return value raises. Same rule as ``tracking.wrap_tool_handlers``. + result = fn(kwargs) + if inspect.isawaitable(result): + result = await result + return str(result) t = tool_fn( name, diff --git a/packages/langchain-agents/tests/test_handler.py b/packages/langchain-agents/tests/test_handler.py index 1861a1ed..f3cd1854 100644 --- a/packages/langchain-agents/tests/test_handler.py +++ b/packages/langchain-agents/tests/test_handler.py @@ -398,6 +398,24 @@ def _capture_tool(fn: Any, **kw: Any) -> Any: # --------------------------------------------------------------------------- +def _built_tool(config_tools: dict[str, Any], handlers: dict[str, Any]) -> Any: + """The LangChain tool callable ``_build_agent_tools`` registers for the first tool.""" + mocks = _make_langchain_mock() + captured: list[Any] = [] + + def _capture_tool(_name: Any, fn: Any = None, **_kw: Any) -> Any: + captured.append(fn) + return fn + + mocks["langchain_core.tools"].tool = MagicMock(side_effect=_capture_tool) + with _patch_lc(mocks): + from launchdarkly_ai_langchain_agents.handler import _build_agent_tools + + _build_agent_tools(config_tools, handlers) + assert captured, "tool was not registered" + return captured[0] + + class TestToolExecutionLoop: @pytest.mark.asyncio async def test_tool_not_found_throws(self) -> None: @@ -448,6 +466,134 @@ async def _bad_handler(args: Any) -> str: with pytest.raises(RuntimeError, match="tool error"): await captured_fns[0](key="val") + @pytest.mark.asyncio + async def test_sync_handler_returns_string(self) -> None: + def sync_handler(args: dict[str, Any]) -> str: + assert args == {"city": "Paris"} + return "sunny" + + tool = _built_tool( + {"weather": {"description": "d", "parameters": {}}}, + {"weather": sync_handler}, + ) + assert await tool(city="Paris") == "sunny" + + @pytest.mark.asyncio + async def test_async_handler_is_awaited(self) -> None: + seen: dict[str, Any] = {} + + async def async_handler(args: dict[str, Any]) -> str: + seen["args"] = args + return "awaited" + + tool = _built_tool( + {"lookup": {"description": "d", "parameters": {}}}, + {"lookup": async_handler}, + ) + assert await tool(q="hi") == "awaited" + assert seen["args"] == {"q": "hi"} + + @pytest.mark.asyncio + async def test_handoff_handler_records_destination(self) -> None: + chosen: list[str] = [] + + def handoff(_args: dict[str, Any]) -> str: + if not chosen: + chosen.append("billing") + return ( + "Handoff to billing recorded. " + "Finish your own work and provide your final response." + ) + + tool = _built_tool( + {"__handoff_billing": {"description": "transfer", "parameters": {}}}, + {"__handoff_billing": handoff}, + ) + assert await tool() == ( + "Handoff to billing recorded. " + "Finish your own work and provide your final response." + ) + assert chosen == ["billing"] + + @pytest.mark.asyncio + @pytest.mark.parametrize("sync", [True, False]) + async def test_handler_exception_propagates(self, sync: bool) -> None: + def sync_handler(_args: dict[str, Any]) -> str: + raise RuntimeError("sync tool error") + + async def async_handler(_args: dict[str, Any]) -> str: + raise RuntimeError("async tool error") + + handler = sync_handler if sync else async_handler + tool = _built_tool( + {"my-tool": {"description": "d", "parameters": {}}}, + {"my-tool": handler}, + ) + match = "sync tool error" if sync else "async tool error" + with pytest.raises(RuntimeError, match=match): + await tool(key="val") + + @pytest.mark.asyncio + async def test_invoke_and_stream_run_sync_handoff(self) -> None: + """Invoke and stream both install tools from ``_build_agent_tools``.""" + recorded: list[str] = [] + + def handoff(_args: dict[str, Any]) -> str: + recorded.append("billing") + return "Handoff to billing recorded." + + config = _make_config( + instructions="route", + tools={"__handoff_billing": {"description": "transfer", "parameters": {}}}, + ) + handlers = {"__handoff_billing": handoff} + + async def _drive(stream: bool) -> None: + mocks = _make_langchain_mock() + built: dict[str, Any] = {} + + def _create(_model: Any, tools: list[Any], **_kwargs: Any) -> Any: + built["tools"] = tools + return mocks["_agent"] + + mocks["langgraph.prebuilt"].create_react_agent = MagicMock( + side_effect=_create + ) + + async def _call_handoff() -> str: + assert built.get("tools"), "route did not receive built tools" + return await built["tools"][0]() + + async def _ainvoke(*_a: Any, **_kw: Any) -> dict[str, Any]: + built["text"] = await _call_handoff() + return {"messages": [mocks["_ai_msg"]]} + + async def _astream(*_a: Any, **_kw: Any) -> AsyncIterator[Any]: + built["text"] = await _call_handoff() + yield {"agent": {"messages": [mocks["_ai_msg"]]}} + + mocks["_agent"].ainvoke = _ainvoke + mocks["_agent"].astream = _astream + + with _patch_lc(mocks), patch.object(spans_mod, "_HAS_OTEL", False): + h = create_langchain_agents_handler(llm=MagicMock()) + if stream: + async for _event in await h.stream(config, "hi", handlers): + pass + else: + await h(config, "hi", handlers) + + assert built["text"] == "Handoff to billing recorded." + + with patch.object( + handler_mod, "_build_agent_tools", wraps=handler_mod._build_agent_tools + ) as build_tools: + await _drive(False) + await _drive(True) + assert build_tools.call_count == 2 + + assert recorded == ["billing", "billing"] + @pytest.mark.asyncio async def test_no_tools_in_config_handler_never_invoked(self) -> None: mocks = _make_langchain_mock()