Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_loop.py1955 linesDownload Raw Back to pregel
1from __future__ import annotations2 3import asyncio4import binascii5import concurrent.futures6from collections import defaultdict, deque7from collections.abc import Callable, Iterator, Mapping, Sequence8from contextlib import (9    AbstractAsyncContextManager,10    AbstractContextManager,11    AsyncExitStack,12    ExitStack,13)14from datetime import datetime, timezone15from inspect import signature16from types import TracebackType17from typing import (18    Any,19    Literal,20    TypeVar,21    cast,22)23 24from langchain_core.callbacks import AsyncParentRunManager, ParentRunManager25from langchain_core.runnables import RunnableConfig26from langgraph.cache.base import BaseCache27from langgraph.checkpoint.base import (28    WRITES_IDX_MAP,29    BaseCheckpointSaver,30    ChannelVersions,31    Checkpoint,32    CheckpointMetadata,33    CheckpointTuple,34    PendingWrite,35)36from langgraph.store.base import BaseStore37from typing_extensions import ParamSpec, Self38 39from langgraph._internal._config import patch_configurable40from langgraph._internal._constants import (41    CONF,42    CONFIG_KEY_CHECKPOINT_ID,43    CONFIG_KEY_CHECKPOINT_MAP,44    CONFIG_KEY_CHECKPOINT_NS,45    CONFIG_KEY_REPLAY_STATE,46    CONFIG_KEY_RESUME_MAP,47    CONFIG_KEY_RESUMING,48    CONFIG_KEY_RUNTIME,49    CONFIG_KEY_SCRATCHPAD,50    CONFIG_KEY_STREAM,51    CONFIG_KEY_TASK_ID,52    CONFIG_KEY_THREAD_ID,53    ERROR,54    ERROR_SOURCE_NODE,55    INPUT,56    INTERRUPT,57    NS_END,58    NS_SEP,59    NULL_TASK_ID,60    PUSH,61    RESUME,62    TASKS,63)64from langgraph._internal._replay import ReplayState65from langgraph._internal._scratchpad import PregelScratchpad66from langgraph._internal._typing import EMPTY_SEQ, MISSING67from langgraph.callbacks import (68    GraphInterruptEvent,69    GraphLifecycleEvent,70    GraphResumeEvent,71)72from langgraph.channels.base import BaseChannel73from langgraph.channels.delta import DeltaChannel74from langgraph.channels.untracked_value import UntrackedValue75from langgraph.constants import TAG_HIDDEN76from langgraph.errors import (77    EmptyInputError,78    GraphInterrupt,79)80from langgraph.managed.base import (81    ManagedValueMapping,82    ManagedValueSpec,83)84from langgraph.pregel._algo import (85    Call,86    GetNextVersion,87    PregelTaskWrites,88    apply_writes,89    checkpoint_null_version,90    increment,91    prepare_next_tasks,92    prepare_node_error_handler_task,93    prepare_single_task,94    sanitize_untracked_values_in_send,95    should_interrupt,96    task_path_str,97)98from langgraph.pregel._checkpoint import (99    achannels_from_checkpoint,100    channels_from_checkpoint,101    copy_checkpoint,102    create_checkpoint,103    delta_channels_to_snapshot,104    empty_checkpoint,105)106from langgraph.pregel._executor import (107    AsyncBackgroundExecutor,108    BackgroundExecutor,109    Submit,110)111from langgraph.pregel._io import (112    map_command,113    map_input,114    map_output_updates,115    map_output_values,116    read_channels,117)118from langgraph.pregel._read import PregelNode119from langgraph.pregel._utils import get_new_channel_versions, is_xxh3_128_hexdigest120from langgraph.pregel.debug import (121    map_debug_checkpoint,122    map_debug_task_results,123    map_debug_tasks,124)125from langgraph.pregel.protocol import StreamChunk, StreamProtocol126from langgraph.runtime import RunControl, Runtime127from langgraph.types import (128    All,129    CachePolicy,130    Command,131    Durability,132    Interrupt,133    PregelExecutableTask,134    RetryPolicy,135    Send,136    StreamMode,137)138 139V = TypeVar("V")140P = ParamSpec("P")141 142 143WritesT = Sequence[tuple[str, Any]]144 145 146def DuplexStream(*streams: StreamProtocol) -> StreamProtocol:147    def __call__(value: StreamChunk) -> None:148        for stream in streams:149            if value[1] in stream.modes:150                stream(value)151 152    return StreamProtocol(__call__, {mode for s in streams for mode in s.modes})153 154 155class PregelLoop:156    config: RunnableConfig157    store: BaseStore | None158    stream: StreamProtocol | None159    step: int160    stop: int161 162    input: Any | None163    cache: BaseCache[WritesT] | None164    checkpointer: BaseCheckpointSaver | None165    nodes: Mapping[str, PregelNode]166    specs: Mapping[str, BaseChannel | ManagedValueSpec]167    input_keys: str | Sequence[str]168    output_keys: str | Sequence[str]169    stream_keys: str | Sequence[str]170    is_replaying: bool171    is_nested: bool172    manager: None | AsyncParentRunManager | ParentRunManager173    interrupt_after: All | Sequence[str]174    interrupt_before: All | Sequence[str]175    durability: Durability176    retry_policy: Sequence[RetryPolicy]177    cache_policy: CachePolicy | None178 179    checkpointer_get_next_version: GetNextVersion180    checkpointer_put_writes: Callable[[RunnableConfig, WritesT, str], Any] | None181    checkpointer_put_writes_accepts_task_path: bool182    _checkpointer_put_after_previous: (183        Callable[184            [185                concurrent.futures.Future | None,186                RunnableConfig,187                Checkpoint,188                str,189                ChannelVersions,190            ],191            Any,192        ]193        | None194    )195    _migrate_checkpoint: Callable[[Checkpoint], None] | None196    submit: Submit197    channels: Mapping[str, BaseChannel]198    # Futures from `checkpointer.put_writes` calls that produced delta-channel199    # writes. `_checkpointer_put_after_previous` drains this list (swap to a200    # local `futs` then reset to `[]` and wait/gather) before putting the201    # next checkpoint, so a checkpoint never becomes durable before the202    # writes that produced it. Initialised to `[]` in both sync and async203    # `__enter__`; stays `None` only when no checkpointer.204    _delta_write_futs: list[Any] | None = None205 206    # Same pattern as `_delta_write_futs` but for error-handler writes.207    # When `put_writes` persists an ERROR_SOURCE_NODE marker, the future is208    # appended here.  `schedule_error_handler` / `aschedule_error_handler`209    # drain this list so the write is durable before the handler starts.210    _error_handler_write_futs: list[Any] | None = None211 212    # Exit-mode accumulator: every delta-channel write produced during this213    # run (input writes from `_first` + per-superstep writes captured in214    # `after_tick`). At exit, `_put_exit_delta_writes` filters out channels215    # that will snapshot, then persists the rest under an anchor parent.216    # `None` when not in exit mode (so the capture sites are no-ops).217    # Each tuple is `(step, task_id, channel, value)` — `step` drives the218    # synthetic step-prefixed task_id used to preserve chronological order219    # under the saver's `ORDER BY task_id, idx` sorting.220    _exit_delta_writes: list[tuple[int, str, str, Any]] | None = None221 222    # The checkpoint_config that points at the parent loaded at `__enter__`223    # (or the synthetic-empty checkpoint, on first run). We capture it224    # eagerly because every `_put_checkpoint` advances `self.checkpoint_config`225    # to the newly-saved checkpoint's id — by exit time the original parent226    # config would otherwise be lost. `_put_exit_delta_writes` uses this:227    # on resumed runs as the anchor for exit delta writes; on first runs228    # to derive the lazy stub's config (its `checkpoint_id` is the229    # synthetic-empty id we want the stub persisted under).230    _initial_checkpoint_config: RunnableConfig231 232    # True iff the saver actually returned a tuple at `__enter__`. False233    # on the first-ever run for a thread (no parent persisted yet).234    # `_put_exit_delta_writes` uses this to decide between anchoring on235    # the existing parent (True) or creating a lazy stub (False).236    _has_persisted_parent: bool = False237 238    managed: ManagedValueMapping239    checkpoint: Checkpoint240    checkpoint_id_saved: str241    checkpoint_ns: tuple[str, ...]242    checkpoint_config: RunnableConfig243    checkpoint_metadata: CheckpointMetadata244    checkpoint_pending_writes: list[PendingWrite]245    checkpoint_previous_versions: dict[str, str | float | int]246    prev_checkpoint_config: RunnableConfig | None247 248    status: Literal[249        "input",250        "pending",251        "done",252        "draining",253        "interrupt_before",254        "interrupt_after",255        "out_of_steps",256    ]257    control: RunControl | None258    tasks: dict[str, PregelExecutableTask]259    output: None | dict[str, Any] | Any = None260    updated_channels: set[str] | None = None261    _graph_lifecycle_events: deque[GraphLifecycleEvent]262    _has_graph_lifecycle_callbacks: bool263 264    # public265 266    def __init__(267        self,268        input: Any | None,269        *,270        stream: StreamProtocol | None,271        config: RunnableConfig,272        store: BaseStore | None,273        cache: BaseCache | None,274        checkpointer: BaseCheckpointSaver | None,275        nodes: Mapping[str, PregelNode],276        specs: Mapping[str, BaseChannel | ManagedValueSpec],277        input_keys: str | Sequence[str],278        output_keys: str | Sequence[str],279        stream_keys: str | Sequence[str],280        trigger_to_nodes: Mapping[str, Sequence[str]],281        durability: Durability,282        interrupt_after: All | Sequence[str] = EMPTY_SEQ,283        interrupt_before: All | Sequence[str] = EMPTY_SEQ,284        manager: None | AsyncParentRunManager | ParentRunManager = None,285        migrate_checkpoint: Callable[[Checkpoint], None] | None = None,286        retry_policy: Sequence[RetryPolicy] = (),287        cache_policy: CachePolicy | None = None,288        has_graph_lifecycle_callbacks: bool = False,289    ) -> None:290        self.stream = stream291        self.config = config292        self.store = store293        self.step = 0294        self.stop = 0295        self.input = input296        self.checkpointer = checkpointer297        self.cache = cache298        self.nodes = nodes299        self.specs = specs300        self.input_keys = input_keys301        self.output_keys = output_keys302        self.stream_keys = stream_keys303        self.interrupt_after = interrupt_after304        self.interrupt_before = interrupt_before305        self.manager = manager306        self.is_nested = CONFIG_KEY_TASK_ID in self.config.get(CONF, {})307        self.is_replaying = CONFIG_KEY_CHECKPOINT_ID in config[CONF]308        self._migrate_checkpoint = migrate_checkpoint309        self.trigger_to_nodes = trigger_to_nodes310        self.retry_policy = retry_policy311        self.cache_policy = cache_policy312        self.durability = durability313        self._has_graph_lifecycle_callbacks = has_graph_lifecycle_callbacks314        self._graph_lifecycle_events = deque()315        if self.stream is not None and CONFIG_KEY_STREAM in config[CONF]:316            self.stream = DuplexStream(self.stream, config[CONF][CONFIG_KEY_STREAM])317        scratchpad: PregelScratchpad | None = config[CONF].get(CONFIG_KEY_SCRATCHPAD)318        if isinstance(scratchpad, PregelScratchpad):319            # if count is > 0, append to checkpoint_ns320            # if count is 0, leave as is321            if cnt := scratchpad.subgraph_counter():322                self.config = patch_configurable(323                    self.config,324                    {325                        CONFIG_KEY_CHECKPOINT_NS: NS_SEP.join(326                            (327                                config[CONF][CONFIG_KEY_CHECKPOINT_NS],328                                str(cnt),329                            )330                        )331                    },332                )333        if not self.is_nested and config[CONF].get(CONFIG_KEY_CHECKPOINT_NS):334            self.config = patch_configurable(335                self.config,336                {CONFIG_KEY_CHECKPOINT_NS: "", CONFIG_KEY_CHECKPOINT_ID: None},337            )338        if (339            CONFIG_KEY_CHECKPOINT_MAP in self.config[CONF]340            and self.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS)341            in self.config[CONF][CONFIG_KEY_CHECKPOINT_MAP]342        ):343            self.checkpoint_config = patch_configurable(344                self.config,345                {346                    CONFIG_KEY_CHECKPOINT_ID: self.config[CONF][347                        CONFIG_KEY_CHECKPOINT_MAP348                    ][self.config[CONF][CONFIG_KEY_CHECKPOINT_NS]]349                },350            )351        else:352            self.checkpoint_config = self.config353        if thread_id := self.checkpoint_config[CONF].get(CONFIG_KEY_THREAD_ID):354            if not isinstance(thread_id, str):355                self.checkpoint_config = patch_configurable(356                    self.checkpoint_config,357                    {CONFIG_KEY_THREAD_ID: str(thread_id)},358                )359        self.checkpoint_ns = (360            tuple(cast(str, self.config[CONF][CONFIG_KEY_CHECKPOINT_NS]).split(NS_SEP))361            if self.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS)362            else ()363        )364        self.prev_checkpoint_config = None365        runtime = self.config[CONF].get(CONFIG_KEY_RUNTIME)366        self.control = runtime.control if isinstance(runtime, Runtime) else None367 368    def _push_graph_lifecycle_event(369        self,370        kind: Literal["resume", "interrupt"],371        *,372        interrupts: tuple[Interrupt, ...] = (),373    ) -> None:374        # drain status never reaches lifecycle events: tick() returns False375        # before pushing, and interrupts are raised through GraphInterrupt376        if self.status == "draining":377            raise RuntimeError("Draining status cannot emit lifecycle events")378        status = self.status379        if kind == "resume":380            self._graph_lifecycle_events.append(381                GraphResumeEvent(382                    run_id=None,383                    status=status,384                    checkpoint_id=self.checkpoint["id"],385                    checkpoint_ns=self.checkpoint_ns,386                )387            )388        elif kind == "interrupt":389            self._graph_lifecycle_events.append(390                GraphInterruptEvent(391                    run_id=None,392                    status=status,393                    checkpoint_id=self.checkpoint["id"],394                    checkpoint_ns=self.checkpoint_ns,395                    interrupts=interrupts,396                )397            )398        else:399            msg = f"Unknown graph lifecycle event type: {kind}"400            raise AssertionError(msg)401 402    def _pop_lifecycle_event(self) -> GraphLifecycleEvent | None:403        if not self._graph_lifecycle_events:404            return None405        return self._graph_lifecycle_events.popleft()406 407    def put_writes(self, task_id: str, writes: WritesT) -> None:408        """Put writes for a task, to be read by the next tick."""409        if not writes:410            return411        # deduplicate writes to special channels, last write wins412        if all(w[0] in WRITES_IDX_MAP for w in writes):413            writes = list({w[0]: w for w in writes}.values())414        if task_id == NULL_TASK_ID:415            # writes for the null task are accumulated416            self.checkpoint_pending_writes = [417                w418                for w in self.checkpoint_pending_writes419                if w[0] != task_id or w[1] not in WRITES_IDX_MAP420            ]421            writes_to_save: WritesT = [422                w[1:] for w in self.checkpoint_pending_writes if w[0] == task_id423            ] + list(writes)424        else:425            # remove existing writes for this task426            self.checkpoint_pending_writes = [427                w for w in self.checkpoint_pending_writes if w[0] != task_id428            ]429            writes_to_save = writes430 431        # check if any writes are to an UntrackedValue channel432        if any(433            isinstance(channel, UntrackedValue) for channel in self.channels.values()434        ):435            # we do not persist untracked values in checkpoints436            writes_to_save = [437                # sanitize UntrackedValues that are nested within Send packets438                (439                    (c, sanitize_untracked_values_in_send(v, self.channels))440                    if c == TASKS and isinstance(v, Send)441                    else (c, v)442                )443                for c, v in writes_to_save444                # dont persist UntrackedValue channel writes445                if not isinstance(self.specs.get(c), UntrackedValue)446            ]447 448        # save writes449        self.checkpoint_pending_writes.extend((task_id, c, v) for c, v in writes)450        if self.durability != "exit" and self.checkpointer_put_writes is not None:451            config = patch_configurable(452                self.checkpoint_config,453                {454                    CONFIG_KEY_CHECKPOINT_NS: self.config[CONF].get(455                        CONFIG_KEY_CHECKPOINT_NS, ""456                    ),457                    CONFIG_KEY_CHECKPOINT_ID: self.checkpoint["id"],458                },459            )460            if self.checkpointer_put_writes_accepts_task_path:461                if hasattr(self, "tasks"):462                    task = self.tasks.get(task_id)463                else:464                    task = None465                fut = self.submit(466                    self.checkpointer_put_writes,467                    config,468                    writes_to_save,469                    task_id,470                    task_path_str(task.path) if task else "",471                )472            else:473                fut = self.submit(474                    self.checkpointer_put_writes,475                    config,476                    writes_to_save,477                    task_id,478                )479            if self._delta_write_futs is not None and any(480                isinstance(self.specs.get(c), DeltaChannel) for c, _ in writes_to_save481            ):482                self._delta_write_futs.append(fut)483            # ERROR_SOURCE_NODE is only appended by commit() when the task484            # has an error handler (_should_route_to_error_handler), so this485            # check naturally limits future collection to those tasks.486            if self._error_handler_write_futs is not None and any(487                c == ERROR_SOURCE_NODE for c, _ in writes488            ):489                self._error_handler_write_futs.append(fut)490        # output writes491        if hasattr(self, "tasks"):492            self.output_writes(task_id, writes)493 494    def _put_pending_writes(self) -> None:495        if self.checkpointer_put_writes is None:496            return497        if not self.checkpoint_pending_writes:498            return499        # patch config500        config = patch_configurable(501            self.checkpoint_config,502            {503                CONFIG_KEY_CHECKPOINT_NS: self.config[CONF].get(504                    CONFIG_KEY_CHECKPOINT_NS, ""505                ),506                CONFIG_KEY_CHECKPOINT_ID: self.checkpoint["id"],507            },508        )509        # group by task id510        by_task = defaultdict(list)511        for task_id, channel, value in self.checkpoint_pending_writes:512            by_task[task_id].append((channel, value))513        # submit writes to checkpointer514        for task_id, writes in by_task.items():515            if self.checkpointer_put_writes_accepts_task_path and hasattr(516                self, "tasks"517            ):518                task = self.tasks.get(task_id)519                self.submit(520                    self.checkpointer_put_writes,521                    config,522                    writes,523                    task_id,524                    task_path_str(task.path) if task else "",525                )526            else:527                self.submit(528                    self.checkpointer_put_writes,529                    config,530                    writes,531                    task_id,532                )533 534    def accept_push(535        self, task: PregelExecutableTask, write_idx: int, call: Call | None = None536    ) -> PregelExecutableTask | None:537        """Accept a PUSH from a task, potentially returning a new task to start."""538        checkpoint_id_bytes = binascii.unhexlify(self.checkpoint["id"].replace("-", ""))539        null_version = checkpoint_null_version(self.checkpoint)540        if pushed := cast(541            PregelExecutableTask | None,542            prepare_single_task(543                (PUSH, task.path, write_idx, task.id, call),544                None,545                checkpoint=self.checkpoint,546                checkpoint_id_bytes=checkpoint_id_bytes,547                checkpoint_null_version=null_version,548                pending_writes=self.checkpoint_pending_writes,549                processes=self.nodes,550                channels=self.channels,551                managed=self.managed,552                config=task.config,553                step=self.step,554                stop=self.stop,555                for_execution=True,556                store=self.store,557                checkpointer=self.checkpointer,558                manager=self.manager,559                retry_policy=self.retry_policy,560                cache_policy=self.cache_policy,561            ),562        ):563            # produce debug output564            self._emit("tasks", map_debug_tasks, [pushed])565            # save the new task566            self.tasks[pushed.id] = pushed567            # match any pending writes to the new task568            if not self.is_replaying:569                self._reapply_writes_to_succeeded_nodes({pushed.id: pushed})570            # return the new task, to be started if not run before571            return pushed572 573    def schedule_error_handler(574        self, failed_task: PregelExecutableTask, error: BaseException575    ) -> PregelExecutableTask | None:576        raise NotImplementedError577 578    async def aschedule_error_handler(579        self, failed_task: PregelExecutableTask, error: BaseException580    ) -> PregelExecutableTask | None:581        raise NotImplementedError582 583    def tick(self) -> bool:584        """Execute a single iteration of the Pregel loop.585 586        Returns:587            True if more iterations are needed.588        """589 590        # check if iteration limit is reached591        if self.step > self.stop:592            self.status = "out_of_steps"593            return False594 595        # prepare next tasks596        self.tasks = prepare_next_tasks(597            self.checkpoint,598            self.checkpoint_pending_writes,599            self.nodes,600            self.channels,601            self.managed,602            self.config,603            self.step,604            self.stop,605            for_execution=True,606            manager=self.manager,607            store=self.store,608            checkpointer=self.checkpointer,609            trigger_to_nodes=self.trigger_to_nodes,610            updated_channels=self.updated_channels,611            retry_policy=self.retry_policy,612            cache_policy=self.cache_policy,613        )614 615        # produce debug output616        if self._checkpointer_put_after_previous is not None:617            self._emit(618                "checkpoints",619                map_debug_checkpoint,620                {621                    **self.checkpoint_config,622                    CONF: {623                        **self.checkpoint_config[CONF],624                        CONFIG_KEY_CHECKPOINT_ID: self.checkpoint["id"],625                    },626                },627                self.channels,628                self.stream_keys,629                self.checkpoint_metadata,630                self.tasks.values(),631                self.checkpoint_pending_writes,632                self.prev_checkpoint_config,633                self.output_keys,634            )635 636        # if no more tasks, we're done637        if not self.tasks:638            self.status = "done"639            return False640 641        if self.control is not None and self.control.drain_requested:642            self.status = "draining"643            return False644 645        # if there are pending writes from a previous loop, apply them646        if not self.is_replaying and self.checkpoint_pending_writes:647            self._reapply_writes_to_succeeded_nodes(self.tasks)648            self._resume_error_handlers_if_applicable()649 650        # before execution, check if we should interrupt651        if self.interrupt_before and should_interrupt(652            self.checkpoint, self.interrupt_before, self.tasks.values()653        ):654            self.status = "interrupt_before"655            raise GraphInterrupt()656 657        # produce debug output658        self._emit("tasks", map_debug_tasks, self.tasks.values())659 660        # print output for any tasks we applied previous writes to661        for task in self.tasks.values():662            if task.writes:663                self.output_writes(task.id, task.writes, cached=True)664 665        return True666 667    def after_tick(self) -> None:668        # finish superstep669        writes = [w for t in self.tasks.values() for w in t.writes]670        # all tasks have finished671        self.updated_channels = apply_writes(672            self.checkpoint,673            self.channels,674            self.tasks.values(),675            self.checkpointer_get_next_version,676            self.trigger_to_nodes,677        )678        # produce values output679        if not self.updated_channels.isdisjoint(680            (self.output_keys,)681            if isinstance(self.output_keys, str)682            else self.output_keys683        ):684            self._emit(685                "values", map_output_values, self.output_keys, writes, self.channels686            )687        # capture delta-channel writes for exit-mode accumulator before clearing688        if self._exit_delta_writes is not None:689            for tid, ch, v in self.checkpoint_pending_writes:690                if isinstance(self.specs.get(ch), DeltaChannel):691                    self._exit_delta_writes.append((self.step, tid, ch, v))692        # clear pending writes693        self.checkpoint_pending_writes.clear()694        # only replay (re-execute) done tasks on the first tick695        self.is_replaying = False696        # save checkpoint697        self._put_checkpoint({"source": "loop"})698        # after execution, check if we should interrupt699        if self.interrupt_after and should_interrupt(700            self.checkpoint, self.interrupt_after, self.tasks.values()701        ):702            self.status = "interrupt_after"703            raise GraphInterrupt()704        # unset resuming flag705        self.config[CONF].pop(CONFIG_KEY_RESUMING, None)706 707    def match_cached_writes(self) -> Sequence[PregelExecutableTask]:708        raise NotImplementedError709 710    async def amatch_cached_writes(self) -> Sequence[PregelExecutableTask]:711        raise NotImplementedError712 713    # private714 715    def _reapply_writes_to_succeeded_nodes(716        self, tasks: Mapping[str, PregelExecutableTask]717    ) -> None:718        """Restore successful channel writes from checkpoint to in-memory tasks.719 720        Skips control signals (ERROR, ERROR_SOURCE_NODE, INTERRUPT, RESUME)721        so that failed/interrupted tasks remain with empty writes and will be722        re-executed (or routed to error handlers) by the runner.723        """724        for tid, k, v in self.checkpoint_pending_writes:725            if k in (ERROR, ERROR_SOURCE_NODE, INTERRUPT, RESUME):726                continue727            if task := tasks.get(tid):728                task.writes.append((k, v))729 730    def _resume_error_handlers_if_applicable(self) -> None:731        """On resume, schedule error handlers for tasks that failed in a prior run.732 733        Called right after ``_reapply_writes_to_succeeded_nodes`` during ``tick()``.734        At that point, ``_reapply_writes_to_succeeded_nodes`` has already skipped735        ERROR / ERROR_SOURCE_NODE writes, so a previously-failed task still has736        empty ``writes``.  Without intervention the runner (which executes only737        tasks where ``not t.writes``) would re-run the original node.738 739        This method prevents that re-execution for nodes that have an error740        handler:741 742        1. Scan ``checkpoint_pending_writes`` for ERROR_SOURCE_NODE markers743           persisted by a prior ``commit()``.  Each marker means "this task744           already failed and was routed to an error handler".745        2. For each such task, write ``(ERROR, error)`` into ``task.writes``746           so the task is no longer empty — the runner will skip it.747        3. Prepare a fresh error-handler task and add it to ``self.tasks``.748           Because the handler task starts with empty ``writes``, the runner749           will pick it up and execute it.750        """751        # Phase 1: collect task-ids that have ERROR_SOURCE_NODE + ERROR pairs.752        failed: dict[str, BaseException] = {}753        for tid, chan, val in self.checkpoint_pending_writes:754            if chan == ERROR_SOURCE_NODE:755                error = next(756                    (757                        v758                        for t, c, v in self.checkpoint_pending_writes759                        if t == tid and c == ERROR760                    ),761                    None,762                )763                if error is not None:764                    failed[tid] = error765        # Phase 2: mark originals as done, schedule handler tasks.766        for task_id, error in failed.items():767            task = self.tasks.get(task_id)768            if task is None:769                continue770            handler_node = self.nodes[task.name].error_handler_node771            if not handler_node:772                continue773            # Non-empty writes → runner's `not t.writes` filter skips this task.774            task.writes.append((ERROR, error))775            # The handler task starts with empty writes → runner will execute it.776            handler_task = prepare_node_error_handler_task(777                task,778                handler_node_name=handler_node,779                failed_error=error,780                checkpoint=self.checkpoint,781                pending_writes=self.checkpoint_pending_writes,782                processes=self.nodes,783                channels=self.channels,784                managed=self.managed,785                config=task.config,786                step=self.step,787                stop=self.stop,788                store=self.store,789                checkpointer=self.checkpointer,790                manager=self.manager,791                retry_policy=self.retry_policy,792                cache_policy=self.cache_policy,793            )794            if handler_task is not None:795                self.tasks[handler_task.id] = handler_task796 797    def _pending_interrupts(self) -> set[str]:798        """Return the set of interrupt ids that are pending without corresponding resume values."""799        # mapping of task ids to interrupt ids800        pending_interrupts: dict[str, str] = {}801 802        # set of resume task ids803        pending_resumes: set[str] = set()804 805        for task_id, write_type, value in self.checkpoint_pending_writes:806            if write_type == INTERRUPT:807                # interrupts is always a list, but there should only be one element808                pending_interrupts[task_id] = value[0].id809            elif write_type == RESUME:810                pending_resumes.add(task_id)811 812        resumed_interrupt_ids = {813            pending_interrupts[task_id]814            for task_id in pending_resumes815            if task_id in pending_interrupts816        }817 818        # Keep only interrupts whose interrupt_id is not resumed819        hanging_interrupts: set[str] = {820            interrupt_id821            for interrupt_id in pending_interrupts.values()822            if interrupt_id not in resumed_interrupt_ids823        }824 825        return hanging_interrupts826 827    def _first(828        self, *, input_keys: str | Sequence[str], updated_channels: set[str] | None829    ) -> set[str] | None:830        # Resuming from a previous checkpoint requires two things:831        # 1. A prior checkpoint exists (channel_versions is non-empty)832        # 2. The input signals continuation (not a fresh run with new input)833        # For subgraphs, the parent explicitly sets CONFIG_KEY_RESUMING.834        # For the outer graph, we infer from the input:835        #   - None input: resume after interrupt (invoke(None, config))836        #   - Command input: any Command operates on existing state837        #   - Same run_id: re-entry into an ongoing run (e.g. stream reconnect)838        configurable = self.config.get(CONF, {})839        input_is_command = isinstance(self.input, Command)840        is_resuming = bool(self.checkpoint["channel_versions"]) and bool(841            configurable.get(842                CONFIG_KEY_RESUMING,843                self.input is None844                or input_is_command845                or (846                    not self.is_nested847                    and self.config.get("metadata", {}).get("run_id")848                    == self.checkpoint_metadata.get("run_id", MISSING)849                ),850            )851        )852 853        # When replaying from a specific checkpoint, drop cached RESUME854        # writes so that interrupt() calls re-fire instead of returning855        # stale values. But if we're actively resuming, keep them —856        # multi-interrupt scenarios need previously resolved values preserved.857        is_time_traveling = self.is_replaying and (858            # Time-travel to a subgraph checkpoint: the parent sets859            # RESUMING=True (it can't distinguish time-travel from resume),860            # so we check if this subgraph's own ns is in checkpoint_map.861            # Normally the map only has ancestor entries (_algo.py); the862            # subgraph's own entry only appears via get_state(subgraphs=True).863            (864                self.is_nested865                and configurable.get(CONFIG_KEY_CHECKPOINT_NS, "")866                in configurable.get(CONFIG_KEY_CHECKPOINT_MAP, {})867            )868            or not (869                # Outer graph: resume arrives as Command(resume=...)870                (input_is_command and cast(Command, self.input).resume is not None)871                # Subgraphs: resume arrives via config flag from parent872                # (subgraph input is a Send arg, not a Command)873                or configurable.get(CONFIG_KEY_RESUMING, False)874            )875        )876        if is_time_traveling:877            self.checkpoint_pending_writes = [878                w for w in self.checkpoint_pending_writes if w[1] != RESUME879            ]880 881        # map command to writes882        if input_is_command:883            if (resume := cast(Command, self.input).resume) is not None:884                if not self.checkpointer:885                    raise RuntimeError(886                        "Cannot use Command(resume=...) without checkpointer"887                    )888 889                if resume_is_map := (890                    isinstance(resume, dict)891                    and all(is_xxh3_128_hexdigest(k) for k in resume)892                ):893                    self.config[CONF][CONFIG_KEY_RESUME_MAP] = resume894                else:895                    if len(self._pending_interrupts()) > 1:896                        raise RuntimeError(897                            "When there are multiple pending interrupts, you must specify the interrupt id when resuming. "898                            "Docs: https://docs.langchain.com/oss/python/langgraph/add-human-in-the-loop#resume-multiple-interrupts-with-one-invocation."899                        )900 901            writes: defaultdict[str, list[tuple[str, Any]]] = defaultdict(list)902            # group writes by task ID903            for tid, c, v in map_command(cmd=cast(Command, self.input)):904                if not (c == RESUME and resume_is_map):905                    writes[tid].append((c, v))906            if not writes and not resume_is_map:907                raise EmptyInputError("Received empty Command input")908            # save writes909            for tid, ws in writes.items():910                self.put_writes(tid, ws)911        # apply NULL writes912        if null_writes := [913            w[1:] for w in self.checkpoint_pending_writes if w[0] == NULL_TASK_ID914        ]:915            null_updated_channels = apply_writes(916                self.checkpoint,917                self.channels,918                [PregelTaskWrites((), INPUT, null_writes, [])],919                self.checkpointer_get_next_version,920                self.trigger_to_nodes,921            )922            if updated_channels is not None:923                updated_channels.update(null_updated_channels)924        # proceed past previous checkpoint925        if is_resuming:926            self.checkpoint["versions_seen"].setdefault(INTERRUPT, {})927            for k in self.channels:928                if k in self.checkpoint["channel_versions"]:929                    version = self.checkpoint["channel_versions"][k]930                    self.checkpoint["versions_seen"][INTERRUPT][k] = version931            # When time-traveling (replaying from a specific checkpoint),932            # save a fork checkpoint so the replayed execution creates a933            # new branch. Without this, if the execution hits an interrupt934            # before after_tick() runs, no new checkpoint is created —935            # the parent's latest checkpoint remains the old one and936            # subsequent resumes load the wrong state.937            # Skip for update_state forks (source=update/fork) since they938            # already have their own fork checkpoint.939            if is_time_traveling and self.checkpoint_metadata.get("source") not in (940                "update",941                "fork",942            ):943                # Clear old INTERRUPT writes from the loaded checkpoint.944                # The fork will have a new checkpoint_id which changes945                # task IDs — stale interrupt writes would accumulate and946                # confuse the multiple-interrupt check in future resumes.947                self.checkpoint_pending_writes = [948                    w for w in self.checkpoint_pending_writes if w[1] != INTERRUPT949                ]950                self._put_checkpoint({"source": "fork"})951            # produce values output952            self._emit(953                "values", map_output_values, self.output_keys, True, self.channels954            )955        # map inputs to channel updates956        elif input_writes := deque(map_input(input_keys, self.input)):957            # discard any unfinished tasks from previous checkpoint958            discard_tasks = prepare_next_tasks(959                self.checkpoint,960                self.checkpoint_pending_writes,961                self.nodes,962                self.channels,963                self.managed,964                self.config,965                self.step,966                self.stop,967                for_execution=True,968                store=None,969                checkpointer=None,970                manager=None,971                updated_channels=updated_channels,972            )973            # apply input writes974            updated_channels = apply_writes(975                self.checkpoint,976                self.channels,977                [978                    *discard_tasks.values(),979                    PregelTaskWrites((), INPUT, input_writes, []),980                ],981                self.checkpointer_get_next_version,982                self.trigger_to_nodes,983            )984            # Input writes go through `apply_writes` directly (above) — they985            # never enter `checkpoint_pending_writes`, so the after_tick986            # capture site does not see them. In exit mode, capture them987            # here so `_exit_delta_writes` includes the input's delta writes988            # alongside per-superstep writes; otherwise the input would be989            # lost on read (it's not in final_checkpoint.channel_values for990            # sub-freq channels, and walks ignore target.pending_writes).991            if self._exit_delta_writes is not None:992                for c, v in input_writes:993                    if isinstance(self.specs.get(c), DeltaChannel):994                        self._exit_delta_writes.append((self.step, NULL_TASK_ID, c, v))995            # Persist delta-channel input writes so sub-freq inputs are996            # recoverable via ancestor walk (mirrors the Command input path).997            if self.durability != "exit":998                delta_input = [999                    (c, v)1000                    for c, v in input_writes1001                    if isinstance(self.specs.get(c), DeltaChannel)1002                ]1003                if delta_input:1004                    self.put_writes(NULL_TASK_ID, delta_input)1005            # save input checkpoint1006            self.updated_channels = updated_channels1007            self._put_checkpoint({"source": "input"})1008        elif CONFIG_KEY_RESUMING not in configurable:1009            raise EmptyInputError(f"Received no input for {input_keys}")1010        # Propagate resuming and replaying flags to subgraphs.1011        if not self.is_nested:1012            # Pass the resolved before-bound checkpoint ID so subgraphs can1013            # find their corresponding checkpoint without re-fetching the1014            # parent. For forks (source=update/fork), use the fork's parent1015            # checkpoint ID since the fork was created after the subgraph's1016            # checkpoints from the original execution.1017            #1018            # Only gate on is_time_traveling (not is_replaying). When the1019            # client resumes with an explicit checkpoint_id that happens to1020            # point at the current head (e.g. LangGraph Studio sending1021            # `checkpoint: {checkpoint_id}` alongside Command(resume=...)),1022            # is_replaying is True but is_time_traveling is False. In that1023            # case subgraphs should load their latest checkpoint normally,1024            # not go through ReplayState's before-bound lookup which would1025            # miss subgraph checkpoints created during processing of the1026            # current parent step.1027            replay_state: ReplayState | None = None1028            if is_time_traveling:1029                replay_checkpoint_id = self.checkpoint["id"]1030                if (1031                    self.checkpoint_metadata.get("source")1032                    in (1033                        "update",1034                        "fork",1035                    )1036                    and self.prev_checkpoint_config1037                ):1038                    replay_checkpoint_id = self.prev_checkpoint_config[CONF].get(1039                        CONFIG_KEY_CHECKPOINT_ID, replay_checkpoint_id1040                    )1041                replay_state = ReplayState(replay_checkpoint_id)1042            self.config = patch_configurable(1043                self.config,1044                {1045                    CONFIG_KEY_RESUMING: is_resuming,1046                    CONFIG_KEY_REPLAY_STATE: replay_state,1047                },1048            )1049        # set flag1050        self.status = "pending"1051        if is_resuming:1052            self._push_graph_lifecycle_event("resume")1053        return updated_channels1054 1055    def _put_checkpoint(self, metadata: CheckpointMetadata) -> None:1056        # `is` (object identity) — not `==`. Three of four call sites pass a1057        # fresh dict ({"source":"input"|"loop"|"fork"}); only1058        # `_suppress_interrupt`(will rename to _on_loop_exit soon)1059        # at exit reuses the existing `self.checkpoint_metadata` instance. So1060        # `metadata is self.checkpoint_metadata` is True only on the exit call,1061        # which is what we use to gate exit-only behaviour (skip count-bump,1062        # don't replace metadata). Could be replaced by an explicit1063        # `exiting: bool = False` parameter; left as-is to match the existing1064        # idiom in this file.1065        # TODO: replace with an explicit `exiting: bool = False` parameter.1066        exiting = metadata is self.checkpoint_metadata1067        if exiting and self.checkpoint["id"] == self.checkpoint_id_saved:1068            # checkpoint already saved1069            return1070        # Per-delta-channel counter bookkeeping.1071        #1072        # Each delta channel tracks a (updates, supersteps) tuple:1073        # - `updates` increments only when the channel is written this step.1074        # - `supersteps` increments every superstep regardless.1075        #1076        # `_put_checkpoint` is called once per superstep with a fresh1077        # metadata dict (source="input"|"loop"|"fork") — those are the1078        # intermediate calls that bump counters. In exit mode,1079        # `_suppress_interrupt`(will rename to _on_loop_exit soon)1080        # additionally calls `_put_checkpoint(self.checkpoint_metadata)` AT1081        # EXIT to commit the final checkpoint — this runs *after* the last1082        # intermediate call already counted the last superstep. So the1083        # exit call must NOT bump again or it would double-count the last1084        # superstep.1085        if not exiting:1086            prev_counters = dict(1087                self.checkpoint_metadata.get("counters_since_delta_snapshot") or {}1088            )1089            new_counters: dict[str, tuple[int, int]] = {}1090            updated = self.updated_channels or set()1091            for ch_name, ch in self.channels.items():1092                if not isinstance(ch, DeltaChannel):1093                    continue1094                u, s = prev_counters.get(ch_name, (0, 0))1095                s += 11096                if ch_name in updated:1097                    u += 11098                new_counters[ch_name] = (u, s)1099            metadata["step"] = self.step1100            metadata["parents"] = self.config[CONF].get(CONFIG_KEY_CHECKPOINT_MAP, {})1101            self.checkpoint_metadata = metadata1102        else:1103            new_counters = dict(1104                self.checkpoint_metadata.get("counters_since_delta_snapshot") or {}1105            )1106        # do checkpoint?1107        do_checkpoint = self._checkpointer_put_after_previous is not None and (1108            exiting or self.durability != "exit"1109        )1110        # create new checkpoint1111        channels_to_snapshot = (1112            delta_channels_to_snapshot(self.channels, new_counters)1113            if do_checkpoint1114            else set()1115        )1116        self.checkpoint = create_checkpoint(1117            self.checkpoint,1118            self.channels if do_checkpoint else None,1119            self.step,1120            id=self.checkpoint["id"] if exiting else None,1121            updated_channels=self.updated_channels,1122            get_next_version=self.checkpointer_get_next_version1123            if do_checkpoint1124            else None,1125            channels_to_snapshot=channels_to_snapshot,1126        )1127        for k in channels_to_snapshot:1128            new_counters[k] = (0, 0)1129        non_zero = {k: v for k, v in new_counters.items() if v != (0, 0)}1130        if non_zero:1131            self.checkpoint_metadata["counters_since_delta_snapshot"] = non_zero1132        elif "counters_since_delta_snapshot" in self.checkpoint_metadata:1133            del self.checkpoint_metadata["counters_since_delta_snapshot"]1134        # sanitize TASK channel in the checkpoint before saving (durability=="exit")1135        if TASKS in self.checkpoint["channel_values"] and any(1136            isinstance(channel, UntrackedValue) for channel in self.channels.values()1137        ):1138            sanitized_tasks = [1139                sanitize_untracked_values_in_send(value, self.channels)1140                if isinstance(value, Send)1141                else value1142                for value in self.checkpoint["channel_values"][TASKS]1143            ]1144            self.checkpoint["channel_values"][TASKS] = sanitized_tasks1145        # bail if no checkpointer1146 1147        if do_checkpoint and self._checkpointer_put_after_previous is not None:1148            self.prev_checkpoint_config = (1149                self.checkpoint_config1150                if CONFIG_KEY_CHECKPOINT_ID in self.checkpoint_config[CONF]1151                and self.checkpoint_config[CONF][CONFIG_KEY_CHECKPOINT_ID]1152                else None1153            )1154            self.checkpoint_config = {1155                **self.checkpoint_config,1156                CONF: {1157                    **self.checkpoint_config[CONF],1158                    CONFIG_KEY_CHECKPOINT_NS: self.config[CONF].get(1159                        CONFIG_KEY_CHECKPOINT_NS, ""1160                    ),1161                },1162            }1163 1164            channel_versions = self.checkpoint["channel_versions"].copy()1165            new_versions = get_new_channel_versions(1166                self.checkpoint_previous_versions, channel_versions1167            )1168            self.checkpoint_previous_versions = channel_versions1169 1170            # save it, without blocking1171            # if there's a previous checkpoint save in progress, wait for it1172            # ensuring checkpointers receive checkpoints in order1173            self._put_checkpoint_fut = self.submit(1174                self._checkpointer_put_after_previous,1175                getattr(self, "_put_checkpoint_fut", None),1176                self.checkpoint_config,1177                copy_checkpoint(self.checkpoint),1178                self.checkpoint_metadata,1179                new_versions,1180            )1181            self.checkpoint_config = {1182                **self.checkpoint_config,1183                CONF: {1184                    **self.checkpoint_config[CONF],1185                    CONFIG_KEY_CHECKPOINT_ID: self.checkpoint["id"],1186                },1187            }1188        if not exiting:1189            # increment step1190            self.step += 11191 1192    def _put_exit_delta_writes(self) -> None:1193        """Stage stub + accumulated delta writes so final_checkpoint's put1194        waits on them (visibility invariant: both must be durable before1195        final_checkpoint becomes visible to readers).1196 1197        Stub is created lazily — only when no persisted parent exists AND at1198        least one delta channel has writes that won't be snapshotted.1199        """1200        if (

Showing the first 1,200 of 1955 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai