mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-20 06:35:46 +02:00
Merge pull request #317 from langchain-ai/nc/17apr/serde
Introduce json-based checkpoint serialization
This commit is contained in:
@@ -1,4 +1,3 @@
|
||||
import pickle
|
||||
from contextlib import AbstractAsyncContextManager
|
||||
from types import TracebackType
|
||||
from typing import AsyncIterator, Optional
|
||||
@@ -14,10 +13,11 @@ from langgraph.checkpoint.base import (
|
||||
CheckpointTuple,
|
||||
SerializerProtocol,
|
||||
)
|
||||
from langgraph.checkpoint.sqlite import JsonPlusSerializerCompat
|
||||
|
||||
|
||||
class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
|
||||
serde = pickle
|
||||
serde = JsonPlusSerializerCompat()
|
||||
|
||||
conn: aiosqlite.Connection
|
||||
|
||||
|
||||
@@ -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,
|
||||
@@ -8,12 +7,13 @@ from typing import (
|
||||
Iterator,
|
||||
NamedTuple,
|
||||
Optional,
|
||||
Protocol,
|
||||
TypedDict,
|
||||
)
|
||||
|
||||
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()},
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -102,18 +105,10 @@ CheckpointThreadTs = ConfigurableFieldSpec(
|
||||
)
|
||||
|
||||
|
||||
class SerializerProtocol(Protocol):
|
||||
def dumps(self, obj: Any) -> bytes:
|
||||
...
|
||||
|
||||
def loads(self, data: bytes) -> Any:
|
||||
...
|
||||
|
||||
|
||||
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__(
|
||||
|
||||
@@ -2,7 +2,7 @@ import pickle
|
||||
import sqlite3
|
||||
from contextlib import AbstractContextManager, contextmanager
|
||||
from types import TracebackType
|
||||
from typing import Iterator, Optional
|
||||
from typing import Any, Iterator, Optional
|
||||
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from typing_extensions import Self
|
||||
@@ -14,10 +14,19 @@ from langgraph.checkpoint.base import (
|
||||
CheckpointTuple,
|
||||
SerializerProtocol,
|
||||
)
|
||||
from langgraph.serde.jsonplus import JsonPlusSerializer
|
||||
|
||||
|
||||
# for backwards compat we continue to support loading pickled checkpoints
|
||||
class JsonPlusSerializerCompat(JsonPlusSerializer):
|
||||
def loads(self, data: bytes) -> Any:
|
||||
if data.startswith(b"\x80") and data.endswith(b"."):
|
||||
return pickle.loads(data)
|
||||
return super().loads(data)
|
||||
|
||||
|
||||
class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
|
||||
serde = pickle
|
||||
serde = JsonPlusSerializerCompat()
|
||||
|
||||
conn: sqlite3.Connection
|
||||
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
from typing import Any, Protocol
|
||||
|
||||
|
||||
class SerializerProtocol(Protocol):
|
||||
"""Protocol for serialization and deserialization of objects.
|
||||
|
||||
- `dumps`: Serialize an object to bytes.
|
||||
- `loads`: Deserialize an object from bytes.
|
||||
|
||||
Valid implementations include the `pickle`, `json` and `orjson` modules.
|
||||
"""
|
||||
|
||||
def dumps(self, obj: Any) -> bytes:
|
||||
...
|
||||
|
||||
def loads(self, data: bytes) -> Any:
|
||||
...
|
||||
@@ -0,0 +1,96 @@
|
||||
import dataclasses
|
||||
import importlib
|
||||
import json
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from enum import Enum
|
||||
from typing import Any, Optional
|
||||
from uuid import UUID
|
||||
|
||||
from langchain_core.load.load import Reviver
|
||||
from langchain_core.load.serializable import Serializable
|
||||
from langchain_core.pydantic_v1 import BaseModel as LcBaseModel
|
||||
from pydantic import BaseModel
|
||||
|
||||
from langgraph.serde.base import SerializerProtocol
|
||||
|
||||
LC_REVIVER = Reviver()
|
||||
|
||||
|
||||
class JsonPlusSerializer(SerializerProtocol):
|
||||
def _encode_constructor_args(
|
||||
self,
|
||||
constructor: type[Any],
|
||||
*,
|
||||
method: Optional[str] = None,
|
||||
args: Optional[list[Any]] = None,
|
||||
kwargs: Optional[dict[str, Any]] = None,
|
||||
):
|
||||
return {
|
||||
"lc": 2,
|
||||
"type": "constructor",
|
||||
"id": [*constructor.__module__.split("."), constructor.__name__],
|
||||
"method": method,
|
||||
"args": args if args is not None else [],
|
||||
"kwargs": kwargs if kwargs is not None else {},
|
||||
}
|
||||
|
||||
def _default(self, obj):
|
||||
if isinstance(obj, Serializable):
|
||||
return obj.to_json()
|
||||
elif isinstance(obj, (BaseModel, LcBaseModel)):
|
||||
return self._encode_constructor_args(obj.__class__, kwargs=obj.dict())
|
||||
elif isinstance(obj, UUID):
|
||||
return self._encode_constructor_args(UUID, args=[obj.hex])
|
||||
elif isinstance(obj, (set, frozenset)):
|
||||
return self._encode_constructor_args(type(obj), args=[list(obj)])
|
||||
elif isinstance(obj, datetime):
|
||||
return self._encode_constructor_args(
|
||||
datetime, method="fromisoformat", args=[obj.isoformat()]
|
||||
)
|
||||
elif isinstance(obj, timezone):
|
||||
return self._encode_constructor_args(timezone, args=obj.__getinitargs__())
|
||||
elif isinstance(obj, timedelta):
|
||||
return self._encode_constructor_args(
|
||||
timedelta, args=[obj.days, obj.seconds, obj.microseconds]
|
||||
)
|
||||
elif dataclasses.is_dataclass(obj):
|
||||
return self._encode_constructor_args(
|
||||
obj.__class__,
|
||||
kwargs={
|
||||
field.name: getattr(obj, field.name)
|
||||
for field in dataclasses.fields(obj)
|
||||
},
|
||||
)
|
||||
elif isinstance(obj, Enum):
|
||||
return self._encode_constructor_args(obj.__class__, args=[obj.value])
|
||||
else:
|
||||
raise TypeError(
|
||||
f"Object of type {obj.__class__.__name__} is not JSON serializable"
|
||||
)
|
||||
|
||||
def _reviver(self, value: dict[str, Any]) -> Any:
|
||||
if (
|
||||
value.get("lc", None) == 2
|
||||
and value.get("type", None) == "constructor"
|
||||
and value.get("id", None) is not None
|
||||
):
|
||||
# Get module and class name
|
||||
[*module, name] = value["id"]
|
||||
# Import module
|
||||
mod = importlib.import_module(".".join(module))
|
||||
# Import class
|
||||
cls = getattr(mod, name)
|
||||
# Instantiate class
|
||||
if value["method"] is not None:
|
||||
method = getattr(cls, value["method"])
|
||||
return method(*value["args"], **value["kwargs"])
|
||||
else:
|
||||
return cls(*value["args"], **value["kwargs"])
|
||||
|
||||
return LC_REVIVER(value)
|
||||
|
||||
def dumps(self, obj: Any) -> bytes:
|
||||
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)
|
||||
Generated
+817
-832
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,90 @@
|
||||
import dataclasses
|
||||
import sys
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from enum import Enum
|
||||
|
||||
import dataclasses_json
|
||||
from langchain_core.pydantic_v1 import BaseModel as LcBaseModel
|
||||
from langchain_core.runnables import RunnableMap
|
||||
from pydantic import BaseModel
|
||||
|
||||
from langgraph.serde.jsonplus import JsonPlusSerializer
|
||||
|
||||
|
||||
class MyPydantic(BaseModel):
|
||||
foo: str
|
||||
bar: int
|
||||
|
||||
|
||||
class MyFunnyPydantic(LcBaseModel):
|
||||
foo: str
|
||||
bar: int
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class MyDataclass:
|
||||
foo: str
|
||||
bar: int
|
||||
|
||||
def something(self) -> None:
|
||||
pass
|
||||
|
||||
|
||||
if sys.version_info < (3, 10):
|
||||
|
||||
class MyDataclassWSlots(MyDataclass):
|
||||
pass
|
||||
else:
|
||||
|
||||
@dataclasses.dataclass(slots=True)
|
||||
class MyDataclassWSlots:
|
||||
foo: str
|
||||
bar: int
|
||||
|
||||
def something(self) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class MyEnum(Enum):
|
||||
FOO = "foo"
|
||||
BAR = "bar"
|
||||
|
||||
|
||||
@dataclasses_json.dataclass_json
|
||||
@dataclasses.dataclass
|
||||
class Person:
|
||||
name: str
|
||||
|
||||
|
||||
def test_serde_jsonplus() -> None:
|
||||
uid = uuid.UUID(int=1)
|
||||
current_time = datetime(2024, 4, 19, 23, 4, 57, 51022, timezone.max)
|
||||
|
||||
to_serialize = {
|
||||
"uid": uid,
|
||||
"time": current_time,
|
||||
"my_slotted_class": MyDataclassWSlots("bar", 2),
|
||||
"my_dataclass": MyDataclass("foo", 1),
|
||||
"my_enum": MyEnum.FOO,
|
||||
"my_pydantic": MyPydantic(foo="foo", bar=1),
|
||||
"my_funny_pydantic": MyFunnyPydantic(foo="foo", bar=1),
|
||||
"person": Person(name="foo"),
|
||||
"a_bool": True,
|
||||
"a_none": None,
|
||||
"a_str": "foo",
|
||||
"an_int": 1,
|
||||
"a_float": 1.1,
|
||||
"runnable_map": RunnableMap({}),
|
||||
}
|
||||
|
||||
serde = JsonPlusSerializer()
|
||||
|
||||
dumped = serde.dumps(to_serialize)
|
||||
|
||||
assert (
|
||||
dumped
|
||||
== 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