Accept raw functions (#466)

This commit is contained in:
William FH
2024-05-15 17:03:51 -07:00
committed by GitHub
parent aef06f17ed
commit d8c3e70900
3 changed files with 74 additions and 8 deletions
+7 -3
View File
@@ -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
+8 -3
View File
@@ -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
View File
@@ -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"