Files
langgraph/libs/scheduler-kafka/tests/test_push.py
T
2025-01-14 15:49:57 -08:00

207 lines
6.2 KiB
Python

import operator
from typing import (
Annotated,
Literal,
Union,
)
import pytest
from aiokafka import AIOKafkaProducer
from langgraph.checkpoint.base import BaseCheckpointSaver
from langgraph.constants import START
from langgraph.errors import NodeInterrupt
from langgraph.graph.state import CompiledStateGraph, StateGraph
from langgraph.scheduler.kafka import serde
from langgraph.scheduler.kafka.types import MessageToOrchestrator, Topics
from langgraph.types import Command, Send
from tests.any import AnyDict
from tests.drain import drain_topics_async
pytestmark = pytest.mark.anyio
def mk_push_graph(
checkpointer: BaseCheckpointSaver,
) -> CompiledStateGraph:
# copied from test_send_dedupe_on_resume
class InterruptOnce:
ticks: int = 0
def __call__(self, state):
self.ticks += 1
if self.ticks == 1:
raise NodeInterrupt("Bahh")
return ["|".join(("flaky", str(state)))]
class Node:
def __init__(self, name: str):
self.name = name
self.ticks = 0
self.__name__ = name
def __call__(self, state):
self.ticks += 1
update = (
[self.name]
if isinstance(state, list)
else ["|".join((self.name, str(state)))]
)
if isinstance(state, Command):
return state.copy(update=update)
else:
return update
def send_for_fun(state):
return [
Send("2", Command(goto=Send("2", 3))),
Send("2", Command(goto=Send("flaky", 4))),
"3.1",
]
def route_to_three(state) -> Literal["3"]:
return "3"
builder = StateGraph(Annotated[list, operator.add])
builder.add_node(Node("1"))
builder.add_node(Node("2"))
builder.add_node(Node("3"))
builder.add_node(Node("3.1"))
builder.add_node("flaky", InterruptOnce())
builder.add_edge(START, "1")
builder.add_conditional_edges("1", send_for_fun)
builder.add_conditional_edges("2", route_to_three)
return builder.compile(checkpointer=checkpointer)
@pytest.mark.skip("TODO: re-enable in next PR")
async def test_push_graph(topics: Topics, acheckpointer: BaseCheckpointSaver) -> None:
input = ["0"]
config = {"configurable": {"thread_id": "1"}}
graph = mk_push_graph(acheckpointer)
graph_compare = mk_push_graph(acheckpointer)
# start a new run
async with AIOKafkaProducer(value_serializer=serde.dumps) as producer:
await producer.send_and_wait(
topics.orchestrator,
MessageToOrchestrator(input=input, config=config),
)
# drain topics
orch_msgs, exec_msgs = await drain_topics_async(topics, graph)
# check state
state = await graph.aget_state(config)
assert all(not t.error for t in state.tasks)
assert state.next == ("flaky",)
assert (
state.values
== await graph_compare.ainvoke(input, {"configurable": {"thread_id": "2"}})
== [
"0",
"1",
"2|Control(goto=Send(node='2', arg=3))",
"2|Control(goto=Send(node='flaky', arg=4))",
"2|3",
]
)
# check history
history = [c async for c in graph.aget_state_history(config)]
assert len(history) == 2
# check messages
assert orch_msgs == [MessageToOrchestrator(input=input, config=config)] + [
{
"config": {
"callbacks": None,
"configurable": {
"__pregel_ensure_latest": True,
"__pregel_dedupe_tasks": True,
"__pregel_resuming": False,
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_ns": "",
"thread_id": "1",
},
"metadata": AnyDict(),
"recursion_limit": 25,
"tags": [],
},
"input": None,
"finally_send": None,
}
for c in reversed(history)
for _ in c.tasks
]
assert exec_msgs == [
{
"config": {
"callbacks": None,
"configurable": {
"__pregel_ensure_latest": True,
"__pregel_dedupe_tasks": True,
"__pregel_resuming": False,
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_ns": "",
"thread_id": "1",
},
"metadata": AnyDict(),
"recursion_limit": 25,
"tags": [],
},
"task": {
"id": t.id,
"path": _convert_path(t.path),
},
"finally_send": None,
}
for c in reversed(history)
for t in c.tasks
]
# resume the thread
async with AIOKafkaProducer(value_serializer=serde.dumps) as producer:
await producer.send_and_wait(
topics.orchestrator,
MessageToOrchestrator(input=None, config=config),
)
orch_msgs, exec_msgs = await drain_topics_async(topics, graph)
# check final state
state = await graph.aget_state(config)
assert state.next == ()
assert (
state.values
== await graph_compare.ainvoke(None, {"configurable": {"thread_id": "2"}})
== [
"0",
"1",
"2|Control(goto=Send(node='2', arg=3))",
"2|Control(goto=Send(node='flaky', arg=4))",
"2|3",
"flaky|4",
"3",
"3.1",
]
)
# check history
history = [c async for c in graph.aget_state_history(config)]
assert len(history) == 4
# check executions
# node "2" doesn't get called again, as we recover writes saved before
assert graph.builder.nodes["2"].runnable.func.ticks == 3
# node "flaky" gets called again, as it was interrupted
assert graph.builder.nodes["flaky"].runnable.func.ticks == 2
def _convert_path(
path: tuple[Union[str, int, tuple], ...],
) -> list[Union[str, int, list]]:
return list(_convert_path(p) if isinstance(p, tuple) else p for p in path)