Add support for single key state, eg just list of messages

This commit is contained in:
Nuno Campos
2024-01-20 15:04:35 -08:00
parent 6236fb086d
commit ee66e11eae
4 changed files with 285 additions and 40 deletions
+119 -4
View File
@@ -1,9 +1,10 @@
import json
import operator
import time
import warnings
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from typing import Annotated, Generator, Optional, TypedDict, Union
from typing import Annotated, Generator, Optional, Self, TypedDict, Union
import pytest
from langchain_core.runnables import RunnablePassthrough
@@ -17,6 +18,7 @@ from langgraph.channels.topic import Topic
from langgraph.checkpoint.memory import MemorySaver
from langgraph.graph import END, Graph
from langgraph.graph.state import StateGraph
from langgraph.prebuilt.chat_agent_executor import create_function_calling_executor
from langgraph.pregel import Channel, GraphRecursionError, Pregel
from langgraph.pregel.reserved import ReservedChannels
@@ -788,8 +790,6 @@ def test_conditional_graph() -> None:
def test_conditional_graph_state() -> None:
from copy import deepcopy
from langchain.llms.fake import FakeStreamingListLLM
from langchain_community.tools import tool
from langchain_core.agents import AgentAction, AgentFinish
@@ -894,7 +894,7 @@ def test_conditional_graph_state() -> None:
),
}
assert [deepcopy(c) for c in app.stream({"input": "what is weather in sf"})] == [
assert [*app.stream({"input": "what is weather in sf"})] == [
{
"agent": {
"agent_outcome": AgentAction(
@@ -973,3 +973,118 @@ def test_conditional_graph_state() -> None:
}
},
]
def test_prebuilt_chat() -> None:
from langchain.chat_models.fake import FakeMessagesListChatModel
from langchain_community.tools import tool
from langchain_core.messages import AIMessage, FunctionMessage, HumanMessage
class FakeFuntionChatModel(FakeMessagesListChatModel):
def bind_functions(self, functions: list) -> Self:
return self
@tool()
def search_api(query: str) -> str:
"""Searches the API for the query."""
return f"result for {query}"
tools = [search_api]
app = create_function_calling_executor(
FakeFuntionChatModel(
responses=[
AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": json.dumps("query"),
}
},
),
AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": json.dumps("another"),
}
},
),
AIMessage(content="answer"),
]
),
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.stream([HumanMessage(content="what is weather in sf")])] == [
{
"agent": [
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": [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"),
]
},
]
+121 -5
View File
@@ -1,4 +1,5 @@
import asyncio
import json
import operator
from contextlib import asynccontextmanager, contextmanager
from typing import (
@@ -8,6 +9,7 @@ from typing import (
AsyncIterator,
Generator,
Optional,
Self,
TypedDict,
Union,
)
@@ -23,6 +25,7 @@ from langgraph.channels.last_value import LastValue
from langgraph.channels.topic import Topic
from langgraph.checkpoint.memory import MemorySaver
from langgraph.graph import END, Graph, StateGraph
from langgraph.prebuilt.chat_agent_executor import create_function_calling_executor
from langgraph.pregel import Channel, GraphRecursionError, Pregel
from langgraph.pregel.reserved import ReservedChannels
@@ -834,8 +837,6 @@ async def test_conditional_graph() -> None:
async def test_conditional_graph_state() -> None:
from copy import deepcopy
from langchain.llms.fake import FakeStreamingListLLM
from langchain_community.tools import tool
from langchain_core.agents import AgentAction, AgentFinish
@@ -940,9 +941,7 @@ async def test_conditional_graph_state() -> None:
),
}
assert [
deepcopy(c) async for c in app.astream({"input": "what is weather in sf"})
] == [
assert [c async for c in app.astream({"input": "what is weather in sf"})] == [
{
"agent": {
"agent_outcome": AgentAction(
@@ -1021,3 +1020,120 @@ async def test_conditional_graph_state() -> None:
}
},
]
async def test_prebuilt_chat() -> None:
from langchain.chat_models.fake import FakeMessagesListChatModel
from langchain_community.tools import tool
from langchain_core.messages import AIMessage, FunctionMessage, HumanMessage
class FakeFuntionChatModel(FakeMessagesListChatModel):
def bind_functions(self, functions: list) -> Self:
return self
@tool()
def search_api(query: str) -> str:
"""Searches the API for the query."""
return f"result for {query}"
tools = [search_api]
app = create_function_calling_executor(
FakeFuntionChatModel(
responses=[
AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": json.dumps("query"),
}
},
),
AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": json.dumps("another"),
}
},
),
AIMessage(content="answer"),
]
),
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 [
c async for c in app.astream([HumanMessage(content="what is weather in sf")])
] == [
{
"agent": [
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": [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"),
]
},
]