Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from __future__ import annotations

import asyncio
import inspect
import json
from collections.abc import AsyncGenerator
from typing import Any
Expand Down Expand Up @@ -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,
Expand Down
146 changes: 146 additions & 0 deletions packages/langchain-agents/tests/test_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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()
Expand Down
Loading