Merge pull request #1228 from langchain-ai/vb/update-sqlite

checkpoint: stop using sqlite checkpointers as context managers, make memorysaver a context manager
This commit is contained in:
Nuno Campos
2024-08-05 14:09:42 -07:00
committed by GitHub
5 changed files with 82 additions and 108 deletions
@@ -1,12 +1,10 @@
import sqlite3
import threading
from contextlib import AbstractContextManager, contextmanager
from contextlib import contextmanager
from hashlib import md5
from types import TracebackType
from typing import Any, AsyncIterator, Dict, Iterator, Optional, Sequence, Tuple
from langchain_core.runnables import RunnableConfig
from typing_extensions import Self
from langgraph.checkpoint.base import (
BaseCheckpointSaver,
@@ -32,7 +30,7 @@ _AIO_ERROR_MSG = (
)
class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
class SqliteSaver(BaseCheckpointSaver):
"""A checkpoint saver that stores checkpoints in a SQLite database.
Note:
@@ -82,43 +80,34 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
self.lock = threading.Lock()
@classmethod
def from_conn_string(cls, conn_string: str) -> "SqliteSaver":
@contextmanager
def from_conn_string(cls, conn_string: str) -> Iterator["SqliteSaver"]:
"""Create a new SqliteSaver instance from a connection string.
Args:
conn_string (str): The SQLite connection string.
Returns:
Yields:
SqliteSaver: A new SqliteSaver instance.
Examples:
In memory:
memory = SqliteSaver.from_conn_string(":memory:")
with SqliteSaver.from_conn_string(":memory:") as memory:
...
To disk:
memory = SqliteSaver.from_conn_string("checkpoints.sqlite")
with SqliteSaver.from_conn_string("checkpoints.sqlite") as memory:
...
"""
return SqliteSaver(
conn=sqlite3.connect(
conn_string,
# https://ricardoanderegg.com/posts/python-sqlite-thread-safety/
check_same_thread=False,
)
)
def __enter__(self) -> Self:
return self
def __exit__(
self,
__exc_type: Optional[type[BaseException]],
__exc_value: Optional[BaseException],
__traceback: Optional[TracebackType],
) -> Optional[bool]:
return self.conn.close()
with sqlite3.connect(
conn_string,
# https://ricardoanderegg.com/posts/python-sqlite-thread-safety/
check_same_thread=False,
) as conn:
yield SqliteSaver(conn)
def setup(self) -> None:
"""Set up the checkpoint database.
@@ -1,7 +1,6 @@
import asyncio
import functools
from contextlib import AbstractAsyncContextManager
from types import TracebackType
from contextlib import asynccontextmanager
from typing import (
Any,
AsyncIterator,
@@ -15,7 +14,6 @@ from typing import (
import aiosqlite
from langchain_core.runnables import RunnableConfig
from typing_extensions import Self
from langgraph.checkpoint.base import (
BaseCheckpointSaver,
@@ -45,7 +43,7 @@ def not_implemented_sync_method(func: T) -> T:
return wrapper
class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
class AsyncSqliteSaver(BaseCheckpointSaver):
"""An asynchronous checkpoint saver that stores checkpoints in a SQLite database.
This class provides an asynchronous interface for saving and retrieving checkpoints
@@ -135,29 +133,20 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
self.is_setup = False
@classmethod
def from_conn_string(cls, conn_string: str) -> "AsyncSqliteSaver":
@asynccontextmanager
async def from_conn_string(
cls, conn_string: str
) -> AsyncIterator["AsyncSqliteSaver"]:
"""Create a new AsyncSqliteSaver instance from a connection string.
Args:
conn_string (str): The SQLite connection string.
Returns:
Yields:
AsyncSqliteSaver: A new AsyncSqliteSaver instance.
"""
return AsyncSqliteSaver(conn=aiosqlite.connect(conn_string))
async def __aenter__(self) -> Self:
return self
async def __aexit__(
self,
__exc_type: Optional[type[BaseException]],
__exc_value: Optional[BaseException],
__traceback: Optional[TracebackType],
) -> Optional[bool]:
if self.is_setup:
return await self.conn.close()
async with aiosqlite.connect(conn_string) as conn:
yield AsyncSqliteSaver(conn)
@not_implemented_sync_method
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
+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.
+6 -25
View File
@@ -558,7 +558,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 +758,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 +774,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 +927,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 +1262,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 +1334,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:
@@ -7550,8 +7541,7 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None:
def test_nested_graph_interrupts(
checkpointer_fct: Callable[[], BaseCheckpointSaver],
) -> None:
try:
checkpointer = checkpointer_fct()
with checkpointer_fct() as checkpointer:
class InnerState(TypedDict):
my_key: str
@@ -8729,9 +8719,6 @@ def test_nested_graph_interrupts(
parent_config=None,
),
]
finally:
if hasattr(checkpointer, "__exit__"):
checkpointer.__exit__(None, None, None)
@pytest.mark.parametrize(
@@ -8746,7 +8733,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]
@@ -8861,9 +8848,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
@@ -8879,7 +8863,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
@@ -8965,9 +8949,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:
+25 -36
View File
@@ -2,7 +2,8 @@ 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,
@@ -66,6 +67,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 +250,7 @@ async def test_step_timeout_on_stream_hang() -> None:
[
MemorySaverAssertImmutable(),
AsyncSqliteSaver.from_conn_string(":memory:"),
None,
NoneContextManager(),
],
ids=[
"memory",
@@ -247,7 +261,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 +325,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 +332,7 @@ async def test_cancel_graph_astream(
[
MemorySaverAssertImmutable(),
AsyncSqliteSaver.from_conn_string(":memory:"),
None,
NoneContextManager(),
],
ids=[
"memory",
@@ -332,7 +343,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 +413,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 +687,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 +888,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 +904,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 +1066,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 +1408,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 +1480,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:
@@ -6143,8 +6142,7 @@ async def test_nested_graph(snapshot: SnapshotAssertion) -> None:
async def test_nested_graph_interrupts(
checkpointer_fct: Callable[[], BaseCheckpointSaver],
) -> None:
try:
checkpointer = checkpointer_fct()
async with checkpointer_fct() as checkpointer:
class InnerState(TypedDict):
my_key: str
@@ -7335,9 +7333,6 @@ async def test_nested_graph_interrupts(
parent_config=None,
),
]
finally:
if hasattr(checkpointer, "__aexit__"):
await checkpointer.__aexit__(None, None, None)
@pytest.mark.parametrize(
@@ -7354,7 +7349,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]
@@ -7471,9 +7466,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
@@ -7491,7 +7483,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
@@ -7582,9 +7574,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: