This commit is contained in:
Nuno Campos
2024-04-24 15:40:48 -07:00
parent 6125ea68f5
commit 3ad0204a90
2 changed files with 96 additions and 3 deletions
+22 -3
View File
@@ -1,6 +1,8 @@
import dataclasses
import importlib
import json
from datetime import datetime
from datetime import datetime, timedelta, timezone
from enum import Enum
from typing import Any, Optional
from uuid import UUID
@@ -43,14 +45,31 @@ class JsonPlusSerializer(SerializerProtocol):
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(), obj.tzinfo]
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:
print(value)
if (
value.get("lc", None) == 2
and value.get("type", None) == "constructor"
@@ -75,4 +94,4 @@ class JsonPlusSerializer(SerializerProtocol):
return json.dumps(obj, default=self._default, sort_keys=True)
def loads(self, data: bytes) -> Any:
return json.loads(data)
return json.loads(data, object_hook=self._reviver)
+74
View File
@@ -0,0 +1,74 @@
import dataclasses
import uuid
from datetime import datetime, timezone
from enum import Enum
import dataclasses_json
from langchain_core.pydantic_v1 import BaseModel as LcBaseModel
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
@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.uuid4()
current_time = datetime.now(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,
}
serde = JsonPlusSerializer()
assert serde.loads(serde.dumps(to_serialize)) == to_serialize