Update sqlite error for async (#481)

This commit is contained in:
William FH
2024-05-16 09:51:07 -07:00
committed by GitHub
parent f4a2da9095
commit 9da2c7fecc
2 changed files with 175 additions and 23 deletions
+104 -22
View File
@@ -1,7 +1,8 @@
import asyncio
import functools
from contextlib import AbstractAsyncContextManager
from types import TracebackType
from typing import AsyncIterator, Optional
from typing import AsyncIterator, Iterator, Optional, TypeVar
import aiosqlite
from langchain_core.runnables import RunnableConfig
@@ -16,6 +17,22 @@ from langgraph.checkpoint.base import (
)
from langgraph.checkpoint.sqlite import JsonPlusSerializerCompat, search_where
T = TypeVar("T", bound=callable)
def not_implemented_sync_method(func: T) -> T:
@functools.wraps(func)
def wrapper(*args, **kwargs):
raise NotImplementedError(
"The AsyncSqliteSaver does not support synchronous methods. "
"Consider using the SqliteSaver instead.\n"
"from langgraph.checkpoint.sqlite import SqliteSaver\n"
"See https://langchain-ai.github.io/langgraph/reference/checkpoints/#sqlitesaver "
"for more information."
)
return wrapper
class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
"""An asynchronous checkpoint saver that stores checkpoints in a SQLite database.
@@ -34,7 +51,6 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
serde (Optional[SerializerProtocol]): The serializer to use for serializing and deserializing checkpoints. Defaults to JsonPlusSerializerCompat.
Examples:
Usage within a StateGraph:
```pycon
>>> import asyncio
@@ -113,6 +129,54 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
if self.is_setup:
return await self.conn.close()
@not_implemented_sync_method
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
"""Get a checkpoint tuple from the database.
Note:
This method is not implemented for the AsyncSqliteSaver. Use `aget` instead.
Or consider using the [SqliteSaver](#sqlitesaver) checkpointer.
"""
@not_implemented_sync_method
def list(
self,
config: RunnableConfig,
*,
before: Optional[RunnableConfig] = None,
limit: Optional[int] = None,
) -> Iterator[CheckpointTuple]:
"""List checkpoints from the database.
Note:
This method is not implemented for the AsyncSqliteSaver. Use `alist` instead.
Or consider using the [SqliteSaver](#sqlitesaver) checkpointer.
"""
@not_implemented_sync_method
def search(
self,
metadata_filter: CheckpointMetadata,
*,
before: Optional[RunnableConfig] = None,
limit: Optional[int] = None,
) -> Iterator[CheckpointTuple]:
"""Search for checkpoints by metadata.
Note:
This method is not implemented for the AsyncSqliteSaver. Use `asearch` instead.
Or consider using the [SqliteSaver](#sqlitesaver) checkpointer.
"""
@not_implemented_sync_method
def put(
self,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: CheckpointMetadata,
) -> RunnableConfig:
"""Save a checkpoint to the database. FOO"""
async def setup(self) -> None:
"""Set up the checkpoint database asynchronously.
@@ -169,14 +233,16 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
config,
self.serde.loads(value[0]),
self.serde.loads(value[2]) if value[2] is not None else {},
{
"configurable": {
"thread_id": config["configurable"]["thread_id"],
"thread_ts": value[1],
(
{
"configurable": {
"thread_id": config["configurable"]["thread_id"],
"thread_ts": value[1],
}
}
}
if value[1]
else None,
if value[1]
else None
),
)
else:
async with self.conn.execute(
@@ -193,14 +259,16 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
},
self.serde.loads(value[3]),
self.serde.loads(value[4]) if value[4] is not None else {},
{
"configurable": {
"thread_id": value[0],
"thread_ts": value[2],
(
{
"configurable": {
"thread_id": value[0],
"thread_ts": value[2],
}
}
}
if value[2]
else None,
if value[2]
else None
),
)
async def alist(
@@ -247,9 +315,16 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
{"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}},
self.serde.loads(value),
self.serde.loads(metadata) if metadata is not None else {},
{"configurable": {"thread_id": thread_id, "thread_ts": parent_ts}}
if parent_ts
else None,
(
{
"configurable": {
"thread_id": thread_id,
"thread_ts": parent_ts,
}
}
if parent_ts
else None
),
)
async def asearch(
@@ -291,9 +366,16 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
{"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}},
self.serde.loads(value),
self.serde.loads(metadata) if metadata is not None else {},
{"configurable": {"thread_id": thread_id, "thread_ts": parent_ts}}
if parent_ts
else None,
(
{
"configurable": {
"thread_id": thread_id,
"thread_ts": parent_ts,
}
}
if parent_ts
else None
),
)
async def aput(
+71 -1
View File
@@ -4,7 +4,7 @@ import sqlite3
import threading
from contextlib import AbstractContextManager, contextmanager
from types import TracebackType
from typing import Any, Iterator, Optional, Tuple
from typing import Any, AsyncIterator, Iterator, Optional, Tuple
from langchain_core.runnables import RunnableConfig
from typing_extensions import Self
@@ -47,6 +47,17 @@ class JsonPlusSerializerCompat(JsonPlusSerializer):
return super().loads(data)
_AIO_ERROR_MSG = (
"The SqliteSaver does not support async methods. "
"Consider using AsyncSqliteSaver instead.\n"
"from langgraph.checkpoint.aiosqlite import AsyncSqliteSaver\n"
"Note: AsyncSqliteSaver requires the aiosqlite package to use.\n"
"Install with:\n`pip install aiosqlite`\n"
"See https://langchain-ai.github.io/langgraph/reference/checkpoints/#asyncsqlitesaver"
"for more information."
)
class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
"""A checkpoint saver that stores checkpoints in a SQLite database.
@@ -524,3 +535,62 @@ def _metadata_predicate(
# predicate contains an extra trailing space
return (predicate, param_values)
async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
"""Get a checkpoint tuple from the database asynchronously.
Note:
This async method is not supported by the SqliteSaver class.
Use get_tuple() instead, or consider using [AsyncSqliteSaver](#asyncsqlitesaver).
"""
raise NotImplementedError(_AIO_ERROR_MSG)
def alist(
self,
config: RunnableConfig,
*,
before: Optional[RunnableConfig] = None,
limit: Optional[int] = None,
) -> AsyncIterator[CheckpointTuple]:
"""List checkpoints from the database asynchronously.
Note:
This async method is not supported by the SqliteSaver class.
Use list() instead, or consider using [AsyncSqliteSaver](#asyncsqlitesaver).
"""
raise NotImplementedError(_AIO_ERROR_MSG)
yield
def asearch(
self,
metadata_filter: CheckpointMetadata,
*,
before: Optional[RunnableConfig] = None,
limit: Optional[int] = None,
) -> AsyncIterator[CheckpointTuple]:
"""Search for checkpoints by metadata asynchronously.
Note:
This async method is not supported by the SqliteSaver class.
Use search() instead, or consider using [AsyncSqliteSaver](#asyncsqlitesaver).
"""
raise NotImplementedError(_AIO_ERROR_MSG)
yield
async def aput(
self,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: CheckpointMetadata,
) -> RunnableConfig:
"""Save a checkpoint to the database asynchronously.
Note:
This async method is not supported by the SqliteSaver class.
Use put() instead, or consider using [AsyncSqliteSaver](#asyncsqlitesaver).
"""
raise NotImplementedError(_AIO_ERROR_MSG)