Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_serde.py254 linesDownload Raw Back to _internal
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 
codekingpro/portable-devtools · Team Ai