mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-26 17:42:24 +02:00
Nicer NotImplementedError in sqlite checkpointer (#598)
This commit is contained in:
@@ -451,42 +451,60 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
|
||||
}
|
||||
}
|
||||
|
||||
async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
||||
"""Get a checkpoint tuple from the database asynchronously.
|
||||
|
||||
def search_where(
|
||||
metadata_filter: CheckpointMetadata,
|
||||
before: Optional[RunnableConfig] = None,
|
||||
) -> Tuple[str, Tuple[Any, ...]]:
|
||||
"""Return WHERE clause predicates for (a)search() given metadata filter
|
||||
and `before` config.
|
||||
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)
|
||||
|
||||
This method returns a tuple of a string and a tuple of values. The string
|
||||
is the parametered WHERE clause predicate (including the WHERE keyword):
|
||||
"WHERE column1 = ? AND column2 IS ?". The tuple of values contains the
|
||||
values for each of the corresponding parameters.
|
||||
"""
|
||||
where = "WHERE "
|
||||
param_values = ()
|
||||
async def alist(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
*,
|
||||
before: Optional[RunnableConfig] = None,
|
||||
limit: Optional[int] = None,
|
||||
) -> AsyncIterator[CheckpointTuple]:
|
||||
"""List checkpoints from the database asynchronously.
|
||||
|
||||
# construct predicate for metadata filter
|
||||
metadata_predicate, metadata_values = _metadata_predicate(metadata_filter)
|
||||
if metadata_predicate != "":
|
||||
where += metadata_predicate
|
||||
param_values += metadata_values
|
||||
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
|
||||
|
||||
# construct predicate for `before`
|
||||
if before is not None:
|
||||
if metadata_predicate != "":
|
||||
where += "AND thread_ts < ? "
|
||||
else:
|
||||
where += "thread_ts < ? "
|
||||
async def asearch(
|
||||
self,
|
||||
metadata_filter: CheckpointMetadata,
|
||||
*,
|
||||
before: Optional[RunnableConfig] = None,
|
||||
limit: Optional[int] = None,
|
||||
) -> AsyncIterator[CheckpointTuple]:
|
||||
"""Search for checkpoints by metadata asynchronously.
|
||||
|
||||
param_values += (before["configurable"]["thread_ts"],)
|
||||
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
|
||||
|
||||
if where == "WHERE ":
|
||||
# no predicates, return an empty WHERE clause string
|
||||
return ("", ())
|
||||
else:
|
||||
return (where, param_values)
|
||||
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)
|
||||
|
||||
|
||||
def _metadata_predicate(
|
||||
@@ -538,60 +556,38 @@ def _metadata_predicate(
|
||||
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,
|
||||
def search_where(
|
||||
metadata_filter: CheckpointMetadata,
|
||||
*,
|
||||
before: Optional[RunnableConfig] = None,
|
||||
limit: Optional[int] = None,
|
||||
) -> AsyncIterator[CheckpointTuple]:
|
||||
"""Search for checkpoints by metadata asynchronously.
|
||||
) -> Tuple[str, Tuple[Any, ...]]:
|
||||
"""Return WHERE clause predicates for (a)search() given metadata filter
|
||||
and `before` config.
|
||||
|
||||
Note:
|
||||
This async method is not supported by the SqliteSaver class.
|
||||
Use search() instead, or consider using [AsyncSqliteSaver](#asyncsqlitesaver).
|
||||
This method returns a tuple of a string and a tuple of values. The string
|
||||
is the parametered WHERE clause predicate (including the WHERE keyword):
|
||||
"WHERE column1 = ? AND column2 IS ?". The tuple of values contains the
|
||||
values for each of the corresponding parameters.
|
||||
"""
|
||||
raise NotImplementedError(_AIO_ERROR_MSG)
|
||||
yield
|
||||
where = "WHERE "
|
||||
param_values = ()
|
||||
|
||||
# construct predicate for metadata filter
|
||||
metadata_predicate, metadata_values = _metadata_predicate(metadata_filter)
|
||||
if metadata_predicate != "":
|
||||
where += metadata_predicate
|
||||
param_values += metadata_values
|
||||
|
||||
async def aput(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
) -> RunnableConfig:
|
||||
"""Save a checkpoint to the database asynchronously.
|
||||
# construct predicate for `before`
|
||||
if before is not None:
|
||||
if metadata_predicate != "":
|
||||
where += "AND thread_ts < ? "
|
||||
else:
|
||||
where += "thread_ts < ? "
|
||||
|
||||
Note:
|
||||
This async method is not supported by the SqliteSaver class.
|
||||
Use put() instead, or consider using [AsyncSqliteSaver](#asyncsqlitesaver).
|
||||
"""
|
||||
raise NotImplementedError(_AIO_ERROR_MSG)
|
||||
param_values += (before["configurable"]["thread_ts"],)
|
||||
|
||||
if where == "WHERE ":
|
||||
# no predicates, return an empty WHERE clause string
|
||||
return ("", ())
|
||||
else:
|
||||
return (where, param_values)
|
||||
|
||||
@@ -3,7 +3,12 @@ from langchain_core.runnables import RunnableConfig
|
||||
|
||||
from langgraph.channels.base import create_checkpoint
|
||||
from langgraph.checkpoint.base import Checkpoint, CheckpointMetadata, empty_checkpoint
|
||||
from langgraph.checkpoint.sqlite import SqliteSaver, _metadata_predicate, search_where
|
||||
from langgraph.checkpoint.sqlite import (
|
||||
_AIO_ERROR_MSG,
|
||||
SqliteSaver,
|
||||
_metadata_predicate,
|
||||
search_where,
|
||||
)
|
||||
|
||||
|
||||
class TestSqliteSaver:
|
||||
@@ -98,3 +103,16 @@ class TestSqliteSaver:
|
||||
expected_predicate_3,
|
||||
expected_param_values_3,
|
||||
)
|
||||
|
||||
async def test_informative_async_errors(self):
|
||||
# call method / assertions
|
||||
with pytest.raises(NotImplementedError, match=_AIO_ERROR_MSG):
|
||||
await self.sqlite_saver.aget(self.config_1)
|
||||
with pytest.raises(NotImplementedError, match=_AIO_ERROR_MSG):
|
||||
await self.sqlite_saver.aget_tuple(self.config_1)
|
||||
with pytest.raises(NotImplementedError, match=_AIO_ERROR_MSG):
|
||||
async for _ in self.sqlite_saver.alist(self.config_1):
|
||||
pass
|
||||
with pytest.raises(NotImplementedError, match=_AIO_ERROR_MSG):
|
||||
async for _ in self.sqlite_saver.asearch(self.metadata_1):
|
||||
pass
|
||||
|
||||
Reference in New Issue
Block a user