update memory saver + tests

This commit is contained in:
vbarda
2024-08-05 16:53:33 -04:00
parent 64c30508d9
commit 55717c9bde
3 changed files with 67 additions and 73 deletions
+27 -1
View File
@@ -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.
+11 -31
View File
@@ -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:
+29 -41
View File
@@ -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: