mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-07 18:27:52 +02:00
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:
@@ -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]:
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user