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