From f9780330a66d4938eea2f6be6fe24185f33edecf Mon Sep 17 00:00:00 2001 From: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com> Date: Tue, 18 Mar 2025 20:13:40 -0700 Subject: [PATCH 1/2] Add refcount test --- libs/langgraph/tests/test_pregel.py | 28 +++++++++++++++++++++++ libs/langgraph/tests/test_pregel_async.py | 28 +++++++++++++++++++++++ 2 files changed, 56 insertions(+) diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index d8d51a8a0..7de1d8018 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -1,5 +1,6 @@ import enum import functools +import gc import json import logging import operator @@ -62,7 +63,9 @@ from langgraph.graph import END, Graph, StateGraph from langgraph.graph.message import MessageGraph, MessagesState, add_messages from langgraph.prebuilt.tool_node import ToolNode from langgraph.pregel import Channel, GraphRecursionError, Pregel, StateSnapshot +from langgraph.pregel.loop import SyncPregelLoop from langgraph.pregel.retry import RetryPolicy +from langgraph.pregel.runner import PregelRunner from langgraph.store.base import BaseStore from langgraph.types import ( Command, @@ -7613,3 +7616,28 @@ def test_parallel_interrupts_double( assert invokes == 5 assert len(events) == 5 + + +def test_pregel_loop_refcount(): + class State(TypedDict): + messages: Annotated[list, add_messages] + + graph_builder = StateGraph(State) + + def chatbot(state: State): + return {"messages": [("ai", "HIYA")]} + + graph_builder.add_node("chatbot", chatbot) + graph_builder.set_entry_point("chatbot") + graph_builder.set_finish_point("chatbot") + graph = graph_builder.compile() + + for _ in range(5): + graph.invoke({"messages": [{"role": "user", "content": "hi"}]}) + assert ( + len([obj for obj in gc.get_objects() if isinstance(obj, SyncPregelLoop)]) + == 0 + ) + assert ( + len([obj for obj in gc.get_objects() if isinstance(obj, PregelRunner)]) == 0 + ) diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 60cbd2547..dc59abcd9 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -1,5 +1,6 @@ import asyncio import functools +import gc import logging import operator import random @@ -52,7 +53,9 @@ from langgraph.graph import END, Graph, StateGraph from langgraph.graph.message import MessagesState, add_messages from langgraph.prebuilt.tool_node import ToolNode from langgraph.pregel import Channel, GraphRecursionError, Pregel, StateSnapshot +from langgraph.pregel.loop import AsyncPregelLoop from langgraph.pregel.retry import RetryPolicy +from langgraph.pregel.runner import PregelRunner from langgraph.store.base import BaseStore from langgraph.types import ( Command, @@ -7838,3 +7841,28 @@ async def test_handles_multiple_interrupts_from_tasks() -> None: assert len(result) == 2 assert result[0] == "Added James!" assert result[1] == "Added Will!" + + +async def test_pregel_loop_refcount(): + class State(TypedDict): + messages: Annotated[list, add_messages] + + graph_builder = StateGraph(State) + + async def chatbot(state: State): + return {"messages": [("ai", "HIYA")]} + + graph_builder.add_node("chatbot", chatbot) + graph_builder.set_entry_point("chatbot") + graph_builder.set_finish_point("chatbot") + graph = graph_builder.compile() + + for _ in range(5): + await graph.ainvoke({"messages": [{"role": "user", "content": "hi"}]}) + assert ( + len([obj for obj in gc.get_objects() if isinstance(obj, AsyncPregelLoop)]) + == 0 + ) + assert ( + len([obj for obj in gc.get_objects() if isinstance(obj, PregelRunner)]) == 0 + ) From 4bfcd84cee47b4469cfeadf7be27adb2f991fb21 Mon Sep 17 00:00:00 2001 From: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com> Date: Wed, 19 Mar 2025 14:10:01 -0700 Subject: [PATCH 2/2] Cleanup ref count check --- libs/langgraph/tests/test_pregel.py | 45 ++++++++++++--------- libs/langgraph/tests/test_pregel_async.py | 49 ++++++++++++++--------- 2 files changed, 58 insertions(+), 36 deletions(-) diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 7de1d8018..f7a527927 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -7619,25 +7619,34 @@ def test_parallel_interrupts_double( def test_pregel_loop_refcount(): - class State(TypedDict): - messages: Annotated[list, add_messages] + gc.collect() + try: + gc.disable() - graph_builder = StateGraph(State) + class State(TypedDict): + messages: Annotated[list, add_messages] - def chatbot(state: State): - return {"messages": [("ai", "HIYA")]} + graph_builder = StateGraph(State) - graph_builder.add_node("chatbot", chatbot) - graph_builder.set_entry_point("chatbot") - graph_builder.set_finish_point("chatbot") - graph = graph_builder.compile() + def chatbot(state: State): + return {"messages": [("ai", "HIYA")]} - for _ in range(5): - graph.invoke({"messages": [{"role": "user", "content": "hi"}]}) - assert ( - len([obj for obj in gc.get_objects() if isinstance(obj, SyncPregelLoop)]) - == 0 - ) - assert ( - len([obj for obj in gc.get_objects() if isinstance(obj, PregelRunner)]) == 0 - ) + graph_builder.add_node("chatbot", chatbot) + graph_builder.set_entry_point("chatbot") + graph_builder.set_finish_point("chatbot") + graph = graph_builder.compile() + + for _ in range(5): + graph.invoke({"messages": [{"role": "user", "content": "hi"}]}) + assert ( + len( + [obj for obj in gc.get_objects() if isinstance(obj, SyncPregelLoop)] + ) + == 0 + ) + assert ( + len([obj for obj in gc.get_objects() if isinstance(obj, PregelRunner)]) + == 0 + ) + finally: + gc.enable() diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index dc59abcd9..7158edd6e 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -7844,25 +7844,38 @@ async def test_handles_multiple_interrupts_from_tasks() -> None: async def test_pregel_loop_refcount(): - class State(TypedDict): - messages: Annotated[list, add_messages] + gc.collect() + try: + gc.disable() - graph_builder = StateGraph(State) + class State(TypedDict): + messages: Annotated[list, add_messages] - async def chatbot(state: State): - return {"messages": [("ai", "HIYA")]} + graph_builder = StateGraph(State) - graph_builder.add_node("chatbot", chatbot) - graph_builder.set_entry_point("chatbot") - graph_builder.set_finish_point("chatbot") - graph = graph_builder.compile() + async def chatbot(state: State): + return {"messages": [("ai", "HIYA")]} - for _ in range(5): - await graph.ainvoke({"messages": [{"role": "user", "content": "hi"}]}) - assert ( - len([obj for obj in gc.get_objects() if isinstance(obj, AsyncPregelLoop)]) - == 0 - ) - assert ( - len([obj for obj in gc.get_objects() if isinstance(obj, PregelRunner)]) == 0 - ) + graph_builder.add_node("chatbot", chatbot) + graph_builder.set_entry_point("chatbot") + graph_builder.set_finish_point("chatbot") + graph = graph_builder.compile() + + for _ in range(5): + await graph.ainvoke({"messages": [{"role": "user", "content": "hi"}]}) + assert ( + len( + [ + obj + for obj in gc.get_objects() + if isinstance(obj, AsyncPregelLoop) + ] + ) + == 0 + ) + assert ( + len([obj for obj in gc.get_objects() if isinstance(obj, PregelRunner)]) + == 0 + ) + finally: + gc.enable()