diff --git a/langgraph/prebuilt/tool_executor.py b/langgraph/prebuilt/tool_executor.py index 8ce1c65e4..1c44700ee 100644 --- a/langgraph/prebuilt/tool_executor.py +++ b/langgraph/prebuilt/tool_executor.py @@ -1,8 +1,9 @@ -from typing import Any, Sequence, Union +from typing import Any, Callable, Sequence, Union from langchain_core.load.serializable import Serializable from langchain_core.runnables import RunnableConfig from langchain_core.tools import BaseTool +from langchain_core.tools import tool as create_tool from langgraph.utils import RunnableCallable @@ -82,12 +83,15 @@ class ToolExecutor(RunnableCallable): def __init__( self, - tools: Sequence[BaseTool], + tools: Sequence[Union[BaseTool, Callable]], *, invalid_tool_msg_template: str = INVALID_TOOL_MSG_TEMPLATE, ) -> None: super().__init__(self._execute, afunc=self._aexecute, trace=False) - self.tools = tools + tools_ = [ + tool if isinstance(tool, BaseTool) else create_tool(tool) for tool in tools + ] + self.tools = tools_ self.tool_map = {t.name: t for t in tools} self.invalid_tool_msg_template = invalid_tool_msg_template diff --git a/langgraph/prebuilt/tool_node.py b/langgraph/prebuilt/tool_node.py index 314acab21..8b51897f8 100644 --- a/langgraph/prebuilt/tool_node.py +++ b/langgraph/prebuilt/tool_node.py @@ -1,11 +1,12 @@ import asyncio import json -from typing import Any, Literal, Optional, Sequence, Union +from typing import Any, Callable, Dict, Literal, Optional, Sequence, Union from langchain_core.messages import AIMessage, AnyMessage, ToolCall, ToolMessage from langchain_core.runnables import RunnableConfig from langchain_core.runnables.config import get_executor_for_config from langchain_core.tools import BaseTool +from langchain_core.tools import tool as create_tool from langgraph.utils import RunnableCallable @@ -30,13 +31,17 @@ class ToolNode(RunnableCallable): def __init__( self, - tools: Sequence[BaseTool], + tools: Sequence[Union[BaseTool, Callable]], *, name: str = "tools", tags: Optional[list[str]] = None, ) -> None: super().__init__(self._func, self._afunc, name=name, tags=tags, trace=False) - self.tools_by_name = {tool.name: tool for tool in tools} + self.tools_by_name: Dict[str, BaseTool] = {} + for tool_ in tools: + if not isinstance(tool_, BaseTool): + tool_ = create_tool(tool_) + self.tools_by_name[tool_.name] = tool_ def _func( self, input: Union[list[AnyMessage], dict[str, Any]], config: RunnableConfig diff --git a/tests/test_prebuilt.py b/tests/test_prebuilt.py index 97785b3d8..a7d02083d 100644 --- a/tests/test_prebuilt.py +++ b/tests/test_prebuilt.py @@ -7,13 +7,19 @@ from langchain_core.language_models import ( BaseChatModel, LanguageModelInput, ) -from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, SystemMessage +from langchain_core.messages import ( + AIMessage, + BaseMessage, + HumanMessage, + SystemMessage, + ToolMessage, +) from langchain_core.outputs import ChatGeneration, ChatResult from langchain_core.pydantic_v1 import BaseModel from langchain_core.runnables import Runnable, RunnableLambda from langchain_core.tools import BaseTool -from langgraph.prebuilt import create_react_agent +from langgraph.prebuilt import ToolNode, create_react_agent class FakeToolCallingModel(BaseChatModel): @@ -95,3 +101,54 @@ def test_runnable_modifier(): response = agent.invoke({"messages": inputs}) expected_response = {"messages": inputs + [AIMessage(content="Baz", id="0")]} assert response == expected_response + + +async def test_tool_node(): + def tool1(some_val: int, some_other_val: str) -> str: + """Tool 1 docstring.""" + return f"{some_val} - {some_other_val}" + + async def tool2(some_val: int, some_other_val: str) -> str: + """Tool 2 docstring.""" + return f"tool2: {some_val} - {some_other_val}" + + result = ToolNode([tool1]).invoke( + { + "messages": [ + AIMessage( + "hi?", + tool_calls=[ + { + "name": "tool1", + "args": {"some_val": 1, "some_other_val": "foo"}, + "id": "some 0", + } + ], + ) + ] + } + ) + tool_message: ToolMessage = result["messages"][-1] + assert tool_message.type == "tool" + assert tool_message.content == "1 - foo" + assert tool_message.tool_call_id == "some 0" + + result2 = await ToolNode([tool2]).ainvoke( + { + "messages": [ + AIMessage( + "hi?", + tool_calls=[ + { + "name": "tool2", + "args": {"some_val": 2, "some_other_val": "bar"}, + "id": "some 1", + } + ], + ) + ] + } + ) + tool_message: ToolMessage = result2["messages"][-1] + assert tool_message.type == "tool" + assert tool_message.content == "tool2: 2 - bar"