mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-11 12:17:53 +02:00
update memory saver + tests
This commit is contained in:
@@ -1,6 +1,8 @@
|
||||
import asyncio
|
||||
from collections import defaultdict
|
||||
from contextlib import AbstractAsyncContextManager, AbstractContextManager
|
||||
from functools import partial
|
||||
from types import TracebackType
|
||||
from typing import Any, AsyncIterator, Dict, Iterator, List, Optional, Tuple
|
||||
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
@@ -15,7 +17,9 @@ from langgraph.checkpoint.base import (
|
||||
)
|
||||
|
||||
|
||||
class MemorySaver(BaseCheckpointSaver):
|
||||
class MemorySaver(
|
||||
BaseCheckpointSaver, AbstractContextManager, AbstractAsyncContextManager
|
||||
):
|
||||
"""An in-memory checkpoint saver.
|
||||
|
||||
This checkpoint saver stores checkpoints in memory using a defaultdict.
|
||||
@@ -57,6 +61,28 @@ class MemorySaver(BaseCheckpointSaver):
|
||||
self.storage = defaultdict(lambda: defaultdict(dict))
|
||||
self.writes = defaultdict(list)
|
||||
|
||||
def __enter__(self) -> "MemorySaver":
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exc_type: Optional[type[BaseException]],
|
||||
exc_value: Optional[BaseException],
|
||||
traceback: Optional[TracebackType],
|
||||
) -> Optional[bool]:
|
||||
return
|
||||
|
||||
async def __aenter__(self) -> "MemorySaver":
|
||||
return self
|
||||
|
||||
async def __aexit__(
|
||||
self,
|
||||
__exc_type: Optional[type[BaseException]],
|
||||
__exc_value: Optional[BaseException],
|
||||
__traceback: Optional[TracebackType],
|
||||
) -> Optional[bool]:
|
||||
return
|
||||
|
||||
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
||||
"""Get a checkpoint tuple from the in-memory storage.
|
||||
|
||||
|
||||
@@ -8,7 +8,6 @@ from contextlib import contextmanager
|
||||
from typing import (
|
||||
Annotated,
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
Generator,
|
||||
List,
|
||||
@@ -558,7 +557,7 @@ def test_invoke_two_processes_in_out(mocker: MockerFixture) -> None:
|
||||
def test_invoke_two_processes_in_out_interrupt(
|
||||
checkpointer: BaseCheckpointSaver, mocker: MockerFixture
|
||||
) -> None:
|
||||
try:
|
||||
with checkpointer as checkpointer:
|
||||
add_one = mocker.Mock(side_effect=lambda x: x + 1)
|
||||
one = Channel.subscribe_to("input") | add_one | Channel.write_to("inbox")
|
||||
two = Channel.subscribe_to("inbox") | add_one | Channel.write_to("output")
|
||||
@@ -758,9 +757,6 @@ def test_invoke_two_processes_in_out_interrupt(
|
||||
assert [c for c in app.stream(None, fork_config, stream_mode="updates")] == [
|
||||
{"one": {"inbox": 4}}
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__exit__"):
|
||||
checkpointer.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
@@ -777,7 +773,7 @@ def test_invoke_two_processes_in_out_interrupt(
|
||||
def test_fork_always_re_runs_nodes(
|
||||
checkpointer: BaseCheckpointSaver, mocker: MockerFixture
|
||||
) -> None:
|
||||
try:
|
||||
with checkpointer as checkpointer:
|
||||
add_one = mocker.Mock(side_effect=lambda _: 1)
|
||||
|
||||
builder = StateGraph(Annotated[int, operator.add])
|
||||
@@ -930,9 +926,6 @@ def test_fork_always_re_runs_nodes(
|
||||
{"add_one": 1},
|
||||
{"add_one": 1},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__exit__"):
|
||||
checkpointer.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
|
||||
@@ -1268,7 +1261,7 @@ def test_invoke_checkpoint(mocker: MockerFixture) -> None:
|
||||
],
|
||||
)
|
||||
def test_pending_writes_resume(checkpointer: BaseCheckpointSaver) -> None:
|
||||
try:
|
||||
with checkpointer as checkpointer:
|
||||
|
||||
class State(TypedDict):
|
||||
value: Annotated[int, operator.add]
|
||||
@@ -1340,9 +1333,6 @@ def test_pending_writes_resume(checkpointer: BaseCheckpointSaver) -> None:
|
||||
two.rtn = {"value": 3}
|
||||
# both the pending write and the new write were applied, 1 + 2 + 3 = 6
|
||||
assert graph.invoke(None, thread1) == {"value": 6}
|
||||
finally:
|
||||
if getattr(checkpointer, "__exit__", None):
|
||||
checkpointer.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_cond_edge_after_send() -> None:
|
||||
@@ -7514,12 +7504,12 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None:
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.repeat(10)
|
||||
# @pytest.mark.repeat(10)
|
||||
@pytest.mark.parametrize(
|
||||
"checkpointer_fct",
|
||||
"checkpointer",
|
||||
[
|
||||
lambda: MemorySaverAssertImmutable(put_sleep=0.2),
|
||||
lambda: SqliteSaver.from_conn_string(":memory:"),
|
||||
MemorySaverAssertImmutable(put_sleep=0.2),
|
||||
SqliteSaver.from_conn_string(":memory:"),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
@@ -7527,10 +7517,9 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None:
|
||||
],
|
||||
)
|
||||
def test_nested_graph_interrupts(
|
||||
checkpointer_fct: Callable[[], BaseCheckpointSaver],
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
try:
|
||||
checkpointer = checkpointer_fct()
|
||||
with checkpointer as checkpointer:
|
||||
|
||||
class InnerState(TypedDict):
|
||||
my_key: str
|
||||
@@ -8708,9 +8697,6 @@ def test_nested_graph_interrupts(
|
||||
parent_config=None,
|
||||
),
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__exit__"):
|
||||
checkpointer.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
@@ -8725,7 +8711,7 @@ def test_nested_graph_interrupts(
|
||||
],
|
||||
)
|
||||
def test_nested_graph_interrupts_parallel(checkpointer: BaseCheckpointSaver) -> None:
|
||||
try:
|
||||
with checkpointer as checkpointer:
|
||||
|
||||
class InnerState(TypedDict):
|
||||
my_key: Annotated[str, operator.add]
|
||||
@@ -8840,9 +8826,6 @@ def test_nested_graph_interrupts_parallel(checkpointer: BaseCheckpointSaver) ->
|
||||
"my_key": "got here and there and parallel and back again",
|
||||
},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__exit__"):
|
||||
checkpointer.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.skip
|
||||
@@ -8858,7 +8841,7 @@ def test_nested_graph_interrupts_parallel(checkpointer: BaseCheckpointSaver) ->
|
||||
],
|
||||
)
|
||||
def test_doubly_nested_graph_interrupts(checkpointer: BaseCheckpointSaver) -> None:
|
||||
try:
|
||||
with checkpointer as checkpointer:
|
||||
|
||||
class State(TypedDict):
|
||||
my_key: str
|
||||
@@ -8944,9 +8927,6 @@ def test_doubly_nested_graph_interrupts(checkpointer: BaseCheckpointSaver) -> No
|
||||
"my_key": "hi my value here and there and back again",
|
||||
},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__exit__"):
|
||||
checkpointer.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_repeat_condition(snapshot: SnapshotAssertion) -> None:
|
||||
|
||||
@@ -2,13 +2,13 @@ import asyncio
|
||||
import json
|
||||
import operator
|
||||
from collections import Counter
|
||||
from contextlib import asynccontextmanager, contextmanager
|
||||
from contextlib import AbstractAsyncContextManager, asynccontextmanager, contextmanager
|
||||
from types import TracebackType
|
||||
from typing import (
|
||||
Annotated,
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
AsyncIterator,
|
||||
Callable,
|
||||
Dict,
|
||||
Generator,
|
||||
List,
|
||||
@@ -66,6 +66,19 @@ from tests.memory_assert import (
|
||||
from tests.messages import _AnyIdAIMessage, _AnyIdHumanMessage
|
||||
|
||||
|
||||
class NoneContextManager(AbstractAsyncContextManager):
|
||||
async def __aenter__(self) -> None:
|
||||
return None
|
||||
|
||||
async def __aexit__(
|
||||
self,
|
||||
__exc_type: Optional[type[BaseException]],
|
||||
__exc_value: Optional[BaseException],
|
||||
__traceback: Optional[TracebackType],
|
||||
) -> Optional[bool]:
|
||||
return
|
||||
|
||||
|
||||
async def test_checkpoint_errors() -> None:
|
||||
class FaultyGetCheckpointer(MemorySaver):
|
||||
async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
||||
@@ -236,7 +249,7 @@ async def test_step_timeout_on_stream_hang() -> None:
|
||||
[
|
||||
MemorySaverAssertImmutable(),
|
||||
AsyncSqliteSaver.from_conn_string(":memory:"),
|
||||
None,
|
||||
NoneContextManager(),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
@@ -247,7 +260,7 @@ async def test_step_timeout_on_stream_hang() -> None:
|
||||
async def test_cancel_graph_astream(
|
||||
checkpointer: Optional[BaseCheckpointSaver],
|
||||
) -> None:
|
||||
try:
|
||||
async with checkpointer as checkpointer:
|
||||
|
||||
class State(TypedDict):
|
||||
value: Annotated[int, operator.add]
|
||||
@@ -311,9 +324,6 @@ async def test_cancel_graph_astream(
|
||||
"alittlewhile",
|
||||
)
|
||||
assert state.metadata == {"source": "loop", "step": 0, "writes": None}
|
||||
finally:
|
||||
if getattr(checkpointer, "__aexit__", None):
|
||||
await checkpointer.__aexit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
@@ -321,7 +331,7 @@ async def test_cancel_graph_astream(
|
||||
[
|
||||
MemorySaverAssertImmutable(),
|
||||
AsyncSqliteSaver.from_conn_string(":memory:"),
|
||||
None,
|
||||
NoneContextManager(),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
@@ -332,7 +342,7 @@ async def test_cancel_graph_astream(
|
||||
async def test_cancel_graph_astream_events_v2(
|
||||
checkpointer: Optional[BaseCheckpointSaver],
|
||||
) -> None:
|
||||
try:
|
||||
async with checkpointer as checkpointer:
|
||||
|
||||
class State(TypedDict):
|
||||
value: int
|
||||
@@ -402,9 +412,6 @@ async def test_cancel_graph_astream_events_v2(
|
||||
"step": 1,
|
||||
"writes": {"alittlewhile": {"value": 2}},
|
||||
}
|
||||
finally:
|
||||
if getattr(checkpointer, "__aexit__", None):
|
||||
await checkpointer.__aexit__(None, None, None)
|
||||
|
||||
|
||||
async def test_node_schemas_custom_output() -> None:
|
||||
@@ -679,7 +686,7 @@ async def test_invoke_two_processes_in_out(mocker: MockerFixture) -> None:
|
||||
async def test_invoke_two_processes_in_out_interrupt(
|
||||
checkpointer: BaseCheckpointSaver, mocker: MockerFixture
|
||||
) -> None:
|
||||
try:
|
||||
async with checkpointer as checkpointer:
|
||||
add_one = mocker.Mock(side_effect=lambda x: x + 1)
|
||||
one = Channel.subscribe_to("input") | add_one | Channel.write_to("inbox")
|
||||
two = Channel.subscribe_to("inbox") | add_one | Channel.write_to("output")
|
||||
@@ -880,9 +887,6 @@ async def test_invoke_two_processes_in_out_interrupt(
|
||||
assert [
|
||||
c async for c in app.astream(None, fork_config, stream_mode="updates")
|
||||
] == [{"one": {"inbox": 4}}]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__aexit__"):
|
||||
await checkpointer.__aexit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
@@ -899,7 +903,7 @@ async def test_invoke_two_processes_in_out_interrupt(
|
||||
async def test_fork_always_re_runs_nodes(
|
||||
checkpointer: BaseCheckpointSaver, mocker: MockerFixture
|
||||
) -> None:
|
||||
try:
|
||||
async with checkpointer as checkpointer:
|
||||
add_one = mocker.Mock(side_effect=lambda _: 1)
|
||||
|
||||
builder = StateGraph(Annotated[int, operator.add])
|
||||
@@ -1061,9 +1065,6 @@ async def test_fork_always_re_runs_nodes(
|
||||
{"add_one": 1},
|
||||
{"add_one": 1},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__aexit__"):
|
||||
await checkpointer.__aexit__(None, None, None)
|
||||
|
||||
|
||||
async def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
|
||||
@@ -1406,7 +1407,7 @@ async def test_invoke_checkpoint(mocker: MockerFixture) -> None:
|
||||
],
|
||||
)
|
||||
async def test_pending_writes_resume(checkpointer: BaseCheckpointSaver) -> None:
|
||||
try:
|
||||
async with checkpointer as checkpointer:
|
||||
|
||||
class State(TypedDict):
|
||||
value: Annotated[int, operator.add]
|
||||
@@ -1478,9 +1479,6 @@ async def test_pending_writes_resume(checkpointer: BaseCheckpointSaver) -> None:
|
||||
two.rtn = {"value": 3}
|
||||
# both the pending write and the new write were applied, 1 + 2 + 3 = 6
|
||||
assert await graph.ainvoke(None, thread1) == {"value": 6}
|
||||
finally:
|
||||
if getattr(checkpointer, "__aexit__", None):
|
||||
await checkpointer.__aexit__(None, None, None)
|
||||
|
||||
|
||||
async def test_cond_edge_after_send() -> None:
|
||||
@@ -5991,10 +5989,10 @@ async def test_nested_graph(snapshot: SnapshotAssertion) -> None:
|
||||
|
||||
@pytest.mark.repeat(10)
|
||||
@pytest.mark.parametrize(
|
||||
"checkpointer_fct",
|
||||
"checkpointer",
|
||||
[
|
||||
lambda: MemorySaverAssertImmutable(put_sleep=0.2),
|
||||
lambda: AsyncSqliteSaver.from_conn_string(":memory:"),
|
||||
MemorySaverAssertImmutable(put_sleep=0.2),
|
||||
AsyncSqliteSaver.from_conn_string(":memory:"),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
@@ -6002,10 +6000,9 @@ async def test_nested_graph(snapshot: SnapshotAssertion) -> None:
|
||||
],
|
||||
)
|
||||
async def test_nested_graph_interrupts(
|
||||
checkpointer_fct: Callable[[], BaseCheckpointSaver],
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
try:
|
||||
checkpointer = checkpointer_fct()
|
||||
async with checkpointer as checkpointer:
|
||||
|
||||
class InnerState(TypedDict):
|
||||
my_key: str
|
||||
@@ -7196,9 +7193,6 @@ async def test_nested_graph_interrupts(
|
||||
parent_config=None,
|
||||
),
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__aexit__"):
|
||||
await checkpointer.__aexit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
@@ -7215,7 +7209,7 @@ async def test_nested_graph_interrupts(
|
||||
async def test_nested_graph_interrupts_parallel(
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
try:
|
||||
async with checkpointer as checkpointer:
|
||||
|
||||
class InnerState(TypedDict):
|
||||
my_key: Annotated[str, operator.add]
|
||||
@@ -7332,9 +7326,6 @@ async def test_nested_graph_interrupts_parallel(
|
||||
"my_key": "got here and there and parallel and back again",
|
||||
},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__aexit__"):
|
||||
await checkpointer.__aexit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.skip
|
||||
@@ -7352,7 +7343,7 @@ async def test_nested_graph_interrupts_parallel(
|
||||
async def test_doubly_nested_graph_interrupts(
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
try:
|
||||
async with checkpointer as checkpointer:
|
||||
|
||||
class State(TypedDict):
|
||||
my_key: str
|
||||
@@ -7443,9 +7434,6 @@ async def test_doubly_nested_graph_interrupts(
|
||||
"my_key": "hi my value here and there and back again",
|
||||
},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__aexit__"):
|
||||
await checkpointer.__aexit__(None, None, None)
|
||||
|
||||
|
||||
async def test_checkpoint_metadata() -> None:
|
||||
|
||||
Reference in New Issue
Block a user