mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-25 09:02:25 +02:00
Add test
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user