codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import dataclasses4import logging5import sys6import types7from collections import deque8from enum import Enum9from typing import (10 Annotated,11 Any,12 Literal,13 Union,14 get_args,15 get_origin,16 get_type_hints,17)18 19from langchain_core import messages as lc_messages20from langgraph.checkpoint.base import BaseCheckpointSaver21from pydantic import BaseModel22from typing_extensions import NotRequired, Required, is_typeddict23 24try:25 from langgraph.checkpoint.serde._msgpack import ( # noqa: F40126 STRICT_MSGPACK_ENABLED,27 )28except ImportError:29 STRICT_MSGPACK_ENABLED = False30 31_warned_allowlist_unsupported = False32 33logger = logging.getLogger(__name__)34 35 36def _supports_checkpointer_allowlist() -> bool:37 return hasattr(BaseCheckpointSaver, "with_allowlist")38 39 40_SUPPORTS_ALLOWLIST = _supports_checkpointer_allowlist()41 42 43def apply_checkpointer_allowlist(44 checkpointer: Any, allowlist: set[tuple[str, ...]] | None45) -> Any:46 if not checkpointer or allowlist is None or checkpointer in (True, False):47 return checkpointer48 if not _SUPPORTS_ALLOWLIST:49 global _warned_allowlist_unsupported50 if not _warned_allowlist_unsupported:51 logger.warning(52 "Checkpointer does not support with_allowlist; strict msgpack "53 "allowlist will be skipped."54 )55 _warned_allowlist_unsupported = True56 return checkpointer57 return checkpointer.with_allowlist(allowlist)58 59 60def curated_core_allowlist() -> set[tuple[str, ...]]:61 allowlist: set[tuple[str, ...]] = set()62 for name in (63 "BaseMessage",64 "BaseMessageChunk",65 "HumanMessage",66 "HumanMessageChunk",67 "AIMessage",68 "AIMessageChunk",69 "SystemMessage",70 "SystemMessageChunk",71 "ChatMessage",72 "ChatMessageChunk",73 "ToolMessage",74 "ToolMessageChunk",75 "FunctionMessage",76 "FunctionMessageChunk",77 "RemoveMessage",78 ):79 cls = getattr(lc_messages, name, None)80 if cls is None:81 continue82 allowlist.add((cls.__module__, cls.__name__))83 84 return allowlist85 86 87def build_serde_allowlist(88 *,89 schemas: list[type[Any]] | None = None,90 channels: dict[str, Any] | None = None,91) -> set[tuple[str, ...]]:92 allowlist = curated_core_allowlist()93 if schemas:94 schemas = [schema for schema in schemas if schema is not None]95 return allowlist | collect_allowlist_from_schemas(96 schemas=schemas,97 channels=channels,98 )99 100 101def collect_allowlist_from_schemas(102 *,103 schemas: list[type[Any]] | None = None,104 channels: dict[str, Any] | None = None,105) -> set[tuple[str, ...]]:106 allowlist: set[tuple[str, ...]] = set()107 seen: set[Any] = set()108 seen_ids: set[int] = set()109 110 if schemas:111 for schema in schemas:112 _collect_from_type(schema, allowlist, seen, seen_ids)113 114 if channels:115 for channel in channels.values():116 value_type = getattr(channel, "ValueType", None)117 if value_type is not None:118 _collect_from_type(value_type, allowlist, seen, seen_ids)119 update_type = getattr(channel, "UpdateType", None)120 if update_type is not None:121 _collect_from_type(update_type, allowlist, seen, seen_ids)122 123 return allowlist124 125 126def _collect_from_type(127 typ: Any,128 allowlist: set[tuple[str, ...]],129 seen: set[Any],130 seen_ids: set[int],131) -> None:132 if _already_seen(typ, seen, seen_ids):133 return134 135 if typ is Any or typ is None:136 return137 138 if typ is Literal:139 return140 141 if isinstance(typ, types.UnionType):142 for arg in typ.__args__:143 _collect_from_type(arg, allowlist, seen, seen_ids)144 return145 146 origin = get_origin(typ)147 if origin is Union:148 for arg in get_args(typ):149 _collect_from_type(arg, allowlist, seen, seen_ids)150 return151 if origin is Annotated or origin in (Required, NotRequired):152 args = get_args(typ)153 if args:154 _collect_from_type(args[0], allowlist, seen, seen_ids)155 return156 157 if origin is Literal:158 return159 160 if origin in (list, set, tuple, dict, deque, frozenset):161 for arg in get_args(typ):162 _collect_from_type(arg, allowlist, seen, seen_ids)163 return164 165 if hasattr(typ, "__supertype__"):166 _collect_from_type(typ.__supertype__, allowlist, seen, seen_ids)167 return168 169 if is_typeddict(typ):170 for field_type in _safe_get_type_hints(typ).values():171 _collect_from_type(field_type, allowlist, seen, seen_ids)172 return173 174 if _is_pydantic_model(typ):175 allowlist.add((typ.__module__, typ.__name__))176 field_types = _safe_get_type_hints(typ)177 if field_types:178 for field_type in field_types.values():179 _collect_from_type(field_type, allowlist, seen, seen_ids)180 else:181 for field_type in _pydantic_field_types(typ):182 _collect_from_type(field_type, allowlist, seen, seen_ids)183 return184 185 if dataclasses.is_dataclass(typ):186 if typ_name := getattr(typ, "__name__", None):187 allowlist.add((typ.__module__, typ_name))188 field_types = _safe_get_type_hints(typ)189 if field_types:190 for field_type in field_types.values():191 _collect_from_type(field_type, allowlist, seen, seen_ids)192 else:193 for field in dataclasses.fields(typ):194 _collect_from_type(field.type, allowlist, seen, seen_ids)195 return196 197 if isinstance(typ, type) and issubclass(typ, Enum):198 allowlist.add((typ.__module__, typ.__name__))199 return200 201 202def _already_seen(typ: Any, seen: set[Any], seen_ids: set[int]) -> bool:203 try:204 if typ in seen:205 return True206 seen.add(typ)207 return False208 except TypeError:209 typ_id = id(typ)210 if typ_id in seen_ids:211 return True212 seen_ids.add(typ_id)213 return False214 215 216def _safe_get_type_hints(typ: Any) -> dict[str, Any]:217 try:218 module = sys.modules.get(getattr(typ, "__module__", ""))219 globalns = module.__dict__ if module else None220 localns = dict(vars(typ)) if hasattr(typ, "__dict__") else None221 return get_type_hints(222 typ, globalns=globalns, localns=localns, include_extras=True223 )224 except Exception:225 return {}226 227 228def _is_pydantic_model(typ: Any) -> bool:229 if not isinstance(typ, type):230 return False231 if issubclass(typ, BaseModel):232 return True233 try:234 from pydantic.v1 import BaseModel as BaseModelV1235 except Exception:236 return False237 return issubclass(typ, BaseModelV1)238 239 240def _pydantic_field_types(typ: type[Any]) -> list[Any]:241 if hasattr(typ, "model_fields"):242 return [243 field.annotation244 for field in typ.model_fields.values()245 if getattr(field, "annotation", None) is not None246 ]247 if hasattr(typ, "__fields__"):248 return [249 field.outer_type_250 for field in typ.__fields__.values()251 if getattr(field, "outer_type_", None) is not None252 ]253 return []254 