This commit is contained in:
Nuno Campos
2025-01-15 11:02:10 -08:00
parent eed577ee2a
commit 31734fb792
@@ -1,17 +1,8 @@
import asyncio
import random
from collections.abc import AsyncIterator, Iterator, Sequence
from contextlib import asynccontextmanager
from typing import (
Any,
AsyncIterator,
Callable,
Dict,
Iterator,
Optional,
Sequence,
Tuple,
TypeVar,
)
from typing import Any, Callable, Optional, TypeVar
import aiosqlite
from langchain_core.runnables import RunnableConfig
@@ -173,7 +164,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
self,
config: Optional[RunnableConfig],
*,
filter: Optional[Dict[str, Any]] = None,
filter: Optional[dict[str, Any]] = None,
before: Optional[RunnableConfig] = None,
limit: Optional[int] = None,
) -> Iterator[CheckpointTuple]:
@@ -207,7 +198,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
while True:
try:
yield asyncio.run_coroutine_threadsafe(
anext(aiter_),
anext(aiter_), # noqa: F821
self.loop,
).result()
except StopAsyncIteration:
@@ -239,10 +230,14 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
).result()
def put_writes(
self, config: RunnableConfig, writes: Sequence[Tuple[str, Any]], task_id: str
self,
config: RunnableConfig,
writes: Sequence[tuple[str, Any]],
task_id: str,
task_path: str = "",
) -> None:
return asyncio.run_coroutine_threadsafe(
self.aput_writes(config, writes, task_id), self.loop
self.aput_writes(config, writes, task_id, task_path), self.loop
).result()
async def setup(self) -> None:
@@ -372,7 +367,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
self,
config: Optional[RunnableConfig],
*,
filter: Optional[Dict[str, Any]] = None,
filter: Optional[dict[str, Any]] = None,
before: Optional[RunnableConfig] = None,
limit: Optional[int] = None,
) -> AsyncIterator[CheckpointTuple]:
@@ -496,7 +491,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
async def aput_writes(
self,
config: RunnableConfig,
writes: Sequence[Tuple[str, Any]],
writes: Sequence[tuple[str, Any]],
task_id: str,
task_path: str = "",
) -> None: