mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 21:45:08 +02:00
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:
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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__(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user