Merge pull request #317 from langchain-ai/nc/17apr/serde

Introduce json-based checkpoint serialization
This commit is contained in:
Nuno Campos
2024-04-24 16:26:56 -07:00
committed by GitHub
9 changed files with 1041 additions and 852 deletions
+2 -2
View File
@@ -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
+8 -13
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,
@@ -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,
-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__(
+11 -2
View File
@@ -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
View File
+17
View File
@@ -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:
...
+96
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
+90
View File
@@ -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