diff --git a/langgraph/serde/jsonplus.py b/langgraph/serde/jsonplus.py index 3df63d642..3fb3f69be 100644 --- a/langgraph/serde/jsonplus.py +++ b/langgraph/serde/jsonplus.py @@ -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) diff --git a/tests/test_jsonplus.py b/tests/test_jsonplus.py new file mode 100644 index 000000000..3bbe41429 --- /dev/null +++ b/tests/test_jsonplus.py @@ -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