mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-26 09:32:25 +02:00
Merge pull request #2342 from langchain-ai/nc/5nov/send-test-interrupt-before
Add two more test cases for Send + interrupt
This commit is contained in:
@@ -2097,6 +2097,205 @@ def test_send_react_interrupt(
|
||||
}
|
||||
assert foo_called == 0
|
||||
|
||||
# interrupt-update-resume flow, creating new Send in update call
|
||||
foo_called = 0
|
||||
graph = builder.compile(checkpointer=checkpointer, interrupt_before=["foo"])
|
||||
thread1 = {"configurable": {"thread_id": "3"}}
|
||||
assert graph.invoke({"messages": [HumanMessage("hello")]}, thread1) == {
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="hello"),
|
||||
_AnyIdAIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [1, 2, 3]},
|
||||
"id": "",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
),
|
||||
]
|
||||
}
|
||||
assert foo_called == 0
|
||||
|
||||
# get state should show the pending task
|
||||
state = graph.get_state(thread1)
|
||||
assert state == StateSnapshot(
|
||||
values={
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="hello"),
|
||||
_AnyIdAIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [1, 2, 3]},
|
||||
"id": "",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
),
|
||||
]
|
||||
},
|
||||
next=("foo",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "3",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={
|
||||
"step": 1,
|
||||
"source": "loop",
|
||||
"writes": {
|
||||
"agent": {
|
||||
"messages": _AnyIdAIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [1, 2, 3]},
|
||||
"id": "",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
)
|
||||
}
|
||||
},
|
||||
"parents": {},
|
||||
"thread_id": "3",
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
"configurable": {
|
||||
"thread_id": "3",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
tasks=(
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="foo",
|
||||
path=("__pregel_push", 0),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
result=None,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
# replace the tool call, should clear previous send, create new one
|
||||
graph.update_state(
|
||||
thread1,
|
||||
{
|
||||
"messages": AIMessage(
|
||||
"",
|
||||
id=ai_message.id,
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [4, 5, 6]},
|
||||
"id": "tool1",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
# prev tool call no longer in pending tasks, new tool call is
|
||||
assert graph.get_state(thread1) == StateSnapshot(
|
||||
values={
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="hello"),
|
||||
_AnyIdAIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [4, 5, 6]},
|
||||
"id": "tool1",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
),
|
||||
]
|
||||
},
|
||||
next=("foo",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "3",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={
|
||||
"step": 2,
|
||||
"source": "update",
|
||||
"writes": {
|
||||
"agent": {
|
||||
"messages": _AnyIdAIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [4, 5, 6]},
|
||||
"id": "tool1",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
)
|
||||
}
|
||||
},
|
||||
"parents": {},
|
||||
"thread_id": "3",
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
"configurable": {
|
||||
"thread_id": "3",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
tasks=(
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="foo",
|
||||
path=("__pregel_push", 0),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
result=None,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
# prev tool call not executed, new tool call is
|
||||
assert graph.invoke(None, thread1) == {
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="hello"),
|
||||
AIMessage(
|
||||
"",
|
||||
id="ai1",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [4, 5, 6]},
|
||||
"id": "tool1",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
),
|
||||
_AnyIdToolMessage(content="{'hi': [4, 5, 6]}", tool_call_id="tool1"),
|
||||
]
|
||||
}
|
||||
assert foo_called == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
|
||||
def test_send_react_interrupt_control(
|
||||
@@ -5275,6 +5474,8 @@ def test_state_graph_packets(
|
||||
{"agent": {"messages": AIMessage(content="answer", id="ai3")}},
|
||||
]
|
||||
|
||||
# interrupt after agent
|
||||
|
||||
app_w_interrupt = workflow.compile(
|
||||
checkpointer=checkpointer,
|
||||
interrupt_after=["agent"],
|
||||
@@ -5562,6 +5763,296 @@ def test_state_graph_packets(
|
||||
parent_config=[*app_w_interrupt.checkpointer.list(config, limit=2)][-1].config,
|
||||
)
|
||||
|
||||
# interrupt before tools
|
||||
|
||||
app_w_interrupt = workflow.compile(
|
||||
checkpointer=checkpointer,
|
||||
interrupt_before=["tools"],
|
||||
)
|
||||
config = {"configurable": {"thread_id": "2"}}
|
||||
model.i = 0
|
||||
|
||||
assert [
|
||||
c
|
||||
for c in app_w_interrupt.stream(
|
||||
{"messages": HumanMessage(content="what is weather in sf")}, config
|
||||
)
|
||||
] == [
|
||||
{
|
||||
"agent": {
|
||||
"messages": AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "query"},
|
||||
},
|
||||
],
|
||||
)
|
||||
}
|
||||
},
|
||||
{"__interrupt__": ()},
|
||||
]
|
||||
|
||||
assert app_w_interrupt.get_state(config) == StateSnapshot(
|
||||
values={
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="what is weather in sf"),
|
||||
AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "query"},
|
||||
},
|
||||
],
|
||||
),
|
||||
]
|
||||
},
|
||||
tasks=(PregelTask(AnyStr(), "tools", (PUSH, 0)),),
|
||||
next=("tools",),
|
||||
config=(app_w_interrupt.checkpointer.get_tuple(config)).config,
|
||||
created_at=(app_w_interrupt.checkpointer.get_tuple(config)).checkpoint["ts"],
|
||||
metadata={
|
||||
"parents": {},
|
||||
"source": "loop",
|
||||
"step": 1,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"messages": AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "query"},
|
||||
},
|
||||
],
|
||||
)
|
||||
}
|
||||
},
|
||||
"thread_id": "2",
|
||||
},
|
||||
parent_config=[*app_w_interrupt.checkpointer.list(config, limit=2)][-1].config,
|
||||
)
|
||||
|
||||
# modify ai message
|
||||
last_message = (app_w_interrupt.get_state(config)).values["messages"][-1]
|
||||
last_message.tool_calls[0]["args"]["query"] = "a different query"
|
||||
app_w_interrupt.update_state(
|
||||
config, {"messages": last_message, "something_extra": "hi there"}
|
||||
)
|
||||
|
||||
# message was replaced instead of appended
|
||||
assert app_w_interrupt.get_state(config) == StateSnapshot(
|
||||
values={
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="what is weather in sf"),
|
||||
AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a different query"},
|
||||
},
|
||||
],
|
||||
),
|
||||
]
|
||||
},
|
||||
tasks=(PregelTask(AnyStr(), "tools", (PUSH, 0)),),
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
created_at=(app_w_interrupt.checkpointer.get_tuple(config)).checkpoint["ts"],
|
||||
metadata={
|
||||
"parents": {},
|
||||
"source": "update",
|
||||
"step": 2,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"messages": AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a different query"},
|
||||
},
|
||||
],
|
||||
),
|
||||
"something_extra": "hi there",
|
||||
}
|
||||
},
|
||||
"thread_id": "2",
|
||||
},
|
||||
parent_config=[*app_w_interrupt.checkpointer.list(config, limit=2)][-1].config,
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||
{
|
||||
"tools": {
|
||||
"messages": _AnyIdToolMessage(
|
||||
content="result for a different query",
|
||||
name="search_api",
|
||||
tool_call_id="tool_call123",
|
||||
)
|
||||
}
|
||||
},
|
||||
{
|
||||
"agent": {
|
||||
"messages": AIMessage(
|
||||
id="ai2",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call234",
|
||||
"name": "search_api",
|
||||
"args": {"query": "another", "idx": 0},
|
||||
},
|
||||
{
|
||||
"id": "tool_call567",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a third one", "idx": 1},
|
||||
},
|
||||
],
|
||||
)
|
||||
},
|
||||
},
|
||||
{"__interrupt__": ()},
|
||||
]
|
||||
|
||||
assert app_w_interrupt.get_state(config) == StateSnapshot(
|
||||
values={
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="what is weather in sf"),
|
||||
AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a different query"},
|
||||
},
|
||||
],
|
||||
),
|
||||
_AnyIdToolMessage(
|
||||
content="result for a different query",
|
||||
name="search_api",
|
||||
tool_call_id="tool_call123",
|
||||
),
|
||||
AIMessage(
|
||||
id="ai2",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call234",
|
||||
"name": "search_api",
|
||||
"args": {"query": "another", "idx": 0},
|
||||
},
|
||||
{
|
||||
"id": "tool_call567",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a third one", "idx": 1},
|
||||
},
|
||||
],
|
||||
),
|
||||
]
|
||||
},
|
||||
tasks=(
|
||||
PregelTask(AnyStr(), "tools", (PUSH, 0)),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, 1)),
|
||||
),
|
||||
next=("tools", "tools"),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
created_at=(app_w_interrupt.checkpointer.get_tuple(config)).checkpoint["ts"],
|
||||
metadata={
|
||||
"parents": {},
|
||||
"source": "loop",
|
||||
"step": 4,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"messages": AIMessage(
|
||||
id="ai2",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call234",
|
||||
"name": "search_api",
|
||||
"args": {"query": "another", "idx": 0},
|
||||
},
|
||||
{
|
||||
"id": "tool_call567",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a third one", "idx": 1},
|
||||
},
|
||||
],
|
||||
)
|
||||
},
|
||||
},
|
||||
"thread_id": "2",
|
||||
},
|
||||
parent_config=[*app_w_interrupt.checkpointer.list(config, limit=2)][-1].config,
|
||||
)
|
||||
|
||||
app_w_interrupt.update_state(
|
||||
config,
|
||||
{
|
||||
"messages": AIMessage(content="answer", id="ai2"),
|
||||
"something_extra": "hi there",
|
||||
},
|
||||
)
|
||||
|
||||
# replaces message even if object identity is different, as long as id is the same
|
||||
assert app_w_interrupt.get_state(config) == StateSnapshot(
|
||||
values={
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="what is weather in sf"),
|
||||
AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a different query"},
|
||||
},
|
||||
],
|
||||
),
|
||||
_AnyIdToolMessage(
|
||||
content="result for a different query",
|
||||
name="search_api",
|
||||
tool_call_id="tool_call123",
|
||||
),
|
||||
AIMessage(content="answer", id="ai2"),
|
||||
]
|
||||
},
|
||||
tasks=(),
|
||||
next=(),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
created_at=(app_w_interrupt.checkpointer.get_tuple(config)).checkpoint["ts"],
|
||||
metadata={
|
||||
"parents": {},
|
||||
"source": "update",
|
||||
"step": 5,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"messages": AIMessage(content="answer", id="ai2"),
|
||||
"something_extra": "hi there",
|
||||
}
|
||||
},
|
||||
"thread_id": "2",
|
||||
},
|
||||
parent_config=[*app_w_interrupt.checkpointer.list(config, limit=2)][-1].config,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
|
||||
def test_message_graph(
|
||||
|
||||
@@ -2311,6 +2311,205 @@ async def test_send_react_interrupt(checkpointer_name: str) -> None:
|
||||
}
|
||||
assert foo_called == 0
|
||||
|
||||
# interrupt-update-resume flow, creating new Send in update call
|
||||
foo_called = 0
|
||||
graph = builder.compile(checkpointer=checkpointer, interrupt_before=["foo"])
|
||||
thread1 = {"configurable": {"thread_id": "3"}}
|
||||
assert await graph.ainvoke({"messages": [HumanMessage("hello")]}, thread1) == {
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="hello"),
|
||||
_AnyIdAIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [1, 2, 3]},
|
||||
"id": "",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
),
|
||||
]
|
||||
}
|
||||
assert foo_called == 0
|
||||
|
||||
# get state should show the pending task
|
||||
state = await graph.aget_state(thread1)
|
||||
assert state == StateSnapshot(
|
||||
values={
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="hello"),
|
||||
_AnyIdAIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [1, 2, 3]},
|
||||
"id": "",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
),
|
||||
]
|
||||
},
|
||||
next=("foo",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "3",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={
|
||||
"step": 1,
|
||||
"source": "loop",
|
||||
"writes": {
|
||||
"agent": {
|
||||
"messages": _AnyIdAIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [1, 2, 3]},
|
||||
"id": "",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
)
|
||||
}
|
||||
},
|
||||
"parents": {},
|
||||
"thread_id": "3",
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
"configurable": {
|
||||
"thread_id": "3",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
tasks=(
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="foo",
|
||||
path=("__pregel_push", 0),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
result=None,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
# replace the tool call, should clear previous send, create new one
|
||||
await graph.aupdate_state(
|
||||
thread1,
|
||||
{
|
||||
"messages": AIMessage(
|
||||
"",
|
||||
id=ai_message.id,
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [4, 5, 6]},
|
||||
"id": "tool1",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
# prev tool call no longer in pending tasks, new tool call is
|
||||
assert await graph.aget_state(thread1) == StateSnapshot(
|
||||
values={
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="hello"),
|
||||
_AnyIdAIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [4, 5, 6]},
|
||||
"id": "tool1",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
),
|
||||
]
|
||||
},
|
||||
next=("foo",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "3",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={
|
||||
"step": 2,
|
||||
"source": "update",
|
||||
"writes": {
|
||||
"agent": {
|
||||
"messages": _AnyIdAIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [4, 5, 6]},
|
||||
"id": "tool1",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
)
|
||||
}
|
||||
},
|
||||
"parents": {},
|
||||
"thread_id": "3",
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
"configurable": {
|
||||
"thread_id": "3",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
tasks=(
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="foo",
|
||||
path=("__pregel_push", 0),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
result=None,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
# prev tool call not executed, new tool call is
|
||||
assert await graph.ainvoke(None, thread1) == {
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="hello"),
|
||||
AIMessage(
|
||||
"",
|
||||
id="ai1",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "foo",
|
||||
"args": {"hi": [4, 5, 6]},
|
||||
"id": "tool1",
|
||||
"type": "tool_call",
|
||||
}
|
||||
],
|
||||
),
|
||||
_AnyIdToolMessage(content="{'hi': [4, 5, 6]}", tool_call_id="tool1"),
|
||||
]
|
||||
}
|
||||
assert foo_called == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
|
||||
async def test_send_react_interrupt_control(checkpointer_name: str) -> None:
|
||||
@@ -5209,6 +5408,8 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
|
||||
]
|
||||
|
||||
async with awith_checkpointer(checkpointer_name) as checkpointer:
|
||||
# interrupt after agent
|
||||
|
||||
app_w_interrupt = workflow.compile(
|
||||
checkpointer=checkpointer,
|
||||
interrupt_after=["agent"],
|
||||
@@ -5502,6 +5703,302 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
|
||||
][-1].config,
|
||||
)
|
||||
|
||||
# interrupt before tools
|
||||
|
||||
app_w_interrupt = workflow.compile(
|
||||
checkpointer=checkpointer,
|
||||
interrupt_before=["tools"],
|
||||
)
|
||||
config = {"configurable": {"thread_id": "2"}}
|
||||
model.i = 0
|
||||
|
||||
assert [
|
||||
c
|
||||
async for c in app_w_interrupt.astream(
|
||||
{"messages": HumanMessage(content="what is weather in sf")}, config
|
||||
)
|
||||
] == [
|
||||
{
|
||||
"agent": {
|
||||
"messages": AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "query"},
|
||||
},
|
||||
],
|
||||
)
|
||||
}
|
||||
},
|
||||
{"__interrupt__": ()},
|
||||
]
|
||||
|
||||
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
|
||||
values={
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="what is weather in sf"),
|
||||
AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "query"},
|
||||
},
|
||||
],
|
||||
),
|
||||
]
|
||||
},
|
||||
tasks=(PregelTask(AnyStr(), "tools", (PUSH, 0)),),
|
||||
next=("tools",),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
created_at=(
|
||||
await app_w_interrupt.checkpointer.aget_tuple(config)
|
||||
).checkpoint["ts"],
|
||||
metadata={
|
||||
"parents": {},
|
||||
"source": "loop",
|
||||
"step": 1,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"messages": AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "query"},
|
||||
},
|
||||
],
|
||||
)
|
||||
}
|
||||
},
|
||||
"thread_id": "2",
|
||||
},
|
||||
parent_config=[
|
||||
c async for c in app_w_interrupt.checkpointer.alist(config, limit=2)
|
||||
][-1].config,
|
||||
)
|
||||
|
||||
# modify ai message
|
||||
last_message = (await app_w_interrupt.aget_state(config)).values["messages"][-1]
|
||||
last_message.tool_calls[0]["args"]["query"] = "a different query"
|
||||
await app_w_interrupt.aupdate_state(config, {"messages": last_message})
|
||||
|
||||
# message was replaced instead of appended
|
||||
tup = await app_w_interrupt.checkpointer.aget_tuple(config)
|
||||
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
|
||||
values={
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="what is weather in sf"),
|
||||
AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a different query"},
|
||||
},
|
||||
],
|
||||
),
|
||||
]
|
||||
},
|
||||
tasks=(PregelTask(AnyStr(), "tools", (PUSH, 0)),),
|
||||
next=("tools",),
|
||||
config=tup.config,
|
||||
created_at=tup.checkpoint["ts"],
|
||||
metadata={
|
||||
"parents": {},
|
||||
"source": "update",
|
||||
"step": 2,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"messages": AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a different query"},
|
||||
},
|
||||
],
|
||||
)
|
||||
}
|
||||
},
|
||||
"thread_id": "2",
|
||||
},
|
||||
parent_config=[
|
||||
c async for c in app_w_interrupt.checkpointer.alist(config, limit=2)
|
||||
][-1].config,
|
||||
)
|
||||
|
||||
assert [c async for c in app_w_interrupt.astream(None, config)] == [
|
||||
{
|
||||
"tools": {
|
||||
"messages": _AnyIdToolMessage(
|
||||
content="result for a different query",
|
||||
name="search_api",
|
||||
tool_call_id="tool_call123",
|
||||
)
|
||||
}
|
||||
},
|
||||
{
|
||||
"agent": {
|
||||
"messages": AIMessage(
|
||||
id="ai2",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call234",
|
||||
"name": "search_api",
|
||||
"args": {"query": "another", "idx": 0},
|
||||
},
|
||||
{
|
||||
"id": "tool_call567",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a third one", "idx": 1},
|
||||
},
|
||||
],
|
||||
)
|
||||
},
|
||||
},
|
||||
{"__interrupt__": ()},
|
||||
]
|
||||
|
||||
tup = await app_w_interrupt.checkpointer.aget_tuple(config)
|
||||
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
|
||||
values={
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="what is weather in sf"),
|
||||
AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a different query"},
|
||||
},
|
||||
],
|
||||
),
|
||||
_AnyIdToolMessage(
|
||||
content="result for a different query",
|
||||
name="search_api",
|
||||
tool_call_id="tool_call123",
|
||||
),
|
||||
AIMessage(
|
||||
id="ai2",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call234",
|
||||
"name": "search_api",
|
||||
"args": {"query": "another", "idx": 0},
|
||||
},
|
||||
{
|
||||
"id": "tool_call567",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a third one", "idx": 1},
|
||||
},
|
||||
],
|
||||
),
|
||||
]
|
||||
},
|
||||
tasks=(
|
||||
PregelTask(AnyStr(), "tools", (PUSH, 0)),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, 1)),
|
||||
),
|
||||
next=("tools", "tools"),
|
||||
config=tup.config,
|
||||
created_at=tup.checkpoint["ts"],
|
||||
metadata={
|
||||
"parents": {},
|
||||
"source": "loop",
|
||||
"step": 4,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"messages": AIMessage(
|
||||
id="ai2",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call234",
|
||||
"name": "search_api",
|
||||
"args": {"query": "another", "idx": 0},
|
||||
},
|
||||
{
|
||||
"id": "tool_call567",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a third one", "idx": 1},
|
||||
},
|
||||
],
|
||||
)
|
||||
},
|
||||
},
|
||||
"thread_id": "2",
|
||||
},
|
||||
parent_config=[
|
||||
c async for c in app_w_interrupt.checkpointer.alist(config, limit=2)
|
||||
][-1].config,
|
||||
)
|
||||
|
||||
await app_w_interrupt.aupdate_state(
|
||||
config,
|
||||
{"messages": AIMessage(content="answer", id="ai2")},
|
||||
)
|
||||
|
||||
# replaces message even if object identity is different, as long as id is the same
|
||||
tup = await app_w_interrupt.checkpointer.aget_tuple(config)
|
||||
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
|
||||
values={
|
||||
"messages": [
|
||||
_AnyIdHumanMessage(content="what is weather in sf"),
|
||||
AIMessage(
|
||||
id="ai1",
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a different query"},
|
||||
},
|
||||
],
|
||||
),
|
||||
_AnyIdToolMessage(
|
||||
content="result for a different query",
|
||||
name="search_api",
|
||||
tool_call_id="tool_call123",
|
||||
),
|
||||
AIMessage(content="answer", id="ai2"),
|
||||
]
|
||||
},
|
||||
tasks=(),
|
||||
next=(),
|
||||
config=tup.config,
|
||||
created_at=tup.checkpoint["ts"],
|
||||
metadata={
|
||||
"parents": {},
|
||||
"source": "update",
|
||||
"step": 5,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"messages": AIMessage(content="answer", id="ai2"),
|
||||
}
|
||||
},
|
||||
"thread_id": "2",
|
||||
},
|
||||
parent_config=[
|
||||
c async for c in app_w_interrupt.checkpointer.alist(config, limit=2)
|
||||
][-1].config,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
|
||||
async def test_message_graph(checkpointer_name: str) -> None:
|
||||
|
||||
Reference in New Issue
Block a user