codekingpro/portable-devtools
114k
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 