Change chat executor back

This commit is contained in:
Nuno Campos
2024-01-20 17:05:47 -08:00
parent ee66e11eae
commit b328a46f73
3 changed files with 204 additions and 135 deletions
+19 -11
View File
@@ -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))
+92 -62
View File
@@ -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"),
]
}
},
]
+93 -62
View File
@@ -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"),
]
}
},
]