langgraph[patch]: InjectedState annotation (#1067)

Add annotated for injecting state vars into a Tool
This commit is contained in:
Bagatur
2024-07-19 20:07:21 -07:00
committed by GitHub
parent 75f8a33c9e
commit 610b6cc78c
7 changed files with 277 additions and 90 deletions
@@ -3,7 +3,7 @@ from langgraph.prebuilt import chat_agent_executor
from langgraph.prebuilt.agent_executor import create_agent_executor
from langgraph.prebuilt.chat_agent_executor import create_react_agent
from langgraph.prebuilt.tool_executor import ToolExecutor, ToolInvocation
from langgraph.prebuilt.tool_node import ToolNode, tools_condition
from langgraph.prebuilt.tool_node import InjectedState, ToolNode, tools_condition
from langgraph.prebuilt.tool_validator import ValidationNode
__all__ = [
@@ -15,4 +15,5 @@ __all__ = [
"ToolNode",
"tools_condition",
"ValidationNode",
"InjectedState",
]
+149 -10
View File
@@ -1,11 +1,24 @@
import asyncio
from typing import Any, Callable, Dict, Literal, Optional, Sequence, Tuple, Union, cast
from copy import copy
from typing import (
Any,
Callable,
Dict,
List,
Literal,
Optional,
Sequence,
Tuple,
Union,
cast,
)
from langchain_core.messages import AIMessage, AnyMessage, ToolCall, ToolMessage
from langchain_core.runnables import RunnableConfig
from langchain_core.runnables.config import get_config_list, get_executor_for_config
from langchain_core.tools import BaseTool
from langchain_core.tools import BaseTool, InjectedToolArg
from langchain_core.tools import tool as create_tool
from typing_extensions import get_args
from langgraph.utils import RunnableCallable
@@ -60,18 +73,18 @@ class ToolNode(RunnableCallable):
def _func(
self, input: Union[list[AnyMessage], dict[str, Any]], config: RunnableConfig
) -> Any:
message, output_type = self._parse_input(input)
config_list = get_config_list(config, len(message.tool_calls))
tool_calls, output_type = self._parse_input(input)
config_list = get_config_list(config, len(tool_calls))
with get_executor_for_config(config) as executor:
outputs = [*executor.map(self._run_one, message.tool_calls, config_list)]
outputs = [*executor.map(self._run_one, tool_calls, config_list)]
return outputs if output_type == "list" else {"messages": outputs}
async def _afunc(
self, input: Union[list[AnyMessage], dict[str, Any]], config: RunnableConfig
) -> Any:
message, output_type = self._parse_input(input)
tool_calls, output_type = self._parse_input(input)
outputs = await asyncio.gather(
*(self._arun_one(call, config) for call in message.tool_calls)
*(self._arun_one(call, config) for call in tool_calls)
)
return outputs if output_type == "list" else {"messages": outputs}
@@ -102,7 +115,7 @@ class ToolNode(RunnableCallable):
def _parse_input(
self, input: Union[list[AnyMessage], dict[str, Any]]
) -> Tuple[AIMessage, Literal["list", "dict"]]:
) -> Tuple[List[ToolCall], Literal["list", "dict"]]:
if isinstance(input, list):
output_type = "list"
message: AnyMessage = input[-1]
@@ -114,8 +127,12 @@ class ToolNode(RunnableCallable):
if not isinstance(message, AIMessage):
raise ValueError("Last message is not an AIMessage")
else:
return cast(AIMessage, message), output_type
tool_calls = [
self._inject_state(call, input)
for call in cast(AIMessage, message).tool_calls
]
return tool_calls, output_type
def _validate_tool_call(self, call: ToolCall) -> Optional[ToolMessage]:
if (requested_tool := call["name"]) not in self.tools_by_name:
@@ -127,6 +144,39 @@ class ToolNode(RunnableCallable):
else:
return None
def _inject_state(
self, tool_call: ToolCall, input: Union[list[AnyMessage], dict[str, Any]]
) -> ToolCall:
if tool_call["name"] not in self.tools_by_name:
return tool_call
state_args = _get_state_args(self.tools_by_name[tool_call["name"]])
if state_args and not isinstance(input, dict):
required_fields = list(state_args.values())
if (
len(required_fields) == 1
and required_fields[0] == "messages"
or required_fields[0] is None
):
input = {"messages": input}
else:
err_msg = (
f"Invalid input to ToolNode. Tool {tool_call['name']} requires "
f"graph state dict as input."
)
if any(state_field for state_field in state_args.values()):
required_fields_str = ", ".join(f for f in required_fields if f)
err_msg += f" State should contain fields {required_fields_str}."
raise ValueError(err_msg)
tool_call_copy: ToolCall = copy(tool_call)
tool_call_copy["args"] = {
**tool_call_copy["args"],
**{
tool_arg: cast(dict, input)[state_field] if state_field else input
for tool_arg, state_field in state_args.items()
},
}
return tool_call_copy
def tools_condition(
state: Union[list[AnyMessage], dict[str, Any]],
@@ -183,3 +233,92 @@ def tools_condition(
if hasattr(ai_message, "tool_calls") and len(ai_message.tool_calls) > 0:
return "tools"
return "__end__"
class InjectedState(InjectedToolArg):
"""Annotation for a Tool arg that is meant to be populated with the graph state.
Any Tool argument annotated with InjectedState will be hidden from a tool-calling
model, so that the model doesn't attempt to generate the argument. If using
ToolNode, the appropriate graph state field will be automatically injected into
the model-generated tool args.
Args:
field: The key from state to insert. If None, the entire state is expected to
be passed in.
Example:
```python
from typing import List
from typing_extensions import Annotated, TypedDict
from langchain_core.messages import BaseMessage, AIMessage
from langchain_core.tools import tool
from langgraph.prebuilt import InjectedState, ToolNode
class AgentState(TypedDict):
messages: List[BaseMessage]
foo: str
@tool
def state_tool(x: int, state: Annotated[dict, InjectedState]) -> str:
'''Do something with state.'''
if len(state["messages"]) > 2:
return state["foo"] + str(x)
else:
return "not enough messages"
@tool
def foo_tool(x: int, foo: Annotated[str, InjectedState("foo")]) -> str:
'''Do something else with state.'''
return foo + str(x + 1)
node = ToolNode([state_tool, foo_tool])
tool_call1 = {"name": "state_tool", "args": {"x": 1}, "id": "1", "type": "tool_call"}
tool_call2 = {"name": "foo_tool", "args": {"x": 1}, "id": "2", "type": "tool_call"}
state = {
"messages": [AIMessage("", tool_calls=[tool_call1, tool_call2])],
"foo": "bar",
}
node.invoke(state)
```
```pycon
[
ToolMessage(content='not enough messages', name='state_tool', tool_call_id='1'),
ToolMessage(content='bar2', name='foo_tool', tool_call_id='2')
]
```
""" # noqa: E501
def __init__(self, field: Optional[str] = None) -> None:
self.field = field
def _get_state_args(tool: BaseTool) -> Dict[str, Optional[str]]:
full_schema = tool.get_input_schema()
tool_args_to_state_fields: Dict = {}
for name, type_ in full_schema.__annotations__.items():
injections = [
type_arg
for type_arg in get_args(type_)
if isinstance(type_arg, InjectedState)
or (isinstance(type_arg, type) and issubclass(type_arg, InjectedState))
]
if len(injections) > 1:
raise ValueError(
"A tool argument should not be annotated with InjectedState more than "
f"once. Received arg {name} with annotations {injections}."
)
elif len(injections) == 1:
injection = injections[0]
if isinstance(injection, InjectedState) and injection.field:
tool_args_to_state_fields[name] = injection.field
else:
tool_args_to_state_fields[name] = None
else:
pass
return tool_args_to_state_fields
+1 -1
View File
@@ -4165,4 +4165,4 @@ test = ["big-O", "importlib-resources", "jaraco.functools", "jaraco.itertools",
[metadata]
lock-version = "2.0"
python-versions = ">=3.9.0,<4.0"
content-hash = "5fb6190a1b01d0cd351ea9a0023c8c8d6acf4fe831101ab87f9fc41308f74b74"
content-hash = "0d877d3879473de43aca1e1d36a8f420ff3f4b140807cb5cc24935d9114947be"
+1 -1
View File
@@ -9,7 +9,7 @@ repository = "https://www.github.com/langchain-ai/langgraph"
[tool.poetry.dependencies]
python = ">=3.9.0,<4.0"
langchain-core = ">=0.2.19,<0.3"
langchain-core = ">=0.2.22,<0.3"
[tool.poetry.group.dev.dependencies]
+72 -13
View File
@@ -1,15 +1,11 @@
from typing import Any, Callable, Dict, List, Optional, Sequence, Type, Union
from typing import Annotated, Any, Callable, Dict, List, Optional, Sequence, Type, Union
import pytest
from langchain_core.callbacks import (
CallbackManagerForLLMRun,
)
from langchain_core.language_models import (
BaseChatModel,
LanguageModelInput,
)
from langchain_core.callbacks import CallbackManagerForLLMRun
from langchain_core.language_models import BaseChatModel, LanguageModelInput
from langchain_core.messages import (
AIMessage,
AnyMessage,
BaseMessage,
HumanMessage,
SystemMessage,
@@ -23,11 +19,8 @@ from langchain_core.tools import tool as dec_tool
from pydantic import BaseModel as BaseModelV2
from langgraph.checkpoint.base import BaseCheckpointSaver
from langgraph.prebuilt import (
ToolNode,
ValidationNode,
create_react_agent,
)
from langgraph.prebuilt import ToolNode, ValidationNode, create_react_agent
from langgraph.prebuilt.tool_node import InjectedState
from tests.any_str import AnyStr
from tests.memory_assert import MemorySaverAssertImmutable
@@ -453,3 +446,69 @@ async def test_validation_node(tool_schema: Any, use_message_key: bool):
if use_message_key:
result_sync = result_sync["messages"]
check_results(result_sync)
def test_tool_node_inject_state() -> None:
def tool1(some_val: int, state: Annotated[dict, InjectedState]) -> str:
"""Tool 1 docstring."""
return state["foo"]
def tool2(some_val: int, state: Annotated[dict, InjectedState()]) -> str:
"""Tool 1 docstring."""
return state["foo"]
def tool3(
some_val: int,
foo: Annotated[str, InjectedState("foo")],
msgs: Annotated[List[AnyMessage], InjectedState("messages")],
) -> str:
"""Tool 1 docstring."""
return foo
def tool4(
some_val: int, msgs: Annotated[List[AnyMessage], InjectedState("messages")]
) -> str:
"""Tool 1 docstring."""
return msgs[0].content
node = ToolNode([tool1, tool2, tool3, tool4])
for tool_name in ("tool1", "tool2", "tool3"):
tool_call = {
"name": tool_name,
"args": {"some_val": 1},
"id": "some 0",
"type": "tool_call",
}
msg = AIMessage("hi?", tool_calls=[tool_call])
result = node.invoke({"messages": [msg], "foo": "bar"})
tool_message = result["messages"][-1]
assert tool_message.content == "bar"
if tool_name == "tool3":
with pytest.raises(KeyError):
node.invoke({"messages": [msg], "notfoo": "bar"})
with pytest.raises(ValueError):
node.invoke([msg])
else:
tool_message = node.invoke({"messages": [msg], "notfoo": "bar"})[
"messages"
][-1]
assert "KeyError" in tool_message.content
tool_message = node.invoke([msg])[-1]
assert "KeyError" in tool_message.content
tool_call = {
"name": "tool4",
"args": {"some_val": 1},
"id": "some 0",
"type": "tool_call",
}
msg = AIMessage("hi?", tool_calls=[tool_call])
result = node.invoke({"messages": [msg]})
tool_message = result["messages"][-1]
assert tool_message.content == "hi?"
result = node.invoke([msg])
tool_message = result[-1]
assert tool_message.content == "hi?"