From a08eaf1a77756a3a89cd9b7c1f9ef01496e90fe6 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 24 Apr 2024 16:18:02 -0700 Subject: [PATCH] Fix --- langgraph/checkpoint/aiosqlite.py | 20 ++++++-------------- langgraph/checkpoint/sqlite.py | 22 +++++++++++----------- 2 files changed, 17 insertions(+), 25 deletions(-) diff --git a/langgraph/checkpoint/aiosqlite.py b/langgraph/checkpoint/aiosqlite.py index 8a304b1a9..8d9d3fbc0 100644 --- a/langgraph/checkpoint/aiosqlite.py +++ b/langgraph/checkpoint/aiosqlite.py @@ -1,4 +1,3 @@ -import pickle from contextlib import AbstractAsyncContextManager from types import TracebackType from typing import AsyncIterator, Optional @@ -14,14 +13,12 @@ from langgraph.checkpoint.base import ( CheckpointTuple, SerializerProtocol, ) - - -# 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".") +from langgraph.checkpoint.sqlite import JsonPlusSerializerCompat class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): + serde = JsonPlusSerializerCompat() + conn: aiosqlite.Connection is_setup: bool @@ -72,11 +69,6 @@ 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"): @@ -90,7 +82,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): if value := await cursor.fetchone(): return CheckpointTuple( config, - self._loads(value[0]), + self.serde.loads(value[0]), { "configurable": { "thread_id": config["configurable"]["thread_id"], @@ -113,7 +105,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): "thread_ts": value[1], } }, - self._loads(value[3]), + self.serde.loads(value[3]), { "configurable": { "thread_id": value[0], @@ -133,7 +125,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._loads(value), + self.serde.loads(value), {"configurable": {"thread_id": thread_id, "thread_ts": parent_ts}} if parent_ts else None, diff --git a/langgraph/checkpoint/sqlite.py b/langgraph/checkpoint/sqlite.py index 96e2bafcb..118f328fc 100644 --- a/langgraph/checkpoint/sqlite.py +++ b/langgraph/checkpoint/sqlite.py @@ -14,15 +14,20 @@ from langgraph.checkpoint.base import ( CheckpointTuple, SerializerProtocol, ) +from langgraph.serde.jsonplus import JsonPlusSerializer # 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 JsonPlusSerializerCompat(JsonPlusSerializer): + def loads(self, data: bytes) -> pickle.Any: + if data.startswith(b"\x80") and data.endswith(b"."): + return pickle.loads(data) + return super().loads(data) class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): + serde = JsonPlusSerializerCompat + conn: sqlite3.Connection is_setup: bool @@ -82,11 +87,6 @@ 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"): @@ -100,7 +100,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): if value := cur.fetchone(): return CheckpointTuple( config, - self._loads(value[0]), + self.serde.loads(value[0]), { "configurable": { "thread_id": config["configurable"]["thread_id"], @@ -123,7 +123,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): "thread_ts": value[1], } }, - self._loads(value[3]), + self.serde.loads(value[3]), { "configurable": { "thread_id": value[0], @@ -143,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._loads(value), + self.serde.loads(value), { "configurable": { "thread_id": thread_id,