Update stop condition for drain

This commit is contained in:
Nuno Campos
2024-09-10 16:52:09 -07:00
parent bb10acbe28
commit 5de46f7472
3 changed files with 25 additions and 44 deletions
+20 -26
View File
@@ -1,13 +1,11 @@
import asyncio
from typing import Callable, Optional, TypeVar
from typing import Optional, TypeVar
import anyio
from aiokafka import AIOKafkaConsumer
from langchain_core.runnables import RunnableConfig
from typing_extensions import ParamSpec
from langgraph.pregel import Pregel
from langgraph.pregel.types import StateSnapshot
from langgraph.scheduler.kafka.executor import KafkaExecutor
from langgraph.scheduler.kafka.orchestrator import KafkaOrchestrator
from langgraph.scheduler.kafka.types import MessageToOrchestrator, Topics
@@ -17,31 +15,38 @@ R = TypeVar("R")
async def drain_topics(
topics: Topics,
graph: Pregel,
config: RunnableConfig,
*,
until: Callable[[StateSnapshot], bool],
debug: bool = False,
topics: Topics, graph: Pregel, *, debug: bool = False
) -> tuple[list[MessageToOrchestrator], list[MessageToOrchestrator]]:
scope: Optional[anyio.CancelScope] = None
orch_msgs = []
exec_msgs = []
errors = []
def done() -> bool:
return (
len(orch_msgs) > 0
and len(exec_msgs) > 0
and not orch_msgs[-1]
and not exec_msgs[-1]
)
async def orchestrator() -> None:
async with KafkaOrchestrator(graph, topics) as orch:
async for msgs in orch:
orch_msgs.extend(msgs)
orch_msgs.append(msgs)
if debug:
print("\n---\norch", len(msgs), msgs)
if done():
scope.cancel()
async def executor() -> None:
async with KafkaExecutor(graph, topics) as exec:
async for msgs in exec:
exec_msgs.extend(msgs)
exec_msgs.append(msgs)
if debug:
print("\n---\nexec", len(msgs), msgs)
if done():
scope.cancel()
async def error_consumer() -> None:
async with AIOKafkaConsumer(topics.error) as consumer:
@@ -50,18 +55,8 @@ async def drain_topics(
if scope:
scope.cancel()
async def poller(expected_next: tuple[str, ...]) -> None:
while True:
await asyncio.sleep(0.5)
state = await graph.aget_state(config)
if until(state):
break
if scope:
scope.cancel()
# start error consumer and poller
# start error consumer
error_task = asyncio.create_task(error_consumer(), name="error_consumer")
poller_task = asyncio.create_task(poller(()), name="poller")
# run the orchestrator and executor until break_when
async with anyio.create_task_group() as tg:
@@ -70,16 +65,15 @@ async def drain_topics(
tg.start_soon(orchestrator, name="orchestrator")
tg.start_soon(executor, name="executor")
# cancel error consumer and poller
# cancel error consumer
error_task.cancel()
poller_task.cancel()
try:
await asyncio.gather(error_task, poller_task)
await error_task
except asyncio.CancelledError:
pass
# check no errors
assert not errors, errors
return orch_msgs, exec_msgs
return [m for mm in orch_msgs for m in mm], [m for mm in exec_msgs for m in mm]
+3 -9
View File
@@ -99,9 +99,7 @@ async def test_fanout_graph(topics: Topics, checkpointer: BaseCheckpointSaver) -
)
# drain topics
orch_msgs, exec_msgs = await drain_topics(
topics, graph, config, until=lambda s: s.values and s.next == ()
)
orch_msgs, exec_msgs = await drain_topics(topics, graph)
# check state
state = await graph.aget_state(config)
@@ -184,9 +182,7 @@ async def test_fanout_graph_w_interrupt(
MessageToOrchestrator(input=input, config=config),
)
orch_msgs, exec_msgs = await drain_topics(
topics, graph, config, until=lambda s: s.values and s.next == ("qa",)
)
orch_msgs, exec_msgs = await drain_topics(topics, graph)
# check interrupted state
state = await graph.aget_state(config)
@@ -262,9 +258,7 @@ async def test_fanout_graph_w_interrupt(
MessageToOrchestrator(input=None, config=config),
)
orch_msgs, exec_msgs = await drain_topics(
topics, graph, config, until=lambda s: s.values and s.next == ()
)
orch_msgs, exec_msgs = await drain_topics(topics, graph)
# check final state
state = await graph.aget_state(config)
+2 -9
View File
@@ -129,12 +129,7 @@ async def test_subgraph_w_interrupt(
MessageToOrchestrator(input=input, config=config),
)
orch_msgs, exec_msgs = await drain_topics(
topics,
graph,
config,
until=lambda state: state.next == ("weather_graph",),
)
orch_msgs, exec_msgs = await drain_topics(topics, graph)
# check interrupted state
state = await graph.aget_state(config)
@@ -419,9 +414,7 @@ async def test_subgraph_w_interrupt(
MessageToOrchestrator(input=None, config=config),
)
orch_msgs, exec_msgs = await drain_topics(
topics, graph, config, until=lambda state: state.next == ()
)
orch_msgs, exec_msgs = await drain_topics(topics, graph)
# check final state
state = await graph.aget_state(config)