mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-23 16:12:25 +02:00
Merge branch 'wfh/add_refcount_test' into wfh/gc_test_all
This commit is contained in:
@@ -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,
|
||||
@@ -7619,6 +7622,40 @@ def test_parallel_interrupts_double(
|
||||
assert len(events) == 5
|
||||
|
||||
|
||||
def test_pregel_loop_refcount():
|
||||
gc.collect()
|
||||
try:
|
||||
gc.disable()
|
||||
|
||||
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
|
||||
)
|
||||
finally:
|
||||
gc.enable()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_SYNC)
|
||||
def test_bulk_state_updates(
|
||||
request: pytest.FixtureRequest, checkpointer_name: str
|
||||
|
||||
@@ -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,
|
||||
@@ -7844,6 +7847,44 @@ async def test_handles_multiple_interrupts_from_tasks() -> None:
|
||||
assert result[1] == "Added Will!"
|
||||
|
||||
|
||||
async def test_pregel_loop_refcount():
|
||||
gc.collect()
|
||||
try:
|
||||
gc.disable()
|
||||
|
||||
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
|
||||
)
|
||||
finally:
|
||||
gc.enable()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_ASYNC)
|
||||
async def test_bulk_state_updates(checkpointer_name: str) -> None:
|
||||
async with awith_checkpointer(checkpointer_name) as checkpointer:
|
||||
|
||||
Reference in New Issue
Block a user