diff --git a/langgraph/prebuilt/chat_agent_executor.py b/langgraph/prebuilt/chat_agent_executor.py index df6588dc1..253473416 100644 --- a/langgraph/prebuilt/chat_agent_executor.py +++ b/langgraph/prebuilt/chat_agent_executor.py @@ -1,6 +1,6 @@ import json import operator -from typing import Annotated +from typing import Annotated, Sequence, TypedDict from langchain.tools.render import format_tool_to_openai_function from langchain_core.agents import AgentAction @@ -23,7 +23,8 @@ def create_function_calling_executor(model, tools): ) # Define the function that determines whether to continue or not - def should_continue(messages): + def should_continue(state): + messages = state["messages"] last_message = messages[-1] # If there is no function call, then we finish if "function_call" not in last_message.additional_kwargs: @@ -33,18 +34,21 @@ def create_function_calling_executor(model, tools): return "continue" # Define the function that calls the model - def call_model(messages): + def call_model(state): + messages = state["messages"] response = model.invoke(messages) # We return a list, because this will get added to the existing list - return [response] + return {"messages": [response]} - async def acall_model(messages): + async def acall_model(state): + messages = state["messages"] response = await model.ainvoke(messages) # We return a list, because this will get added to the existing list - return [response] + return {"messages": [response]} # Define the function to execute tools - def _get_action(messages): + def _get_action(state): + messages = state["messages"] # Based on the continue condition # we know the last message involves a function call last_message = messages[-1] @@ -64,7 +68,7 @@ def create_function_calling_executor(model, tools): # We use the response to create a FunctionMessage function_message = FunctionMessage(content=str(response), name=action.tool) # We return a list, because this will get added to the existing list - return [function_message] + return {"messages": [function_message]} async def acall_tool(state): action = _get_action(state) @@ -73,13 +77,17 @@ def create_function_calling_executor(model, tools): # We use the response to create a FunctionMessage function_message = FunctionMessage(content=str(response), name=action.tool) # We return a list, because this will get added to the existing list - return [function_message] + return {"messages": [function_message]} - # Define a new graph with state + # We create the AgentState that we will pass around # This simply involves a list of messages # We want steps to return messages to append to the list # So we annotate the messages attribute with operator.add - workflow = StateGraph(Annotated[list[BaseMessage], operator.add]) + class AgentState(TypedDict): + messages: Annotated[Sequence[BaseMessage], operator.add] + + # Define a new graph + workflow = StateGraph(AgentState) # Define the two nodes we will cycle between workflow.add_node("agent", RunnableLambda(call_model, acall_model)) diff --git a/tests/test_pregel.py b/tests/test_pregel.py index 30842bd44..3557a526d 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -1018,73 +1018,103 @@ def test_prebuilt_chat() -> None: tools, ) - assert app.invoke([HumanMessage(content="what is weather in sf")]) == [ - HumanMessage(content="what is weather in sf"), - AIMessage( - content="", - additional_kwargs={ - "function_call": {"name": "search_api", "arguments": '"query"'} - }, - ), - FunctionMessage(content="result for query", name="search_api"), - AIMessage( - content="", - additional_kwargs={ - "function_call": {"name": "search_api", "arguments": '"another"'} - }, - ), - FunctionMessage(content="result for another", name="search_api"), - AIMessage(content="answer"), - ] + assert app.invoke( + {"messages": [HumanMessage(content="what is weather in sf")]} + ) == { + "messages": [ + HumanMessage(content="what is weather in sf"), + AIMessage( + content="", + additional_kwargs={ + "function_call": {"name": "search_api", "arguments": '"query"'} + }, + ), + FunctionMessage(content="result for query", name="search_api"), + AIMessage( + content="", + additional_kwargs={ + "function_call": {"name": "search_api", "arguments": '"another"'} + }, + ), + FunctionMessage(content="result for another", name="search_api"), + AIMessage(content="answer"), + ] + } - assert [*app.stream([HumanMessage(content="what is weather in sf")])] == [ + assert [ + *app.stream({"messages": [HumanMessage(content="what is weather in sf")]}) + ] == [ { - "agent": [ - AIMessage( - content="", - additional_kwargs={ - "function_call": {"name": "search_api", "arguments": '"query"'} - }, - ) - ] + "agent": { + "messages": [ + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"query"', + } + }, + ) + ] + } }, - {"action": [FunctionMessage(content="result for query", name="search_api")]}, { - "agent": [ - AIMessage( - content="", - additional_kwargs={ - "function_call": { - "name": "search_api", - "arguments": '"another"', - } - }, - ) - ] + "action": { + "messages": [ + FunctionMessage(content="result for query", name="search_api") + ] + } }, - {"action": [FunctionMessage(content="result for another", name="search_api")]}, - {"agent": [AIMessage(content="answer")]}, { - "__end__": [ - HumanMessage(content="what is weather in sf"), - AIMessage( - content="", - additional_kwargs={ - "function_call": {"name": "search_api", "arguments": '"query"'} - }, - ), - FunctionMessage(content="result for query", name="search_api"), - AIMessage( - content="", - additional_kwargs={ - "function_call": { - "name": "search_api", - "arguments": '"another"', - } - }, - ), - FunctionMessage(content="result for another", name="search_api"), - AIMessage(content="answer"), - ] + "agent": { + "messages": [ + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"another"', + } + }, + ) + ] + } + }, + { + "action": { + "messages": [ + FunctionMessage(content="result for another", name="search_api") + ] + } + }, + {"agent": {"messages": [AIMessage(content="answer")]}}, + { + "__end__": { + "messages": [ + HumanMessage(content="what is weather in sf"), + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"query"', + } + }, + ), + FunctionMessage(content="result for query", name="search_api"), + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"another"', + } + }, + ), + FunctionMessage(content="result for another", name="search_api"), + AIMessage(content="answer"), + ] + } }, ] diff --git a/tests/test_pregel_async.py b/tests/test_pregel_async.py index 40d3739e1..3adc9e5a0 100644 --- a/tests/test_pregel_async.py +++ b/tests/test_pregel_async.py @@ -1065,75 +1065,106 @@ async def test_prebuilt_chat() -> None: tools, ) - assert await app.ainvoke([HumanMessage(content="what is weather in sf")]) == [ - HumanMessage(content="what is weather in sf"), - AIMessage( - content="", - additional_kwargs={ - "function_call": {"name": "search_api", "arguments": '"query"'} - }, - ), - FunctionMessage(content="result for query", name="search_api"), - AIMessage( - content="", - additional_kwargs={ - "function_call": {"name": "search_api", "arguments": '"another"'} - }, - ), - FunctionMessage(content="result for another", name="search_api"), - AIMessage(content="answer"), - ] + assert await app.ainvoke( + {"messages": [HumanMessage(content="what is weather in sf")]} + ) == { + "messages": [ + HumanMessage(content="what is weather in sf"), + AIMessage( + content="", + additional_kwargs={ + "function_call": {"name": "search_api", "arguments": '"query"'} + }, + ), + FunctionMessage(content="result for query", name="search_api"), + AIMessage( + content="", + additional_kwargs={ + "function_call": {"name": "search_api", "arguments": '"another"'} + }, + ), + FunctionMessage(content="result for another", name="search_api"), + AIMessage(content="answer"), + ] + } assert [ - c async for c in app.astream([HumanMessage(content="what is weather in sf")]) + c + async for c in app.astream( + {"messages": [HumanMessage(content="what is weather in sf")]} + ) ] == [ { - "agent": [ - AIMessage( - content="", - additional_kwargs={ - "function_call": {"name": "search_api", "arguments": '"query"'} - }, - ) - ] + "agent": { + "messages": [ + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"query"', + } + }, + ) + ] + } }, - {"action": [FunctionMessage(content="result for query", name="search_api")]}, { - "agent": [ - AIMessage( - content="", - additional_kwargs={ - "function_call": { - "name": "search_api", - "arguments": '"another"', - } - }, - ) - ] + "action": { + "messages": [ + FunctionMessage(content="result for query", name="search_api") + ] + } }, - {"action": [FunctionMessage(content="result for another", name="search_api")]}, - {"agent": [AIMessage(content="answer")]}, { - "__end__": [ - HumanMessage(content="what is weather in sf"), - AIMessage( - content="", - additional_kwargs={ - "function_call": {"name": "search_api", "arguments": '"query"'} - }, - ), - FunctionMessage(content="result for query", name="search_api"), - AIMessage( - content="", - additional_kwargs={ - "function_call": { - "name": "search_api", - "arguments": '"another"', - } - }, - ), - FunctionMessage(content="result for another", name="search_api"), - AIMessage(content="answer"), - ] + "agent": { + "messages": [ + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"another"', + } + }, + ) + ] + } + }, + { + "action": { + "messages": [ + FunctionMessage(content="result for another", name="search_api") + ] + } + }, + {"agent": {"messages": [AIMessage(content="answer")]}}, + { + "__end__": { + "messages": [ + HumanMessage(content="what is weather in sf"), + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"query"', + } + }, + ), + FunctionMessage(content="result for query", name="search_api"), + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"another"', + } + }, + ), + FunctionMessage(content="result for another", name="search_api"), + AIMessage(content="answer"), + ] + } }, ]