merge with ormsgpack

This commit is contained in:
William Fu-Hinthorn
2025-03-24 16:21:59 -07:00
10 changed files with 416 additions and 521 deletions
@@ -20,7 +20,7 @@ from ipaddress import (
from typing import Any, Callable, Optional, Union, cast
from uuid import UUID
import msgpack # type: ignore[import-untyped]
import ormsgpack
from langchain_core.load.load import Reviver
from langchain_core.load.serializable import Serializable
from zoneinfo import ZoneInfo
@@ -33,7 +33,9 @@ LC_REVIVER = Reviver()
class JsonPlusSerializer(SerializerProtocol):
def __init__(self, *, __unpack_ext_hook__: Optional[Callable] = None) -> None:
def __init__(
self, *, __unpack_ext_hook__: Optional[Callable[[int, bytes], Any]] = None
) -> None:
self._unpack_ext_hook = (
__unpack_ext_hook__
if __unpack_ext_hook__ is not None
@@ -199,8 +201,10 @@ class JsonPlusSerializer(SerializerProtocol):
else:
try:
return "msgpack", _msgpack_enc(obj)
except UnicodeEncodeError:
return "json", self.dumps(obj)
except ormsgpack.MsgpackEncodeError as exc:
if "valid UTF-8" in str(exc):
return "json", self.dumps(obj)
raise exc
def loads(self, data: bytes) -> Any:
return json.loads(data, object_hook=self._reviver)
@@ -214,8 +218,8 @@ class JsonPlusSerializer(SerializerProtocol):
elif type_ == "json":
return self.loads(data_)
elif type_ == "msgpack":
return msgpack.unpackb(
data_, ext_hook=self._unpack_ext_hook, strict_map_key=False
return ormsgpack.unpackb(
data_, ext_hook=self._unpack_ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
)
else:
raise NotImplementedError(f"Unknown serialization type: {type_}")
@@ -231,9 +235,9 @@ EXT_PYDANTIC_V1 = 4
EXT_PYDANTIC_V2 = 5
def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
def _msgpack_default(obj: Any) -> Union[str, ormsgpack.Ext]:
if hasattr(obj, "model_dump") and callable(obj.model_dump): # pydantic v2
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_PYDANTIC_V2,
_msgpack_enc(
(
@@ -245,7 +249,7 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
),
)
elif hasattr(obj, "get_secret_value") and callable(obj.get_secret_value):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_SINGLE_ARG,
_msgpack_enc(
(
@@ -256,7 +260,7 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
),
)
elif hasattr(obj, "dict") and callable(obj.dict): # pydantic v1
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_PYDANTIC_V1,
_msgpack_enc(
(
@@ -267,7 +271,7 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
),
)
elif hasattr(obj, "_asdict") and callable(obj._asdict): # namedtuple
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_KW_ARGS,
_msgpack_enc(
(
@@ -278,56 +282,63 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
),
)
elif isinstance(obj, pathlib.Path):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_POS_ARGS,
_msgpack_enc(
(obj.__class__.__module__, obj.__class__.__name__, obj.parts),
),
)
elif isinstance(obj, re.Pattern):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_POS_ARGS,
_msgpack_enc(
("re", "compile", (obj.pattern, obj.flags)),
),
)
elif isinstance(obj, UUID):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_SINGLE_ARG,
_msgpack_enc(
(obj.__class__.__module__, obj.__class__.__name__, obj.hex),
),
)
elif isinstance(obj, bytearray):
return ormsgpack.Ext(
EXT_CONSTRUCTOR_SINGLE_ARG,
_msgpack_enc(
(obj.__class__.__module__, obj.__class__.__name__, bytes(obj)),
),
)
elif isinstance(obj, decimal.Decimal):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_SINGLE_ARG,
_msgpack_enc(
(obj.__class__.__module__, obj.__class__.__name__, str(obj)),
),
)
elif isinstance(obj, (set, frozenset, deque)):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_SINGLE_ARG,
_msgpack_enc(
(obj.__class__.__module__, obj.__class__.__name__, tuple(obj)),
),
)
elif isinstance(obj, (IPv4Address, IPv4Interface, IPv4Network)):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_SINGLE_ARG,
_msgpack_enc(
(obj.__class__.__module__, obj.__class__.__name__, str(obj)),
),
)
elif isinstance(obj, (IPv6Address, IPv6Interface, IPv6Network)):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_SINGLE_ARG,
_msgpack_enc(
(obj.__class__.__module__, obj.__class__.__name__, str(obj)),
),
)
elif isinstance(obj, datetime):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_METHOD_SINGLE_ARG,
_msgpack_enc(
(
@@ -339,7 +350,7 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
),
)
elif isinstance(obj, timedelta):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_POS_ARGS,
_msgpack_enc(
(
@@ -350,7 +361,7 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
),
)
elif isinstance(obj, date):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_POS_ARGS,
_msgpack_enc(
(
@@ -361,7 +372,7 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
),
)
elif isinstance(obj, time):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_KW_ARGS,
_msgpack_enc(
(
@@ -379,7 +390,7 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
),
)
elif isinstance(obj, timezone):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_POS_ARGS,
_msgpack_enc(
(
@@ -390,21 +401,21 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
),
)
elif isinstance(obj, ZoneInfo):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_SINGLE_ARG,
_msgpack_enc(
(obj.__class__.__module__, obj.__class__.__name__, obj.key),
),
)
elif isinstance(obj, Enum):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_SINGLE_ARG,
_msgpack_enc(
(obj.__class__.__module__, obj.__class__.__name__, obj.value),
),
)
elif isinstance(obj, SendProtocol):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_POS_ARGS,
_msgpack_enc(
(obj.__class__.__module__, obj.__class__.__name__, (obj.node, obj.arg)),
@@ -412,7 +423,7 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
)
elif dataclasses.is_dataclass(obj):
# doesn't use dataclasses.asdict to avoid deepcopy and recursion
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_KW_ARGS,
_msgpack_enc(
(
@@ -426,7 +437,7 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
),
)
elif isinstance(obj, Item):
return msgpack.ExtType(
return ormsgpack.Ext(
EXT_CONSTRUCTOR_KW_ARGS,
_msgpack_enc(
(
@@ -436,7 +447,6 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
),
),
)
elif isinstance(obj, BaseException):
return repr(obj)
else:
@@ -446,8 +456,8 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]:
def _msgpack_ext_hook(code: int, data: bytes) -> Any:
if code == EXT_CONSTRUCTOR_SINGLE_ARG:
try:
tup = msgpack.unpackb(
data, ext_hook=_msgpack_ext_hook, strict_map_key=False
tup = ormsgpack.unpackb(
data, ext_hook=_msgpack_ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
)
# module, name, arg
return getattr(importlib.import_module(tup[0]), tup[1])(tup[2])
@@ -455,8 +465,8 @@ def _msgpack_ext_hook(code: int, data: bytes) -> Any:
return
elif code == EXT_CONSTRUCTOR_POS_ARGS:
try:
tup = msgpack.unpackb(
data, ext_hook=_msgpack_ext_hook, strict_map_key=False
tup = ormsgpack.unpackb(
data, ext_hook=_msgpack_ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
)
# module, name, args
return getattr(importlib.import_module(tup[0]), tup[1])(*tup[2])
@@ -464,8 +474,8 @@ def _msgpack_ext_hook(code: int, data: bytes) -> Any:
return
elif code == EXT_CONSTRUCTOR_KW_ARGS:
try:
tup = msgpack.unpackb(
data, ext_hook=_msgpack_ext_hook, strict_map_key=False
tup = ormsgpack.unpackb(
data, ext_hook=_msgpack_ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
)
# module, name, args
return getattr(importlib.import_module(tup[0]), tup[1])(**tup[2])
@@ -473,8 +483,8 @@ def _msgpack_ext_hook(code: int, data: bytes) -> Any:
return
elif code == EXT_METHOD_SINGLE_ARG:
try:
tup = msgpack.unpackb(
data, ext_hook=_msgpack_ext_hook, strict_map_key=False
tup = ormsgpack.unpackb(
data, ext_hook=_msgpack_ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
)
# module, name, arg, method
return getattr(getattr(importlib.import_module(tup[0]), tup[1]), tup[3])(
@@ -484,8 +494,8 @@ def _msgpack_ext_hook(code: int, data: bytes) -> Any:
return
elif code == EXT_PYDANTIC_V1:
try:
tup = msgpack.unpackb(
data, ext_hook=_msgpack_ext_hook, strict_map_key=False
tup = ormsgpack.unpackb(
data, ext_hook=_msgpack_ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
)
# module, name, kwargs
cls = getattr(importlib.import_module(tup[0]), tup[1])
@@ -502,8 +512,8 @@ def _msgpack_ext_hook(code: int, data: bytes) -> Any:
return
elif code == EXT_PYDANTIC_V2:
try:
tup = msgpack.unpackb(
data, ext_hook=_msgpack_ext_hook, strict_map_key=False
tup = ormsgpack.unpackb(
data, ext_hook=_msgpack_ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
)
# module, name, kwargs, method
cls = getattr(importlib.import_module(tup[0]), tup[1])
@@ -523,8 +533,10 @@ def _msgpack_ext_hook(code: int, data: bytes) -> Any:
def _msgpack_ext_hook_to_json(code: int, data: bytes) -> Any:
if code == EXT_CONSTRUCTOR_SINGLE_ARG:
try:
tup = msgpack.unpackb(
data, ext_hook=_msgpack_ext_hook_to_json, strict_map_key=False
tup = ormsgpack.unpackb(
data,
ext_hook=_msgpack_ext_hook_to_json,
option=ormsgpack.OPT_NON_STR_KEYS,
)
if tup[0] == "uuid" and tup[1] == "UUID":
hex_ = tup[2]
@@ -537,8 +549,10 @@ def _msgpack_ext_hook_to_json(code: int, data: bytes) -> Any:
return
elif code == EXT_CONSTRUCTOR_POS_ARGS:
try:
tup = msgpack.unpackb(
data, ext_hook=_msgpack_ext_hook_to_json, strict_map_key=False
tup = ormsgpack.unpackb(
data,
ext_hook=_msgpack_ext_hook_to_json,
option=ormsgpack.OPT_NON_STR_KEYS,
)
# module, name, args
return tup[2]
@@ -546,8 +560,10 @@ def _msgpack_ext_hook_to_json(code: int, data: bytes) -> Any:
return
elif code == EXT_CONSTRUCTOR_KW_ARGS:
try:
tup = msgpack.unpackb(
data, ext_hook=_msgpack_ext_hook_to_json, strict_map_key=False
tup = ormsgpack.unpackb(
data,
ext_hook=_msgpack_ext_hook_to_json,
option=ormsgpack.OPT_NON_STR_KEYS,
)
# module, name, args
return tup[2]
@@ -555,8 +571,10 @@ def _msgpack_ext_hook_to_json(code: int, data: bytes) -> Any:
return
elif code == EXT_METHOD_SINGLE_ARG:
try:
tup = msgpack.unpackb(
data, ext_hook=_msgpack_ext_hook_to_json, strict_map_key=False
tup = ormsgpack.unpackb(
data,
ext_hook=_msgpack_ext_hook_to_json,
option=ormsgpack.OPT_NON_STR_KEYS,
)
# module, name, arg, method
return tup[2]
@@ -564,8 +582,10 @@ def _msgpack_ext_hook_to_json(code: int, data: bytes) -> Any:
return
elif code == EXT_PYDANTIC_V1:
try:
tup = msgpack.unpackb(
data, ext_hook=_msgpack_ext_hook_to_json, strict_map_key=False
tup = ormsgpack.unpackb(
data,
ext_hook=_msgpack_ext_hook_to_json,
option=ormsgpack.OPT_NON_STR_KEYS,
)
# module, name, kwargs
return tup[2]
@@ -575,8 +595,10 @@ def _msgpack_ext_hook_to_json(code: int, data: bytes) -> Any:
return
elif code == EXT_PYDANTIC_V2:
try:
tup = msgpack.unpackb(
data, ext_hook=_msgpack_ext_hook_to_json, strict_map_key=False
tup = ormsgpack.unpackb(
data,
ext_hook=_msgpack_ext_hook_to_json,
option=ormsgpack.OPT_NON_STR_KEYS,
)
# module, name, kwargs, method
return tup[2]
@@ -584,5 +606,14 @@ def _msgpack_ext_hook_to_json(code: int, data: bytes) -> Any:
return
_option = (
ormsgpack.OPT_NON_STR_KEYS
| ormsgpack.OPT_PASSTHROUGH_DATACLASS
| ormsgpack.OPT_PASSTHROUGH_DATETIME
| ormsgpack.OPT_PASSTHROUGH_ENUM
| ormsgpack.OPT_PASSTHROUGH_UUID
)
def _msgpack_enc(data: Any) -> bytes:
return msgpack.packb(data, default=_msgpack_default)
return ormsgpack.packb(data, default=_msgpack_default, option=_option)