Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
__init__.py705 linesDownload Raw Back to memory
1from __future__ import annotations2 3import logging4import os5import pickle6import random7import shutil8from collections import defaultdict9from collections.abc import AsyncIterator, Iterator, Mapping, Sequence10from contextlib import AbstractAsyncContextManager, AbstractContextManager, ExitStack11from types import TracebackType12from typing import Any13 14from langchain_core.runnables import RunnableConfig15 16from langgraph.checkpoint.base import (17    WRITES_IDX_MAP,18    BaseCheckpointSaver,19    ChannelVersions,20    Checkpoint,21    CheckpointMetadata,22    CheckpointTuple,23    DeltaChannelHistory,24    PendingWrite,25    SerializerProtocol,26    get_checkpoint_id,27    get_checkpoint_metadata,28)29 30logger = logging.getLogger(__name__)31 32 33class InMemorySaver(34    BaseCheckpointSaver[str], AbstractContextManager, AbstractAsyncContextManager35):36    """An in-memory checkpoint saver.37 38    This checkpoint saver stores checkpoints in memory using a `defaultdict`.39 40    Note:41        Only use `InMemorySaver` for debugging or testing purposes.42        For production use cases we recommend installing [langgraph-checkpoint-postgres](https://pypi.org/project/langgraph-checkpoint-postgres/) and using `PostgresSaver` / `AsyncPostgresSaver`.43 44        If you are using LangSmith Deployment, no checkpointer needs to be specified. The correct managed checkpointer will be used automatically.45 46    Args:47        serde: The serializer to use for serializing and deserializing checkpoints.48 49    Example:50        ```python51        import asyncio52 53        from langgraph.checkpoint.memory import InMemorySaver54        from langgraph.graph import StateGraph55 56        builder = StateGraph(int)57        builder.add_node("add_one", lambda x: x + 1)58        builder.set_entry_point("add_one")59        builder.set_finish_point("add_one")60 61        memory = InMemorySaver()62        graph = builder.compile(checkpointer=memory)63        coro = graph.ainvoke(1, {"configurable": {"thread_id": "thread-1"}})64        asyncio.run(coro)  # Output: 265        ```66    """67 68    # thread ID ->  checkpoint NS -> checkpoint ID -> checkpoint mapping69    storage: defaultdict[70        str,71        dict[str, dict[str, tuple[tuple[str, bytes], tuple[str, bytes], str | None]]],72    ]73    # (thread ID, checkpoint NS, checkpoint ID) -> (task ID, write idx)74    writes: defaultdict[75        tuple[str, str, str],76        dict[tuple[str, int], tuple[str, str, tuple[str, bytes], str]],77    ]78    blobs: dict[79        tuple[80            str, str, str, str | int | float81        ],  # thread id, checkpoint ns, channel, version82        tuple[str, bytes],83    ]84 85    def __init__(86        self,87        *,88        serde: SerializerProtocol | None = None,89        factory: type[defaultdict] = defaultdict,90    ) -> None:91        super().__init__(serde=serde)92        self.storage = factory(lambda: defaultdict(dict))93        self.writes = factory(dict)94        self.blobs = factory()95        self.stack = ExitStack()96        if factory is not defaultdict:97            self.stack.enter_context(self.storage)  # type: ignore[arg-type]98            self.stack.enter_context(self.writes)  # type: ignore[arg-type]99            self.stack.enter_context(self.blobs)  # type: ignore[arg-type]100 101    def __enter__(self) -> InMemorySaver:102        self.stack.__enter__()103        return self104 105    def __exit__(106        self,107        exc_type: type[BaseException] | None,108        exc_value: BaseException | None,109        traceback: TracebackType | None,110    ) -> bool | None:111        return self.stack.__exit__(exc_type, exc_value, traceback)112 113    async def __aenter__(self) -> InMemorySaver:114        self.stack.__enter__()115        return self116 117    async def __aexit__(118        self,119        __exc_type: type[BaseException] | None,120        __exc_value: BaseException | None,121        __traceback: TracebackType | None,122    ) -> bool | None:123        return self.stack.__exit__(__exc_type, __exc_value, __traceback)124 125    def _load_blobs(126        self,127        thread_id: str,128        checkpoint_ns: str,129        versions: ChannelVersions,130    ) -> dict[str, Any]:131        result: dict[str, Any] = {}132        for k, ver in versions.items():133            kk = (thread_id, checkpoint_ns, k, ver)134            if kk not in self.blobs:135                continue136            vv = self.blobs[kk]137            if vv[0] == "empty":138                continue139            result[k] = self.serde.loads_typed(vv)140        return result141 142    def get_delta_channel_history(143        self, *, config: RunnableConfig, channels: Sequence[str]144    ) -> Mapping[str, DeltaChannelHistory]:145        """Override: walk the parent chain ONCE for all requested channels.146 147        Each channel terminates independently at the nearest ancestor148        whose stored blob is non-empty. Other channels keep walking until149        they find their own terminator or hit the root.150 151        Pre-delta plain-value blobs subsume their ancestor's pending152        writes (the value already includes them); `_DeltaSnapshot` blobs153        do not (snapshot is the value AT that ancestor, prior to its own154        pending writes that produce the child).155        """156        if not channels:157            return {}158        # Imported lazily to avoid a hard checkpoint→serde-types coupling at159        # module import; only this override needs the runtime check.160        from langgraph.checkpoint.serde.types import _DeltaSnapshot161 162        thread_id = config["configurable"]["thread_id"]163        checkpoint_ns = config["configurable"].get("checkpoint_ns", "")164        checkpoint_id = config["configurable"].get("checkpoint_id", "")165        ns_storage = self.storage.get(thread_id, {}).get(checkpoint_ns, {})166 167        chain: list[str] = []168        target_entry = ns_storage.get(checkpoint_id)169        current: str | None = target_entry[2] if target_entry is not None else None170        while current is not None:171            entry = ns_storage.get(current)172            if entry is None:173                break174            chain.append(current)175            _, _, parent = entry176            current = parent177 178        collected_by_ch: dict[str, list[PendingWrite]] = {c: [] for c in channels}179        seed_by_ch: dict[str, Any] = {}180        remaining: set[str] = set(channels)181 182        for cp_id in chain:183            if not remaining:184                break185            entry = ns_storage.get(cp_id)186            ckpt = self.serde.loads_typed(entry[0]) if entry is not None else None187 188            terminated_here: set[str] = set()189            blob_value_by_ch: dict[str, Any] = {}190            if ckpt is not None:191                versions = ckpt.get("channel_versions", {})192                for ch in remaining:193                    ver = versions.get(ch)194                    if ver is None:195                        continue196                    blob_entry = self.blobs.get((thread_id, checkpoint_ns, ch, ver))197                    if blob_entry is None or blob_entry[0] == "empty":198                        continue199                    blob_value_by_ch[ch] = self.serde.loads_typed(blob_entry)200                    terminated_here.add(ch)201 202            step_writes = self.writes.get((thread_id, checkpoint_ns, cp_id), {})203            for (_task_id, _idx), (tid, ch, serialized, _) in sorted(204                step_writes.items(), reverse=True205            ):206                if ch not in remaining:207                    continue208                blob_value = blob_value_by_ch.get(ch)209                if blob_value is not None and not isinstance(210                    blob_value, _DeltaSnapshot211                ):212                    continue213                collected_by_ch[ch].append(214                    (tid, ch, self.serde.loads_typed(serialized))215                )216 217            for ch in terminated_here:218                seed_by_ch[ch] = blob_value_by_ch[ch]219                remaining.discard(ch)220 221        result: dict[str, DeltaChannelHistory] = {}222        for ch in channels:223            entry_h: DeltaChannelHistory = {224                "writes": list(reversed(collected_by_ch[ch]))225            }226            if ch in seed_by_ch:227                entry_h["seed"] = seed_by_ch[ch]228            result[ch] = entry_h229        return result230 231    async def aget_delta_channel_history(232        self, *, config: RunnableConfig, channels: Sequence[str]233    ) -> Mapping[str, DeltaChannelHistory]:234        return self.get_delta_channel_history(config=config, channels=channels)235 236    def get_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:237        """Get a checkpoint tuple from the in-memory storage.238 239        This method retrieves a checkpoint tuple from the in-memory storage based on the240        provided config. If the config contains a `checkpoint_id` key, the checkpoint with241        the matching thread ID and timestamp is retrieved. Otherwise, the latest checkpoint242        for the given thread ID is retrieved.243 244        Args:245            config: The config to use for retrieving the checkpoint.246 247        Returns:248            The retrieved checkpoint tuple, or None if no matching checkpoint was found.249        """250        thread_id: str = config["configurable"]["thread_id"]251        checkpoint_ns: str = config["configurable"].get("checkpoint_ns", "")252        if checkpoint_id := get_checkpoint_id(config):253            if saved := self.storage[thread_id][checkpoint_ns].get(checkpoint_id):254                checkpoint, metadata, parent_checkpoint_id = saved255                writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)].values()256                checkpoint_: Checkpoint = self.serde.loads_typed(checkpoint)257                return CheckpointTuple(258                    config=config,259                    checkpoint={260                        **checkpoint_,261                        "channel_values": self._load_blobs(262                            thread_id, checkpoint_ns, checkpoint_["channel_versions"]263                        ),264                    },265                    metadata=self.serde.loads_typed(metadata),266                    pending_writes=[267                        (id, c, self.serde.loads_typed(v)) for id, c, v, _ in writes268                    ],269                    parent_config=(270                        {271                            "configurable": {272                                "thread_id": thread_id,273                                "checkpoint_ns": checkpoint_ns,274                                "checkpoint_id": parent_checkpoint_id,275                            }276                        }277                        if parent_checkpoint_id278                        else None279                    ),280                )281        else:282            if checkpoints := self.storage[thread_id][checkpoint_ns]:283                checkpoint_id = max(checkpoints.keys())284                checkpoint, metadata, parent_checkpoint_id = checkpoints[checkpoint_id]285                writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)].values()286                checkpoint_ = self.serde.loads_typed(checkpoint)287                return CheckpointTuple(288                    config={289                        "configurable": {290                            "thread_id": thread_id,291                            "checkpoint_ns": checkpoint_ns,292                            "checkpoint_id": checkpoint_id,293                        }294                    },295                    checkpoint={296                        **checkpoint_,297                        "channel_values": self._load_blobs(298                            thread_id, checkpoint_ns, checkpoint_["channel_versions"]299                        ),300                    },301                    metadata=self.serde.loads_typed(metadata),302                    pending_writes=[303                        (id, c, self.serde.loads_typed(v)) for id, c, v, _ in writes304                    ],305                    parent_config=(306                        {307                            "configurable": {308                                "thread_id": thread_id,309                                "checkpoint_ns": checkpoint_ns,310                                "checkpoint_id": parent_checkpoint_id,311                            }312                        }313                        if parent_checkpoint_id314                        else None315                    ),316                )317 318    def list(319        self,320        config: RunnableConfig | None,321        *,322        filter: dict[str, Any] | None = None,323        before: RunnableConfig | None = None,324        limit: int | None = None,325    ) -> Iterator[CheckpointTuple]:326        """List checkpoints from the in-memory storage.327 328        This method retrieves a list of checkpoint tuples from the in-memory storage based329        on the provided criteria.330 331        Args:332            config: Base configuration for filtering checkpoints.333            filter: Additional filtering criteria for metadata.334            before: List checkpoints created before this configuration.335            limit: Maximum number of checkpoints to return.336 337        Yields:338            An iterator of matching checkpoint tuples.339        """340        thread_ids = (config["configurable"]["thread_id"],) if config else self.storage341        config_checkpoint_ns = (342            config["configurable"].get("checkpoint_ns") if config else None343        )344        config_checkpoint_id = get_checkpoint_id(config) if config else None345        for thread_id in thread_ids:346            for checkpoint_ns in self.storage[thread_id].keys():347                if (348                    config_checkpoint_ns is not None349                    and checkpoint_ns != config_checkpoint_ns350                ):351                    continue352 353                for checkpoint_id, (354                    checkpoint,355                    metadata_b,356                    parent_checkpoint_id,357                ) in sorted(358                    self.storage[thread_id][checkpoint_ns].items(),359                    key=lambda x: x[0],360                    reverse=True,361                ):362                    # filter by checkpoint ID from config363                    if config_checkpoint_id and checkpoint_id != config_checkpoint_id:364                        continue365 366                    # filter by checkpoint ID from `before` config367                    if (368                        before369                        and (before_checkpoint_id := get_checkpoint_id(before))370                        and checkpoint_id >= before_checkpoint_id371                    ):372                        continue373 374                    # filter by metadata375                    metadata = self.serde.loads_typed(metadata_b)376                    if filter and not all(377                        query_value == metadata.get(query_key)378                        for query_key, query_value in filter.items()379                    ):380                        continue381 382                    # limit search results383                    if limit is not None and limit <= 0:384                        break385                    elif limit is not None:386                        limit -= 1387 388                    writes = self.writes[389                        (thread_id, checkpoint_ns, checkpoint_id)390                    ].values()391 392                    checkpoint_: Checkpoint = self.serde.loads_typed(checkpoint)393 394                    yield CheckpointTuple(395                        config={396                            "configurable": {397                                "thread_id": thread_id,398                                "checkpoint_ns": checkpoint_ns,399                                "checkpoint_id": checkpoint_id,400                            }401                        },402                        checkpoint={403                            **checkpoint_,404                            "channel_values": self._load_blobs(405                                thread_id,406                                checkpoint_ns,407                                checkpoint_["channel_versions"],408                            ),409                        },410                        metadata=metadata,411                        parent_config=(412                            {413                                "configurable": {414                                    "thread_id": thread_id,415                                    "checkpoint_ns": checkpoint_ns,416                                    "checkpoint_id": parent_checkpoint_id,417                                }418                            }419                            if parent_checkpoint_id420                            else None421                        ),422                        pending_writes=[423                            (id, c, self.serde.loads_typed(v)) for id, c, v, _ in writes424                        ],425                    )426 427    def put(428        self,429        config: RunnableConfig,430        checkpoint: Checkpoint,431        metadata: CheckpointMetadata,432        new_versions: ChannelVersions,433    ) -> RunnableConfig:434        """Save a checkpoint to the in-memory storage.435 436        This method saves a checkpoint to the in-memory storage. The checkpoint is associated437        with the provided config.438 439        Args:440            config: The config to associate with the checkpoint.441            checkpoint: The checkpoint to save.442            metadata: Additional metadata to save with the checkpoint.443            new_versions: New versions as of this write444 445        Returns:446            RunnableConfig: The updated config containing the saved checkpoint's timestamp.447        """448        c = checkpoint.copy()449        thread_id = config["configurable"]["thread_id"]450        checkpoint_ns = config["configurable"]["checkpoint_ns"]451        values: dict[str, Any] = c.pop("channel_values")  # type: ignore[misc]452        for k, v in new_versions.items():453            self.blobs[(thread_id, checkpoint_ns, k, v)] = (454                self.serde.dumps_typed(values[k]) if k in values else ("empty", b"")455            )456        self.storage[thread_id][checkpoint_ns].update(457            {458                checkpoint["id"]: (459                    self.serde.dumps_typed(c),460                    self.serde.dumps_typed(get_checkpoint_metadata(config, metadata)),461                    config["configurable"].get("checkpoint_id"),  # parent462                )463            }464        )465        return {466            "configurable": {467                "thread_id": thread_id,468                "checkpoint_ns": checkpoint_ns,469                "checkpoint_id": checkpoint["id"],470            }471        }472 473    def put_writes(474        self,475        config: RunnableConfig,476        writes: Sequence[tuple[str, Any]],477        task_id: str,478        task_path: str = "",479    ) -> None:480        """Save a list of writes to the in-memory storage.481 482        This method saves a list of writes to the in-memory storage. The writes are associated483        with the provided config.484 485        Args:486            config: The config to associate with the writes.487            writes: The writes to save.488            task_id: Identifier for the task creating the writes.489            task_path: Path of the task creating the writes.490 491        Returns:492            RunnableConfig: The updated config containing the saved writes' timestamp.493        """494        thread_id = config["configurable"]["thread_id"]495        checkpoint_ns = config["configurable"].get("checkpoint_ns", "")496        checkpoint_id = config["configurable"]["checkpoint_id"]497        outer_key = (thread_id, checkpoint_ns, checkpoint_id)498        outer_writes_ = self.writes.get(outer_key)499        for idx, (c, v) in enumerate(writes):500            inner_key = (task_id, WRITES_IDX_MAP.get(c, idx))501            if inner_key[1] >= 0 and outer_writes_ and inner_key in outer_writes_:502                continue503 504            self.writes[outer_key][inner_key] = (505                task_id,506                c,507                self.serde.dumps_typed(v),508                task_path,509            )510 511    def delete_thread(self, thread_id: str) -> None:512        """Delete all checkpoints and writes associated with a thread ID.513 514        Args:515            thread_id: The thread ID to delete.516 517        Returns:518            None519        """520        if thread_id in self.storage:521            del self.storage[thread_id]522        for k in list(self.writes.keys()):523            if k[0] == thread_id:524                del self.writes[k]525        for k in list(self.blobs.keys()):526            if k[0] == thread_id:527                del self.blobs[k]528 529    async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:530        """Asynchronous version of `get_tuple`.531 532        This method is an asynchronous wrapper around `get_tuple` that runs the synchronous533        method in a separate thread using asyncio.534 535        Args:536            config: The config to use for retrieving the checkpoint.537 538        Returns:539            The retrieved checkpoint tuple, or None if no matching checkpoint was found.540        """541        return self.get_tuple(config)542 543    async def alist(544        self,545        config: RunnableConfig | None,546        *,547        filter: dict[str, Any] | None = None,548        before: RunnableConfig | None = None,549        limit: int | None = None,550    ) -> AsyncIterator[CheckpointTuple]:551        """Asynchronous version of `list`.552 553        This method is an asynchronous wrapper around `list` that runs the synchronous554        method in a separate thread using asyncio.555 556        Args:557            config: The config to use for listing the checkpoints.558 559        Yields:560            An asynchronous iterator of checkpoint tuples.561        """562        for item in self.list(config, filter=filter, before=before, limit=limit):563            yield item564 565    async def aput(566        self,567        config: RunnableConfig,568        checkpoint: Checkpoint,569        metadata: CheckpointMetadata,570        new_versions: ChannelVersions,571    ) -> RunnableConfig:572        """Asynchronous version of `put`.573 574        Args:575            config: The config to associate with the checkpoint.576            checkpoint: The checkpoint to save.577            metadata: Additional metadata to save with the checkpoint.578            new_versions: New versions as of this write579 580        Returns:581            RunnableConfig: The updated config containing the saved checkpoint's timestamp.582        """583        return self.put(config, checkpoint, metadata, new_versions)584 585    async def aput_writes(586        self,587        config: RunnableConfig,588        writes: Sequence[tuple[str, Any]],589        task_id: str,590        task_path: str = "",591    ) -> None:592        """Asynchronous version of `put_writes`.593 594        This method is an asynchronous wrapper around `put_writes` that runs the synchronous595        method in a separate thread using asyncio.596 597        Args:598            config: The config to associate with the writes.599            writes: The writes to save, each as a (channel, value) pair.600            task_id: Identifier for the task creating the writes.601            task_path: Path of the task creating the writes.602 603        Returns:604            None605        """606        return self.put_writes(config, writes, task_id, task_path)607 608    async def adelete_thread(self, thread_id: str) -> None:609        """Delete all checkpoints and writes associated with a thread ID.610 611        Args:612            thread_id: The thread ID to delete.613 614        Returns:615            None616        """617        return self.delete_thread(thread_id)618 619    def get_next_version(self, current: str | None, channel: None) -> str:620        if current is None:621            current_v = 0622        elif isinstance(current, int):623            current_v = current624        else:625            current_v = int(current.split(".")[0])626        next_v = current_v + 1627        next_h = random.random()628        return f"{next_v:032}.{next_h:016}"629 630 631MemorySaver = InMemorySaver  # Kept for backwards compatibility632 633 634class PersistentDict(defaultdict):635    """Persistent dictionary with an API compatible with shelve and anydbm.636 637    The dict is kept in memory, so the dictionary operations run as fast as638    a regular dictionary.639 640    Write to disk is delayed until close or sync (similar to gdbm's fast mode).641 642    Input file format is automatically discovered.643    Output file format is selectable between pickle, json, and csv.644    All three serialization formats are backed by fast C implementations.645 646    Adapted from https://code.activestate.com/recipes/576642-persistent-dict-with-multiple-standard-file-format/647 648    """649 650    def __init__(self, *args: Any, filename: str, **kwds: Any) -> None:651        self.flag = "c"  # r=readonly, c=create, or n=new652        self.mode = None  # None or an octal triple like 0644653        self.format = "pickle"  # 'csv', 'json', or 'pickle'654        self.filename = filename655        super().__init__(*args, **kwds)656 657    def sync(self) -> None:658        "Write dict to disk"659        if self.flag == "r":660            return661        tempname = self.filename + ".tmp"662        fileobj = open(tempname, "wb" if self.format == "pickle" else "w")663        try:664            self.dump(fileobj)665        except Exception:666            os.remove(tempname)667            raise668        finally:669            fileobj.close()670        shutil.move(tempname, self.filename)  # atomic commit671        if self.mode is not None:672            os.chmod(self.filename, self.mode)673 674    def close(self) -> None:675        self.sync()676        self.clear()677 678    def __enter__(self) -> PersistentDict:679        return self680 681    def __exit__(self, *exc_info: Any) -> None:682        self.close()683 684    def dump(self, fileobj: Any) -> None:685        if self.format == "pickle":686            pickle.dump(dict(self), fileobj, 2)687        else:688            raise NotImplementedError("Unknown format: " + repr(self.format))689 690    def load(self) -> None:691        # try formats from most restrictive to least restrictive692        if self.flag == "n":693            return694        with open(self.filename, "rb" if self.format == "pickle" else "r") as fileobj:695            for loader in (pickle.load,):696                fileobj.seek(0)697                try:698                    return self.update(loader(fileobj))699                except EOFError:700                    return701                except Exception:702                    logger.error(f"Failed to load file: {fileobj.name}")703                    raise704            raise ValueError("File not in a supported format")705 
codekingpro/portable-devtools · Team Ai