diff --git a/langgraph/checkpoint/aiosqlite.py b/langgraph/checkpoint/aiosqlite.py index b508cefb9..8a304b1a9 100644 --- a/langgraph/checkpoint/aiosqlite.py +++ b/langgraph/checkpoint/aiosqlite.py @@ -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, diff --git a/langgraph/checkpoint/base.py b/langgraph/checkpoint/base.py index 28b485845..b8c0a9ec7 100644 --- a/langgraph/checkpoint/base.py +++ b/langgraph/checkpoint/base.py @@ -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, diff --git a/langgraph/checkpoint/memory.py b/langgraph/checkpoint/memory.py index c3348ed29..3f3ba5e63 100644 --- a/langgraph/checkpoint/memory.py +++ b/langgraph/checkpoint/memory.py @@ -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__( diff --git a/langgraph/checkpoint/sqlite.py b/langgraph/checkpoint/sqlite.py index 829874646..96e2bafcb 100644 --- a/langgraph/checkpoint/sqlite.py +++ b/langgraph/checkpoint/sqlite.py @@ -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, diff --git a/langgraph/serde/jsonplus.py b/langgraph/serde/jsonplus.py index 1428bfb3a..9fe4b9d3c 100644 --- a/langgraph/serde/jsonplus.py +++ b/langgraph/serde/jsonplus.py @@ -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) diff --git a/tests/test_jsonplus.py b/tests/test_jsonplus.py index 4dacfafb4..df4956a15 100644 --- a/tests/test_jsonplus.py +++ b/tests/test_jsonplus.py @@ -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