Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
jsonplus.py904 linesDownload Raw Back to serde
1from __future__ import annotations2 3import copy4import dataclasses5import decimal6import importlib7import json8import logging9import pathlib10import pickle11import re12import sys13from collections import deque14from collections.abc import Callable, Iterable, Sequence15from datetime import date, datetime, time, timedelta, timezone16from enum import Enum17from inspect import isclass18from ipaddress import (19    IPv4Address,20    IPv4Interface,21    IPv4Network,22    IPv6Address,23    IPv6Interface,24    IPv6Network,25)26from typing import TYPE_CHECKING, Any, Literal, cast27from uuid import UUID28from zoneinfo import ZoneInfo29 30import ormsgpack31from langchain_core.load.load import Reviver32 33from langgraph.checkpoint.serde import _msgpack as _lg_msgpack34from langgraph.checkpoint.serde.base import SerializerProtocol35from langgraph.checkpoint.serde.event_hooks import emit_serde_event36from langgraph.checkpoint.serde.types import (37    SendProtocol,38    _DeltaSnapshot,39)40from langgraph.store.base import Item41 42if TYPE_CHECKING:43    from langgraph.checkpoint.serde._msgpack import (44        AllowedMsgpackModules,45    )46 47LC_REVIVER = Reviver(allowed_objects="core")48EMPTY_BYTES = b""49logger = logging.getLogger(__name__)50 51# Dedup log warnings across process lifetime; cap bounds state if types are52# dynamically generated (also acts as a circuit breaker on warning volume).53# Dedup is best-effort: racing threads may each emit once for the same key,54# and warnings are silently dropped once _MAX_WARNED_TYPES is reached.55_MAX_WARNED_TYPES = 100056_warned_unregistered_types: set[tuple[str, str]] = set()57_warned_blocked_types: set[tuple[str, str]] = set()58 59 60def _is_safe_json_type(id_list: list[str]) -> bool:61    """Return True if an lc=2 id refers to a type in SAFE_MSGPACK_TYPES.62 63    Safe types bypass the ``allowed_json_modules`` gate so that old "json" format64    checkpoints (written before the msgpack migration) can be resumed without65    requiring users to configure an explicit allowlist.66    """67    if len(id_list) < 2:68        return False69    module_name = ".".join(id_list[:-1])70    return (module_name, id_list[-1]) in _lg_msgpack.SAFE_MSGPACK_TYPES71 72 73def _warn_once(74    seen: set[tuple[str, str]], key: tuple[str, str], msg: str, *args: object75) -> None:76    if key in seen or len(seen) >= _MAX_WARNED_TYPES:77        return78    seen.add(key)79    logger.warning(msg, *args)80 81 82class JsonPlusSerializer(SerializerProtocol):83    """Serializer that uses ormsgpack, with optional fallbacks.84 85    !!! warning86 87        Security note: This serializer is intended for use within the `BaseCheckpointSaver`88        class and called within the Pregel loop. It should not be used on untrusted89        python objects. If an attacker can write directly to your checkpoint database,90        they may be able to trigger code execution when data is deserialized.91 92        Set the environment variable ``LANGGRAPH_STRICT_MSGPACK=true`` to restrict93        deserialization to a built-in allowlist of safe types.  You can also pass94        an explicit ``allowed_msgpack_modules`` to the constructor.95    """96 97    def __init__(98        self,99        *,100        pickle_fallback: bool = False,101        allowed_json_modules: Iterable[tuple[str, ...]] | Literal[True] | None = None,102        allowed_msgpack_modules: (103            AllowedMsgpackModules | Literal[True] | None104        ) = _lg_msgpack._SENTINEL,105        __unpack_ext_hook__: Callable[[int, bytes], Any] | None = None,106    ) -> None:107        if allowed_msgpack_modules is _lg_msgpack._SENTINEL:108            if _lg_msgpack.STRICT_MSGPACK_ENABLED:109                # Strict: only SAFE_MSGPACK_TYPES are allowed.110                allowed_msgpack_modules = None111            else:112                # Permissive (default): all types allowed with a warning.113                # Set LANGGRAPH_STRICT_MSGPACK=true to lock this down.114                allowed_msgpack_modules = True115        self.pickle_fallback = pickle_fallback116        self._allowed_json_modules: set[tuple[str, ...]] | Literal[True] | None = (117            _normalize_allowlist(allowed_json_modules)118        )119        self._allowed_msgpack_modules = _normalize_allowlist(allowed_msgpack_modules)120 121        self._custom_unpack_ext_hook = __unpack_ext_hook__ is not None122        self._unpack_ext_hook = (123            __unpack_ext_hook__124            if __unpack_ext_hook__ is not None125            else _create_msgpack_ext_hook(self._allowed_msgpack_modules)126        )127 128    def with_msgpack_allowlist(129        self, extra_allowlist: Iterable[tuple[str, ...] | type]130    ) -> JsonPlusSerializer:131        """Return a new serializer with a merged msgpack allowlist."""132        base_allowlist = self._allowed_msgpack_modules133        if base_allowlist is True or base_allowlist is False:134            return self135        elif base_allowlist:136            base_allowlist = set(base_allowlist)137        else:138            base_allowlist = set()139        extra = _normalize_module_keys(tuple(extra_allowlist))140        merged = base_allowlist | extra141        if merged == base_allowlist:142            return self143        allowed_msgpack_modules: AllowedMsgpackModules | Literal[True] | None144        if merged:145            allowed_msgpack_modules = tuple(merged)146        elif isinstance(self._allowed_msgpack_modules, set):147            allowed_msgpack_modules = tuple(self._allowed_msgpack_modules)148        else:149            allowed_msgpack_modules = self._allowed_msgpack_modules150 151        clone = copy.copy(self)152        clone._allowed_json_modules = _normalize_allowlist(self._allowed_json_modules)153        clone._allowed_msgpack_modules = _normalize_allowlist(allowed_msgpack_modules)154        if not clone._custom_unpack_ext_hook:155            clone._unpack_ext_hook = _create_msgpack_ext_hook(156                clone._allowed_msgpack_modules157            )158        return clone159 160    def _encode_constructor_args(161        self,162        constructor: Callable | type[Any],163        *,164        method: None | str | Sequence[None | str] = None,165        args: Sequence[Any] | None = None,166        kwargs: dict[str, Any] | None = None,167    ) -> dict[str, Any]:168        out = {169            "lc": 2,170            "type": "constructor",171            "id": (*constructor.__module__.split("."), constructor.__name__),172        }173        if method is not None:174            out["method"] = method175        if args is not None:176            out["args"] = args177        if kwargs is not None:178            out["kwargs"] = kwargs179        return out180 181    def _reviver(self, value: dict[str, Any]) -> Any:182        if (183            value.get("lc", None) == 2184            and value.get("type", None) == "constructor"185            and value.get("id", None) is not None186        ):187            id_list = value["id"]188            is_safe = _is_safe_json_type(id_list)189            if self._allowed_json_modules or is_safe:190                try:191                    return self._revive_lc2(value)192                except InvalidModuleError as e:193                    if not is_safe:194                        logger.warning(195                            "Object %s is not in the deserialization allowlist.\n%s",196                            value["id"],197                            e.message,198                        )199 200        return LC_REVIVER(value)201 202    def _revive_lc2(self, value: dict[str, Any]) -> Any:203        self._check_allowed_json_modules(value)204 205        [*module, name] = value["id"]206        try:207            mod = importlib.import_module(".".join(module))208            cls = getattr(mod, name)209            method = value.get("method")210            if isinstance(method, str):211                methods = [getattr(cls, method)]212            elif isinstance(method, list):213                methods = [cls if m is None else getattr(cls, m) for m in method]214            else:215                methods = [cls]216            args = value.get("args")217            kwargs = value.get("kwargs")218            for method in methods:219                try:220                    if isclass(method) and issubclass(method, BaseException):221                        return None222                    if args and kwargs:223                        return method(*args, **kwargs)224                    elif args:225                        return method(*args)226                    elif kwargs:227                        return method(**kwargs)228                    else:229                        return method()230                except Exception:231                    continue232        except Exception:233            return None234 235    def _check_allowed_json_modules(self, value: dict[str, Any]) -> None:236        needed = tuple(value["id"])237        method = value.get("method")238        if isinstance(method, list):239            method_display = ",".join(m or "<init>" for m in method)240        elif isinstance(method, str):241            method_display = method242        else:243            method_display = "<init>"244 245        dotted = ".".join(needed)246        # Safe types (the same set already allowed for msgpack deserialization) are247        # permitted without an explicit allowlist — they are known-safe LangGraph and248        # LangChain types.  This restores backwards-compat for old "json" checkpoints249        # that pre-date the msgpack migration without reopening the broader security gate.250        if _is_safe_json_type(list(needed)):251            return252 253        if not self._allowed_json_modules:254            raise InvalidModuleError(255                f"Refused to deserialize JSON constructor: {dotted} (method: {method_display}). "256                "No allowed_json_modules configured.\n\n"257                "Unblock with ONE of:\n"258                f"  • JsonPlusSerializer(allowed_json_modules=[{needed!r}, ...])\n"259                "  • (DANGEROUS) JsonPlusSerializer(allowed_json_modules=True)\n\n"260                "Note: Prefix allowlists are intentionally unsupported; prefer exact symbols "261                "or plain-JSON representations revived without import-time side effects."262            )263 264        if self._allowed_json_modules is True:265            return266        if needed in self._allowed_json_modules:267            return268 269        raise InvalidModuleError(270            f"Refused to deserialize JSON constructor: {dotted} (method: {method_display}). "271            "Symbol is not in the deserialization allowlist.\n\n"272            "Add exactly this symbol to unblock:\n"273            f"  JsonPlusSerializer(allowed_json_modules=[{needed!r}, ...])\n"274            "Or, as a last resort (DANGEROUS):\n"275            "  JsonPlusSerializer(allowed_json_modules=True)"276        )277 278    def dumps_typed(self, obj: Any) -> tuple[str, bytes]:279        if obj is None:280            return "null", EMPTY_BYTES281        elif isinstance(obj, bytes):282            return "bytes", obj283        elif isinstance(obj, bytearray):284            return "bytearray", obj285        else:286            try:287                return "msgpack", _msgpack_enc(obj)288            except ormsgpack.MsgpackEncodeError as exc:289                if self.pickle_fallback:290                    return "pickle", pickle.dumps(obj)291                raise exc292 293    def loads_typed(self, data: tuple[str, bytes]) -> Any:294        type_, data_ = data295        if type_ == "null":296            return None297        elif type_ == "bytes":298            return data_299        elif type_ == "bytearray":300            return bytearray(data_)301        elif type_ == "json":302            return json.loads(data_, object_hook=self._reviver)303        elif type_ == "msgpack":304            return ormsgpack.unpackb(305                data_, ext_hook=self._unpack_ext_hook, option=ormsgpack.OPT_NON_STR_KEYS306            )307        elif self.pickle_fallback and type_ == "pickle":308            return pickle.loads(data_)309        else:310            raise NotImplementedError(f"Unknown serialization type: {type_}")311 312 313# --- msgpack ---314 315EXT_CONSTRUCTOR_SINGLE_ARG = 0316EXT_CONSTRUCTOR_POS_ARGS = 1317EXT_CONSTRUCTOR_KW_ARGS = 2318EXT_METHOD_SINGLE_ARG = 3319EXT_PYDANTIC_V1 = 4320EXT_PYDANTIC_V2 = 5321EXT_NUMPY_ARRAY = 6322EXT_DELTA_SNAPSHOT = 7323 324 325def _msgpack_default(obj: Any) -> str | ormsgpack.Ext:326    if isinstance(obj, _DeltaSnapshot):327        return ormsgpack.Ext(EXT_DELTA_SNAPSHOT, _msgpack_enc(obj.value))328    elif hasattr(obj, "model_dump") and callable(obj.model_dump):  # pydantic v2329        return ormsgpack.Ext(330            EXT_PYDANTIC_V2,331            _msgpack_enc(332                (333                    obj.__class__.__module__,334                    obj.__class__.__name__,335                    obj.model_dump(),336                    "model_validate_json",337                ),338            ),339        )340    elif hasattr(obj, "get_secret_value") and callable(obj.get_secret_value):341        return ormsgpack.Ext(342            EXT_CONSTRUCTOR_SINGLE_ARG,343            _msgpack_enc(344                (345                    obj.__class__.__module__,346                    obj.__class__.__name__,347                    obj.get_secret_value(),348                ),349            ),350        )351    elif hasattr(obj, "dict") and callable(obj.dict):  # pydantic v1352        return ormsgpack.Ext(353            EXT_PYDANTIC_V1,354            _msgpack_enc(355                (356                    obj.__class__.__module__,357                    obj.__class__.__name__,358                    obj.dict(),359                ),360            ),361        )362    elif hasattr(obj, "_asdict") and callable(obj._asdict):  # namedtuple363        return ormsgpack.Ext(364            EXT_CONSTRUCTOR_KW_ARGS,365            _msgpack_enc(366                (367                    obj.__class__.__module__,368                    obj.__class__.__name__,369                    obj._asdict(),370                ),371            ),372        )373    elif isinstance(obj, pathlib.Path):374        return ormsgpack.Ext(375            EXT_CONSTRUCTOR_POS_ARGS,376            _msgpack_enc(377                (obj.__class__.__module__, obj.__class__.__name__, obj.parts),378            ),379        )380    elif isinstance(obj, re.Pattern):381        return ormsgpack.Ext(382            EXT_CONSTRUCTOR_POS_ARGS,383            _msgpack_enc(384                ("re", "compile", (obj.pattern, obj.flags)),385            ),386        )387    elif isinstance(obj, UUID):388        return ormsgpack.Ext(389            EXT_CONSTRUCTOR_SINGLE_ARG,390            _msgpack_enc(391                (obj.__class__.__module__, obj.__class__.__name__, obj.hex),392            ),393        )394    elif isinstance(obj, decimal.Decimal):395        return ormsgpack.Ext(396            EXT_CONSTRUCTOR_SINGLE_ARG,397            _msgpack_enc(398                (obj.__class__.__module__, obj.__class__.__name__, str(obj)),399            ),400        )401    elif isinstance(obj, (set, frozenset, deque)):402        return ormsgpack.Ext(403            EXT_CONSTRUCTOR_SINGLE_ARG,404            _msgpack_enc(405                (obj.__class__.__module__, obj.__class__.__name__, tuple(obj)),406            ),407        )408    elif isinstance(obj, (IPv4Address, IPv4Interface, IPv4Network)):409        return ormsgpack.Ext(410            EXT_CONSTRUCTOR_SINGLE_ARG,411            _msgpack_enc(412                (obj.__class__.__module__, obj.__class__.__name__, str(obj)),413            ),414        )415    elif isinstance(obj, (IPv6Address, IPv6Interface, IPv6Network)):416        return ormsgpack.Ext(417            EXT_CONSTRUCTOR_SINGLE_ARG,418            _msgpack_enc(419                (obj.__class__.__module__, obj.__class__.__name__, str(obj)),420            ),421        )422    elif isinstance(obj, datetime):423        return ormsgpack.Ext(424            EXT_METHOD_SINGLE_ARG,425            _msgpack_enc(426                (427                    obj.__class__.__module__,428                    obj.__class__.__name__,429                    obj.isoformat(),430                    "fromisoformat",431                ),432            ),433        )434    elif isinstance(obj, timedelta):435        return ormsgpack.Ext(436            EXT_CONSTRUCTOR_POS_ARGS,437            _msgpack_enc(438                (439                    obj.__class__.__module__,440                    obj.__class__.__name__,441                    (obj.days, obj.seconds, obj.microseconds),442                ),443            ),444        )445    elif isinstance(obj, date):446        return ormsgpack.Ext(447            EXT_CONSTRUCTOR_POS_ARGS,448            _msgpack_enc(449                (450                    obj.__class__.__module__,451                    obj.__class__.__name__,452                    (obj.year, obj.month, obj.day),453                ),454            ),455        )456    elif isinstance(obj, time):457        return ormsgpack.Ext(458            EXT_CONSTRUCTOR_KW_ARGS,459            _msgpack_enc(460                (461                    obj.__class__.__module__,462                    obj.__class__.__name__,463                    {464                        "hour": obj.hour,465                        "minute": obj.minute,466                        "second": obj.second,467                        "microsecond": obj.microsecond,468                        "tzinfo": obj.tzinfo,469                        "fold": obj.fold,470                    },471                ),472            ),473        )474    elif isinstance(obj, timezone):475        return ormsgpack.Ext(476            EXT_CONSTRUCTOR_POS_ARGS,477            _msgpack_enc(478                (479                    obj.__class__.__module__,480                    obj.__class__.__name__,481                    obj.__getinitargs__(),  # type: ignore[attr-defined]482                ),483            ),484        )485    elif isinstance(obj, ZoneInfo):486        return ormsgpack.Ext(487            EXT_CONSTRUCTOR_SINGLE_ARG,488            _msgpack_enc(489                (obj.__class__.__module__, obj.__class__.__name__, obj.key),490            ),491        )492    elif isinstance(obj, Enum):493        return ormsgpack.Ext(494            EXT_CONSTRUCTOR_SINGLE_ARG,495            _msgpack_enc(496                (obj.__class__.__module__, obj.__class__.__name__, obj.value),497            ),498        )499    elif isinstance(obj, SendProtocol):500        args: tuple[Any, ...] = (obj.node, obj.arg)501        if (timeout := getattr(obj, "timeout", None)) is not None:502            args = (obj.node, obj.arg, timeout)503        return ormsgpack.Ext(504            EXT_CONSTRUCTOR_POS_ARGS,505            _msgpack_enc(506                (obj.__class__.__module__, obj.__class__.__name__, args),507            ),508        )509    elif dataclasses.is_dataclass(obj):510        # doesn't use dataclasses.asdict to avoid deepcopy and recursion511        return ormsgpack.Ext(512            EXT_CONSTRUCTOR_KW_ARGS,513            _msgpack_enc(514                (515                    obj.__class__.__module__,516                    obj.__class__.__name__,517                    {518                        field.name: getattr(obj, field.name)519                        for field in dataclasses.fields(obj)520                    },521                ),522            ),523        )524    elif isinstance(obj, Item):525        return ormsgpack.Ext(526            EXT_CONSTRUCTOR_KW_ARGS,527            _msgpack_enc(528                (529                    obj.__class__.__module__,530                    obj.__class__.__name__,531                    {k: getattr(obj, k) for k in obj.__slots__},532                ),533            ),534        )535    elif (np_mod := sys.modules.get("numpy")) is not None and isinstance(536        obj, np_mod.ndarray537    ):538        order = "F" if obj.flags.f_contiguous and not obj.flags.c_contiguous else "C"539        if obj.flags.c_contiguous:540            mv = memoryview(obj)541            try:542                meta = (obj.dtype.str, obj.shape, order, mv)543                return ormsgpack.Ext(EXT_NUMPY_ARRAY, _msgpack_enc(meta))544            finally:545                mv.release()546        else:547            buf = obj.tobytes(order="A")548            meta = (obj.dtype.str, obj.shape, order, buf)549            return ormsgpack.Ext(EXT_NUMPY_ARRAY, _msgpack_enc(meta))550 551    elif isinstance(obj, BaseException):552        return repr(obj)553    else:554        raise TypeError(f"Object of type {obj.__class__.__name__} is not serializable")555 556 557def _send_from_args(args: Sequence[Any]) -> Any:558    # ya we have a cyclic import here ¯\_(ツ)_/¯559    from langgraph.types import Send  # type: ignore560 561    if len(args) == 2:562        return Send(*args)563    return Send(args[0], args[1], timeout=args[2])564 565 566def _create_msgpack_ext_hook(567    allowed_modules: set[tuple[str, ...]] | Literal[True] | None,568) -> Callable[[int, bytes], Any]:569    """Create msgpack ext hook with allowlist.570 571    Args:572        allowed_modules: Set of (module, name) tuples that are allowed to be573        deserialized, or True to allow all with warnings for unregistered types, or None to only allow safe types.574 575    Returns:576        An ext_hook function for use with ormsgpack.unpackb.577    """578 579    def _check_allowed(module: str, name: str) -> bool:580        """Check if type is allowed. Returns True if allowed, False if blocked."""581        key = (module, name)582 583        if key in _lg_msgpack.SAFE_MSGPACK_TYPES:584            return True585 586        if allowed_modules is True:587            # default is to warn but allow unregistered types588            emit_serde_event(589                {590                    "kind": "msgpack_unregistered_allowed",591                    "module": module,592                    "name": name,593                }594            )595            _warn_once(596                _warned_unregistered_types,597                key,598                "Deserializing unregistered type %s.%s from checkpoint. "599                "This will be blocked in a future version. "600                "Set LANGGRAPH_STRICT_MSGPACK=true to block now, or add "601                "to allowed_msgpack_modules to allow explicitly: [(%r, %r)]",602                module,603                name,604                module,605                name,606            )607            return True608        if allowed_modules is not None:609            if key in allowed_modules:610                return True611        # strict mode blocks unregistered types612        emit_serde_event(613            {614                "kind": "msgpack_blocked",615                "module": module,616                "name": name,617            }618        )619        _warn_once(620            _warned_blocked_types,621            key,622            "Blocked deserialization of %s.%s - not in allowed_msgpack_modules. "623            "Add to allowed_msgpack_modules to allow: [(%r, %r)]",624            module,625            name,626            module,627            name,628        )629        return False630 631    def _check_allowed_method(module: str, name: str, method: str) -> bool:632        """Check if a method invocation is allowed."""633        key = (module, name, method)634        if key in _lg_msgpack.SAFE_MSGPACK_METHODS:635            return True636        emit_serde_event(637            {638                "kind": "msgpack_method_blocked",639                "module": module,640                "name": name,641                "method": method,642            }643        )644        logger.warning(645            "Blocked deserialization of method call %s.%s.%s - "646            "not in allowed methods set.",647            module,648            name,649            method,650        )651        return False652 653    def ext_hook(code: int, data: bytes) -> Any:654        if code == EXT_DELTA_SNAPSHOT:655            return _DeltaSnapshot(656                ormsgpack.unpackb(657                    data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS658                )659            )660        elif code == EXT_CONSTRUCTOR_SINGLE_ARG:661            try:662                tup = ormsgpack.unpackb(663                    data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS664                )665                if not _check_allowed(tup[0], tup[1]):666                    # We default to returning the raw data. If the user667                    # is using this in the context of a pydantic state, etc., then668                    # it would be validated upon construction.669                    return tup[2]670                # module, name, arg671                return getattr(importlib.import_module(tup[0]), tup[1])(tup[2])672            except Exception:673                return None674        elif code == EXT_CONSTRUCTOR_POS_ARGS:675            try:676                tup = ormsgpack.unpackb(677                    data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS678                )679                if not _check_allowed(tup[0], tup[1]):680                    return tup[2]681                if tup[0] == "langgraph.types" and tup[1] == "Send":682                    return _send_from_args(tup[2])683                # module, name, args684                return getattr(importlib.import_module(tup[0]), tup[1])(*tup[2])685            except Exception:686                return None687        elif code == EXT_CONSTRUCTOR_KW_ARGS:688            try:689                tup = ormsgpack.unpackb(690                    data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS691                )692                if not _check_allowed(tup[0], tup[1]):693                    return tup[2]694                # module, name, kwargs695                return getattr(importlib.import_module(tup[0]), tup[1])(**tup[2])696            except Exception:697                return None698        elif code == EXT_METHOD_SINGLE_ARG:699            try:700                tup = ormsgpack.unpackb(701                    data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS702                )703                if not _check_allowed_method(tup[0], tup[1], tup[3]):704                    return tup[2]705                # module, name, arg, method706                return getattr(707                    getattr(importlib.import_module(tup[0]), tup[1]), tup[3]708                )(tup[2])709            except Exception:710                return None711        elif code == EXT_PYDANTIC_V1:712            try:713                tup = ormsgpack.unpackb(714                    data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS715                )716                if not _check_allowed(tup[0], tup[1]):717                    return tup[2]718                # module, name, kwargs719                cls = getattr(importlib.import_module(tup[0]), tup[1])720                try:721                    return cls(**tup[2])722                except Exception:723                    return cls.construct(**tup[2])724            except Exception:725                # for pydantic objects we can't find/reconstruct726                # let's return the kwargs dict instead727                try:728                    return tup[2]729                except NameError:730                    return None731        elif code == EXT_PYDANTIC_V2:732            try:733                tup = ormsgpack.unpackb(734                    data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS735                )736                if not _check_allowed(tup[0], tup[1]):737                    return tup[2]738                # module, name, kwargs, method739                cls = getattr(importlib.import_module(tup[0]), tup[1])740                try:741                    return cls(**tup[2])742                except Exception:743                    return cls.model_construct(**tup[2])744            except Exception:745                # for pydantic objects we can't find/reconstruct746                # let's return the kwargs dict instead747                try:748                    return tup[2]749                except NameError:750                    return None751        elif code == EXT_NUMPY_ARRAY:752            try:753                import numpy as _np754 755                dtype_str, shape, order, buf = ormsgpack.unpackb(756                    data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS757                )758                arr = _np.frombuffer(buf, dtype=_np.dtype(dtype_str))759                return arr.reshape(shape, order=order)760            except Exception:761                return None762        return None763 764    return ext_hook765 766 767# Aliasing in case anyone imported it directly768_msgpack_ext_hook = _create_msgpack_ext_hook(allowed_modules=None)769 770 771def _msgpack_ext_hook_to_json(code: int, data: bytes) -> Any:772    if code == EXT_CONSTRUCTOR_SINGLE_ARG:773        try:774            tup = ormsgpack.unpackb(775                data,776                ext_hook=_msgpack_ext_hook_to_json,777                option=ormsgpack.OPT_NON_STR_KEYS,778            )779            if tup[0] == "uuid" and tup[1] == "UUID":780                hex_ = tup[2]781                return (782                    f"{hex_[:8]}-{hex_[8:12]}-{hex_[12:16]}-{hex_[16:20]}-{hex_[20:]}"783                )784            # module, name, arg785            return tup[2]786        except Exception:787            return788    elif code == EXT_CONSTRUCTOR_POS_ARGS:789        try:790            tup = ormsgpack.unpackb(791                data,792                ext_hook=_msgpack_ext_hook_to_json,793                option=ormsgpack.OPT_NON_STR_KEYS,794            )795            if tup[0] == "langgraph.types" and tup[1] == "Send":796                return _send_from_args(tup[2])797            # module, name, args798            return tup[2]799        except Exception:800            return801    elif code == EXT_CONSTRUCTOR_KW_ARGS:802        try:803            tup = ormsgpack.unpackb(804                data,805                ext_hook=_msgpack_ext_hook_to_json,806                option=ormsgpack.OPT_NON_STR_KEYS,807            )808            # module, name, args809            return tup[2]810        except Exception:811            return812    elif code == EXT_METHOD_SINGLE_ARG:813        try:814            tup = ormsgpack.unpackb(815                data,816                ext_hook=_msgpack_ext_hook_to_json,817                option=ormsgpack.OPT_NON_STR_KEYS,818            )819            # module, name, arg, method820            return tup[2]821        except Exception:822            return823    elif code == EXT_PYDANTIC_V1:824        try:825            tup = ormsgpack.unpackb(826                data,827                ext_hook=_msgpack_ext_hook_to_json,828                option=ormsgpack.OPT_NON_STR_KEYS,829            )830            # module, name, kwargs831            return tup[2]832        except Exception:833            # for pydantic objects we can't find/reconstruct834            # let's return the kwargs dict instead835            return836    elif code == EXT_PYDANTIC_V2:837        try:838            tup = ormsgpack.unpackb(839                data,840                ext_hook=_msgpack_ext_hook_to_json,841                option=ormsgpack.OPT_NON_STR_KEYS,842            )843            # module, name, kwargs, method844            return tup[2]845        except Exception:846            return847    elif code == EXT_NUMPY_ARRAY:848        try:849            import numpy as _np850 851            dtype_str, shape, order, buf = ormsgpack.unpackb(852                data,853                ext_hook=_msgpack_ext_hook_to_json,854                option=ormsgpack.OPT_NON_STR_KEYS,855            )856            arr = _np.frombuffer(buf, dtype=_np.dtype(dtype_str))857            return arr.reshape(shape, order=order).tolist()858        except Exception:859            return860 861 862class InvalidModuleError(Exception):863    """Exception raised when a module is not in the allowlist."""864 865    def __init__(self, message: str):866        self.message = message867 868 869_option = (870    ormsgpack.OPT_NON_STR_KEYS871    | ormsgpack.OPT_PASSTHROUGH_DATACLASS872    | ormsgpack.OPT_PASSTHROUGH_DATETIME873    | ormsgpack.OPT_PASSTHROUGH_ENUM874    | ormsgpack.OPT_PASSTHROUGH_UUID875    | ormsgpack.OPT_REPLACE_SURROGATES876)877 878 879def _msgpack_enc(data: Any) -> bytes:880    return ormsgpack.packb(data, default=_msgpack_default, option=_option)881 882 883def _normalize_allowlist(884    allowlist: AllowedMsgpackModules | Literal[True] | None,885) -> set[tuple[str, ...]] | Literal[True] | None:886    if allowlist is True:887        return allowlist888    elif allowlist:889        return _normalize_module_keys(allowlist)890    else:891        return None892 893 894def _normalize_module_keys(895    modules: AllowedMsgpackModules,896) -> set[tuple[str, ...]]:897    normalized: set[tuple[str, ...]] = set()898    for module in modules:899        if isclass(module):900            normalized.add((module.__module__, module.__name__))901        else:902            normalized.add(cast(tuple[str, ...], module))903    return normalized904 
codekingpro/portable-devtools · Team Ai