Simplify tool executor

This commit is contained in:
Nuno Campos
2024-04-02 14:46:14 -07:00
parent e3a9760bb2
commit 2f8ae4fdb4
2 changed files with 16 additions and 23 deletions
+12 -19
View File
@@ -1,9 +1,11 @@
from typing import Any, Sequence, Union from typing import Any, Sequence, Union
from langchain_core.load.serializable import Serializable from langchain_core.load.serializable import Serializable
from langchain_core.runnables import RunnableBinding, RunnableConfig, RunnableLambda from langchain_core.runnables import RunnableConfig
from langchain_core.tools import BaseTool from langchain_core.tools import BaseTool
from langgraph.utils import RunnableCallable
INVALID_TOOL_MSG_TEMPLATE = ( INVALID_TOOL_MSG_TEMPLATE = (
"{requested_tool_name} is not a valid tool, " "{requested_tool_name} is not a valid tool, "
"try one of [{available_tool_names_str}]." "try one of [{available_tool_names_str}]."
@@ -26,29 +28,20 @@ class ToolInvocation(Serializable):
"""The input to pass in to the Tool.""" """The input to pass in to the Tool."""
class ToolExecutor(RunnableBinding): class ToolExecutor(RunnableCallable):
tools: Sequence[BaseTool]
tool_map: dict
invalid_tool_msg_template: str
def __init__( def __init__(
self, self,
tools: Sequence[BaseTool], tools: Sequence[BaseTool],
*, *,
invalid_tool_msg_template: str = INVALID_TOOL_MSG_TEMPLATE, invalid_tool_msg_template: str = INVALID_TOOL_MSG_TEMPLATE,
**kwargs: Any,
) -> None: ) -> None:
bound = RunnableLambda(self._execute, afunc=self._aexecute) super().__init__(self._execute, afunc=self._aexecute, trace=False)
super().__init__( self.tools = tools
bound=bound, self.tool_map = {t.name: t for t in tools}
tools=tools, self.invalid_tool_msg_template = invalid_tool_msg_template
tool_map={t.name: t for t in tools},
invalid_tool_msg_template=invalid_tool_msg_template,
**kwargs,
)
def _execute( def _execute(
self, tool_invocation: ToolInvocationInterface, *, config: RunnableConfig self, tool_invocation: ToolInvocationInterface, config: RunnableConfig
) -> Any: ) -> Any:
if tool_invocation.tool not in self.tool_map: if tool_invocation.tool not in self.tool_map:
return self.invalid_tool_msg_template.format( return self.invalid_tool_msg_template.format(
@@ -57,11 +50,11 @@ class ToolExecutor(RunnableBinding):
) )
else: else:
tool = self.tool_map[tool_invocation.tool] tool = self.tool_map[tool_invocation.tool]
output = tool.invoke(tool_invocation.tool_input, config=config) output = tool.invoke(tool_invocation.tool_input, config)
return output return output
async def _aexecute( async def _aexecute(
self, tool_invocation: ToolInvocationInterface, *, config: RunnableConfig self, tool_invocation: ToolInvocationInterface, config: RunnableConfig
) -> Any: ) -> Any:
if tool_invocation.tool not in self.tool_map: if tool_invocation.tool not in self.tool_map:
return self.invalid_tool_msg_template.format( return self.invalid_tool_msg_template.format(
@@ -70,5 +63,5 @@ class ToolExecutor(RunnableBinding):
) )
else: else:
tool = self.tool_map[tool_invocation.tool] tool = self.tool_map[tool_invocation.tool]
output = await tool.ainvoke(tool_invocation.tool_input, config=config) output = await tool.ainvoke(tool_invocation.tool_input, config)
return output return output
+4 -4
View File
@@ -2307,7 +2307,7 @@ def test_message_graph(
FunctionMessage( FunctionMessage(
content="result for query", content="result for query",
name="search_api", name="search_api",
id="00000000-0000-4000-8000-000000000014", id="00000000-0000-4000-8000-000000000013",
), ),
AIMessage( AIMessage(
content="", content="",
@@ -2319,7 +2319,7 @@ def test_message_graph(
FunctionMessage( FunctionMessage(
content="result for another", content="result for another",
name="search_api", name="search_api",
id="00000000-0000-4000-8000-000000000026", id="00000000-0000-4000-8000-000000000024",
), ),
AIMessage(content="answer", id="ai3"), AIMessage(content="answer", id="ai3"),
] ]
@@ -2338,7 +2338,7 @@ def test_message_graph(
"action": FunctionMessage( "action": FunctionMessage(
content="result for query", content="result for query",
name="search_api", name="search_api",
id="00000000-0000-4000-8000-000000000046", id="00000000-0000-4000-8000-000000000043",
) )
}, },
{ {
@@ -2354,7 +2354,7 @@ def test_message_graph(
"action": FunctionMessage( "action": FunctionMessage(
content="result for another", content="result for another",
name="search_api", name="search_api",
id="00000000-0000-4000-8000-000000000058", id="00000000-0000-4000-8000-000000000054",
) )
}, },
{"agent": AIMessage(content="answer", id="ai3")}, {"agent": AIMessage(content="answer", id="ai3")},