mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 20:45:05 +02:00
Simplify tool executor
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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")},
|
||||||
|
|||||||
Reference in New Issue
Block a user