Support updating existing messages in messagegraph

This commit is contained in:
Nuno Campos
2024-03-15 12:13:58 -07:00
parent ba5dbad75f
commit 7548f412fb
4 changed files with 429 additions and 31 deletions
+18 -1
View File
@@ -1,3 +1,4 @@
import uuid
from typing import Annotated, Union
from langchain_core.messages import AnyMessage
@@ -8,11 +9,27 @@ Messages = Union[list[AnyMessage], AnyMessage]
def add_messages(left: Messages, right: Messages) -> Messages:
# coerce to list
if not isinstance(left, list):
left = [left]
if not isinstance(right, list):
right = [right]
return left + right
# assign missing ids
for m in left:
if m.id is None:
m.id = str(uuid.uuid4())
for m in right:
if m.id is None:
m.id = str(uuid.uuid4())
# merge
left_idx_by_id = {m.id: i for i, m in enumerate(left)}
merged = left.copy()
for m in right:
if (existing_idx := left_idx_by_id.get(m.id)) is not None:
merged[existing_idx] = m
else:
merged.append(m)
return merged
class MessageGraph(StateGraph):
+12
View File
@@ -0,0 +1,12 @@
from uuid import UUID
import pytest
from pytest_mock import MockerFixture
@pytest.fixture()
def deterministic_uuids(mocker: MockerFixture) -> MockerFixture:
side_effect = (
UUID(f"00000000-0000-4000-8000-{i:012}", version=4) for i in range(10000)
)
return mocker.patch("uuid.uuid4", side_effect=side_effect)
+201 -15
View File
@@ -2321,7 +2321,9 @@ def test_prebuilt_chat(snapshot: SnapshotAssertion) -> None:
]
def test_message_graph(snapshot: SnapshotAssertion) -> None:
def test_message_graph(
snapshot: SnapshotAssertion, deterministic_uuids: MockerFixture
) -> None:
from langchain.chat_models.fake import FakeMessagesListChatModel
from langchain_community.tools import tool
from langchain_core.agents import AgentAction
@@ -2348,6 +2350,7 @@ def test_message_graph(snapshot: SnapshotAssertion) -> None:
"arguments": json.dumps("query"),
}
},
id="ai1",
),
AIMessage(
content="",
@@ -2357,8 +2360,9 @@ def test_message_graph(snapshot: SnapshotAssertion) -> None:
"arguments": json.dumps("another"),
}
},
id="ai2",
),
AIMessage(content="answer"),
AIMessage(content="answer", id="ai3"),
]
)
@@ -2438,22 +2442,35 @@ def test_message_graph(snapshot: SnapshotAssertion) -> None:
assert app.get_graph().draw_ascii() == snapshot
assert app.invoke(HumanMessage(content="what is weather in sf")) == [
HumanMessage(content="what is weather in sf"),
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000002", # adds missing ids
),
AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
id="ai1", # respects ids passed in
),
FunctionMessage(
content="result for query",
name="search_api",
id="00000000-0000-4000-8000-000000000014",
),
FunctionMessage(content="result for query", name="search_api"),
AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"another"'}
},
id="ai2",
),
FunctionMessage(content="result for another", name="search_api"),
AIMessage(content="answer"),
FunctionMessage(
content="result for another",
name="search_api",
id="00000000-0000-4000-8000-000000000026",
),
AIMessage(content="answer", id="ai3"),
]
assert [*app.stream([HumanMessage(content="what is weather in sf")])] == [
@@ -2463,29 +2480,51 @@ def test_message_graph(snapshot: SnapshotAssertion) -> None:
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
id="ai1",
)
},
{
"action": FunctionMessage(
content="result for query",
name="search_api",
id="00000000-0000-4000-8000-000000000047",
)
},
{"action": FunctionMessage(content="result for query", name="search_api")},
{
"agent": AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"another"'}
},
id="ai2",
)
},
{"action": FunctionMessage(content="result for another", name="search_api")},
{"agent": AIMessage(content="answer")},
{
"action": FunctionMessage(
content="result for another",
name="search_api",
id="00000000-0000-4000-8000-000000000059",
)
},
{"agent": AIMessage(content="answer", id="ai3")},
{
"__end__": [
HumanMessage(content="what is weather in sf"),
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000035",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
id="ai1",
),
FunctionMessage(
content="result for query",
name="search_api",
id="00000000-0000-4000-8000-000000000047",
),
FunctionMessage(content="result for query", name="search_api"),
AIMessage(
content="",
additional_kwargs={
@@ -2494,9 +2533,14 @@ def test_message_graph(snapshot: SnapshotAssertion) -> None:
"arguments": '"another"',
}
},
id="ai2",
),
FunctionMessage(content="result for another", name="search_api"),
AIMessage(content="answer"),
FunctionMessage(
content="result for another",
name="search_api",
id="00000000-0000-4000-8000-000000000059",
),
AIMessage(content="answer", id="ai3"),
]
},
]
@@ -2518,25 +2562,167 @@ def test_message_graph(snapshot: SnapshotAssertion) -> None:
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
id="ai1",
)
}
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
values=[
HumanMessage(content="what is weather in sf"),
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
id="ai1",
),
],
next=("agent:edges",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
)
# TODO use update_state once we have message ids
# modify ai message
last_message = app_w_interrupt.get_state(config).values[-1]
last_message.additional_kwargs["function_call"]["arguments"] = '"a different query"'
app_w_interrupt.update_state(config, last_message)
# message was replaced instead of appended
assert app_w_interrupt.get_state(config) == StateSnapshot(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": '"a different query"',
}
},
id="ai1",
),
],
next=("agent:edges",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
{
"action": FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
)
},
{
"agent": AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"another"'}
},
id="ai2",
)
},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": '"a different query"',
}
},
id="ai1",
),
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"another"'}
},
id="ai2",
),
],
next=("agent:edges",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
)
app_w_interrupt.update_state(
config,
AIMessage(content="answer", id="ai2"),
)
# replaces message even if object identity is different, as long as id is the same
assert app_w_interrupt.get_state(config) == StateSnapshot(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": '"a different query"',
}
},
id="ai1",
),
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
),
AIMessage(content="answer", id="ai2"),
],
next=("agent:edges",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
{
"__end__": [
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": '"a different query"',
}
},
id="ai1",
),
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
),
AIMessage(content="answer", id="ai2"),
]
}
]
def test_in_one_fan_out_out_one_graph_state() -> None:
+198 -15
View File
@@ -2356,7 +2356,7 @@ async def test_prebuilt_chat() -> None:
]
async def test_message_graph() -> None:
async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
from langchain.chat_models.fake import FakeMessagesListChatModel
from langchain_community.tools import tool
from langchain_core.agents import AgentAction
@@ -2383,6 +2383,7 @@ async def test_message_graph() -> None:
"arguments": json.dumps("query"),
}
},
id="ai1",
),
AIMessage(
content="",
@@ -2392,8 +2393,9 @@ async def test_message_graph() -> None:
"arguments": json.dumps("another"),
}
},
id="ai2",
),
AIMessage(content="answer"),
AIMessage(content="answer", id="ai3"),
]
)
@@ -2468,22 +2470,34 @@ async def test_message_graph() -> None:
app = workflow.compile()
assert await app.ainvoke(HumanMessage(content="what is weather in sf")) == [
HumanMessage(content="what is weather in sf"),
HumanMessage(
content="what is weather in sf", id="00000000-0000-4000-8000-000000000002"
),
AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
id="ai1",
),
FunctionMessage(
content="result for query",
name="search_api",
id="00000000-0000-4000-8000-000000000014",
),
FunctionMessage(content="result for query", name="search_api"),
AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"another"'}
},
id="ai2",
),
FunctionMessage(content="result for another", name="search_api"),
AIMessage(content="answer"),
FunctionMessage(
content="result for another",
name="search_api",
id="00000000-0000-4000-8000-000000000026",
),
AIMessage(content="answer", id="ai3"),
]
assert [
@@ -2495,29 +2509,51 @@ async def test_message_graph() -> None:
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
id="ai1",
)
},
{
"action": FunctionMessage(
content="result for query",
name="search_api",
id="00000000-0000-4000-8000-000000000047",
)
},
{"action": FunctionMessage(content="result for query", name="search_api")},
{
"agent": AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"another"'}
},
id="ai2",
)
},
{"action": FunctionMessage(content="result for another", name="search_api")},
{"agent": AIMessage(content="answer")},
{
"action": FunctionMessage(
content="result for another",
name="search_api",
id="00000000-0000-4000-8000-000000000059",
)
},
{"agent": AIMessage(content="answer", id="ai3")},
{
"__end__": [
HumanMessage(content="what is weather in sf"),
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000035",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
id="ai1",
),
FunctionMessage(
content="result for query",
name="search_api",
id="00000000-0000-4000-8000-000000000047",
),
FunctionMessage(content="result for query", name="search_api"),
AIMessage(
content="",
additional_kwargs={
@@ -2526,9 +2562,14 @@ async def test_message_graph() -> None:
"arguments": '"another"',
}
},
id="ai2",
),
FunctionMessage(content="result for another", name="search_api"),
AIMessage(content="answer"),
FunctionMessage(
content="result for another",
name="search_api",
id="00000000-0000-4000-8000-000000000059",
),
AIMessage(content="answer", id="ai3"),
]
},
]
@@ -2550,25 +2591,167 @@ async def test_message_graph() -> None:
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
id="ai1",
)
}
]
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
values=[
HumanMessage(content="what is weather in sf"),
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
id="ai1",
),
],
next=("agent:edges",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
)
# TODO use update_state once we have message ids
# modify ai message
last_message = (await app_w_interrupt.aget_state(config)).values[-1]
last_message.additional_kwargs["function_call"]["arguments"] = '"a different query"'
await app_w_interrupt.aupdate_state(config, last_message)
# message was replaced instead of appended
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": '"a different query"',
}
},
id="ai1",
),
],
next=("agent:edges",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
{
"action": FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
)
},
{
"agent": AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"another"'}
},
id="ai2",
)
},
]
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": '"a different query"',
}
},
id="ai1",
),
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"another"'}
},
id="ai2",
),
],
next=("agent:edges",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
)
await app_w_interrupt.aupdate_state(
config,
AIMessage(content="answer", id="ai2"),
)
# replaces message even if object identity is different, as long as id is the same
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": '"a different query"',
}
},
id="ai1",
),
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
),
AIMessage(content="answer", id="ai2"),
],
next=("agent:edges",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
{
"__end__": [
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
),
AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": '"a different query"',
}
},
id="ai1",
),
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
),
AIMessage(content="answer", id="ai2"),
]
}
]
async def test_in_one_fan_out_out_one_graph_state() -> None: