mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-09 11:17:53 +02:00
Accept raw functions (#466)
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
+59
-2
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user