Add one more test for send-react-interrupt flow with replacing tool call

This commit is contained in:
Nuno Campos
2024-11-05 09:27:02 -08:00
parent 639501809c
commit f283dac325
2 changed files with 402 additions and 4 deletions
+199
View File
@@ -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(
+203 -4
View File
@@ -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:
@@ -5579,7 +5778,7 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
)
}
},
"thread_id": "2a",
"thread_id": "2",
},
parent_config=[
c async for c in app_w_interrupt.checkpointer.alist(config, limit=2)
@@ -5633,7 +5832,7 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
)
}
},
"thread_id": "2a",
"thread_id": "2",
},
parent_config=[
c async for c in app_w_interrupt.checkpointer.alist(config, limit=2)
@@ -5743,7 +5942,7 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
)
},
},
"thread_id": "2a",
"thread_id": "2",
},
parent_config=[
c async for c in app_w_interrupt.checkpointer.alist(config, limit=2)
@@ -5793,7 +5992,7 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
"messages": AIMessage(content="answer", id="ai2"),
}
},
"thread_id": "2a",
"thread_id": "2",
},
parent_config=[
c async for c in app_w_interrupt.checkpointer.alist(config, limit=2)