Use by default in sqlite adapters

- Support loading pickled checkpoints for backwards compat
- Update copy_checkpoint to coerce values to defaultdicts where needed
This commit is contained in:
Nuno Campos
2024-04-24 16:08:41 -07:00
parent e45c461d99
commit f9646b6a57
6 changed files with 36 additions and 20 deletions
+13 -5
View File
@@ -16,9 +16,12 @@ from langgraph.checkpoint.base import (
)
class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
serde = pickle
# for backwards compat we continue to support loading pickled checkpoints
def is_pickled(value: bytes) -> bool:
return value.startswith(b"\x80") and value.endswith(b".")
class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
conn: aiosqlite.Connection
is_setup: bool
@@ -69,6 +72,11 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
self.is_setup = True
def _loads(self, value: bytes) -> Checkpoint:
if is_pickled(value):
return pickle.loads(value)
return self.serde.loads(value)
async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
await self.setup()
if config["configurable"].get("thread_ts"):
@@ -82,7 +90,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
if value := await cursor.fetchone():
return CheckpointTuple(
config,
self.serde.loads(value[0]),
self._loads(value[0]),
{
"configurable": {
"thread_id": config["configurable"]["thread_id"],
@@ -105,7 +113,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
"thread_ts": value[1],
}
},
self.serde.loads(value[3]),
self._loads(value[3]),
{
"configurable": {
"thread_id": value[0],
@@ -125,7 +133,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
async for thread_id, thread_ts, parent_ts, value in cursor:
yield CheckpointTuple(
{"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}},
self.serde.loads(value),
self._loads(value),
{"configurable": {"thread_id": thread_id, "thread_ts": parent_ts}}
if parent_ts
else None,
+7 -4
View File
@@ -1,6 +1,5 @@
from abc import ABC
from collections import defaultdict
from copy import deepcopy
from datetime import datetime, timezone
from typing import (
Any,
@@ -14,6 +13,7 @@ from typing import (
from langchain_core.runnables import ConfigurableFieldSpec, RunnableConfig
from langgraph.serde.base import SerializerProtocol
from langgraph.serde.jsonplus import JsonPlusSerializer
from langgraph.utils import StrEnum
@@ -63,8 +63,11 @@ def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint:
v=checkpoint["v"],
ts=checkpoint["ts"],
channel_values=checkpoint["channel_values"].copy(),
channel_versions=checkpoint["channel_versions"].copy(),
versions_seen=deepcopy(checkpoint["versions_seen"]),
channel_versions=defaultdict(int, checkpoint["channel_versions"]),
versions_seen=defaultdict(
_seen_dict,
{k: defaultdict(int, v) for k, v in checkpoint["versions_seen"].items()},
),
)
@@ -105,7 +108,7 @@ CheckpointThreadTs = ConfigurableFieldSpec(
class BaseCheckpointSaver(ABC):
at: CheckpointAt = CheckpointAt.END_OF_STEP
serde: SerializerProtocol
serde: SerializerProtocol = JsonPlusSerializer()
def __init__(
self,
-3
View File
@@ -1,5 +1,4 @@
import asyncio
import pickle
from collections import defaultdict
from typing import AsyncIterator, Iterator, Optional
@@ -15,8 +14,6 @@ from langgraph.checkpoint.base import (
class MemorySaver(BaseCheckpointSaver):
serde = pickle
storage: defaultdict[str, dict[str, Checkpoint]]
def __init__(
+14 -5
View File
@@ -16,9 +16,13 @@ from langgraph.checkpoint.base import (
)
class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
serde = pickle
# for backwards compat we continue to support loading pickled checkpoints
def is_pickled(value: bytes) -> bool:
print(value, type(value))
return value.startswith(b"\x80") and value.endswith(b".")
class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
conn: sqlite3.Connection
is_setup: bool
@@ -78,6 +82,11 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
self.conn.commit()
cur.close()
def _loads(self, value: bytes) -> Checkpoint:
if is_pickled(value):
return pickle.loads(value)
return self.serde.loads(value)
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
with self.cursor(transaction=False) as cur:
if config["configurable"].get("thread_ts"):
@@ -91,7 +100,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
if value := cur.fetchone():
return CheckpointTuple(
config,
self.serde.loads(value[0]),
self._loads(value[0]),
{
"configurable": {
"thread_id": config["configurable"]["thread_id"],
@@ -114,7 +123,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
"thread_ts": value[1],
}
},
self.serde.loads(value[3]),
self._loads(value[3]),
{
"configurable": {
"thread_id": value[0],
@@ -134,7 +143,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
for thread_id, thread_ts, parent_ts, value in cur:
yield CheckpointTuple(
{"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}},
self.serde.loads(value),
self._loads(value),
{
"configurable": {
"thread_id": thread_id,
+1 -1
View File
@@ -90,7 +90,7 @@ class JsonPlusSerializer(SerializerProtocol):
return LC_REVIVER(value)
def dumps(self, obj: Any) -> bytes:
return json.dumps(obj, default=self._default, sort_keys=True)
return json.dumps(obj, default=self._default, sort_keys=True).encode()
def loads(self, data: bytes) -> Any:
return json.loads(data, object_hook=self._reviver)
+1 -2
View File
@@ -5,7 +5,6 @@ from datetime import datetime, timezone
from enum import Enum
import dataclasses_json
import pytest
from langchain_core.pydantic_v1 import BaseModel as LcBaseModel
from langchain_core.runnables import RunnableMap
from pydantic import BaseModel
@@ -83,7 +82,7 @@ def test_serde_jsonplus() -> None:
assert (
dumped
== """{"a_bool": true, "a_float": 1.1, "a_none": null, "a_str": "foo", "an_int": 1, "my_dataclass": {"args": [], "id": ["tests", "test_jsonplus", "MyDataclass"], "kwargs": {"bar": 1, "foo": "foo"}, "lc": 2, "method": null, "type": "constructor"}, "my_enum": {"args": ["foo"], "id": ["tests", "test_jsonplus", "MyEnum"], "kwargs": {}, "lc": 2, "method": null, "type": "constructor"}, "my_funny_pydantic": {"args": [], "id": ["tests", "test_jsonplus", "MyFunnyPydantic"], "kwargs": {"bar": 1, "foo": "foo"}, "lc": 2, "method": null, "type": "constructor"}, "my_pydantic": {"args": [], "id": ["tests", "test_jsonplus", "MyPydantic"], "kwargs": {"bar": 1, "foo": "foo"}, "lc": 2, "method": null, "type": "constructor"}, "my_slotted_class": {"args": [], "id": ["tests", "test_jsonplus", "MyDataclassWSlots"], "kwargs": {"bar": 2, "foo": "bar"}, "lc": 2, "method": null, "type": "constructor"}, "person": {"args": [], "id": ["tests", "test_jsonplus", "Person"], "kwargs": {"name": "foo"}, "lc": 2, "method": null, "type": "constructor"}, "runnable_map": {"graph": {"edges": [], "nodes": [{"data": "Parallel<>Input", "id": 0, "type": "schema"}, {"data": "Parallel<>Output", "id": 1, "type": "schema"}]}, "id": ["langchain", "schema", "runnable", "RunnableParallel"], "kwargs": {"steps__": {}}, "lc": 1, "name": "RunnableParallel<>", "type": "constructor"}, "time": {"args": ["2024-04-19T23:04:57.051022+23:59"], "id": ["datetime", "datetime"], "kwargs": {}, "lc": 2, "method": "fromisoformat", "type": "constructor"}, "uid": {"args": ["00000000000000000000000000000001"], "id": ["uuid", "UUID"], "kwargs": {}, "lc": 2, "method": null, "type": "constructor"}}"""
== b"""{"a_bool": true, "a_float": 1.1, "a_none": null, "a_str": "foo", "an_int": 1, "my_dataclass": {"args": [], "id": ["tests", "test_jsonplus", "MyDataclass"], "kwargs": {"bar": 1, "foo": "foo"}, "lc": 2, "method": null, "type": "constructor"}, "my_enum": {"args": ["foo"], "id": ["tests", "test_jsonplus", "MyEnum"], "kwargs": {}, "lc": 2, "method": null, "type": "constructor"}, "my_funny_pydantic": {"args": [], "id": ["tests", "test_jsonplus", "MyFunnyPydantic"], "kwargs": {"bar": 1, "foo": "foo"}, "lc": 2, "method": null, "type": "constructor"}, "my_pydantic": {"args": [], "id": ["tests", "test_jsonplus", "MyPydantic"], "kwargs": {"bar": 1, "foo": "foo"}, "lc": 2, "method": null, "type": "constructor"}, "my_slotted_class": {"args": [], "id": ["tests", "test_jsonplus", "MyDataclassWSlots"], "kwargs": {"bar": 2, "foo": "bar"}, "lc": 2, "method": null, "type": "constructor"}, "person": {"args": [], "id": ["tests", "test_jsonplus", "Person"], "kwargs": {"name": "foo"}, "lc": 2, "method": null, "type": "constructor"}, "runnable_map": {"graph": {"edges": [], "nodes": [{"data": "Parallel<>Input", "id": 0, "type": "schema"}, {"data": "Parallel<>Output", "id": 1, "type": "schema"}]}, "id": ["langchain", "schema", "runnable", "RunnableParallel"], "kwargs": {"steps__": {}}, "lc": 1, "name": "RunnableParallel<>", "type": "constructor"}, "time": {"args": ["2024-04-19T23:04:57.051022+23:59"], "id": ["datetime", "datetime"], "kwargs": {}, "lc": 2, "method": "fromisoformat", "type": "constructor"}, "uid": {"args": ["00000000000000000000000000000001"], "id": ["uuid", "UUID"], "kwargs": {}, "lc": 2, "method": null, "type": "constructor"}}"""
)
assert serde.loads(dumped) == to_serialize