diff --git a/libs/langgraph/langgraph/constants.py b/libs/langgraph/langgraph/constants.py index 0a748a032..64d4644f6 100644 --- a/libs/langgraph/langgraph/constants.py +++ b/libs/langgraph/langgraph/constants.py @@ -14,11 +14,13 @@ CONFIG_KEY_CHECKPOINT_MAP = "checkpoint_map" INTERRUPT = "__interrupt__" ERROR = "__error__" TASKS = "__pregel_tasks" +SUBSCRIPTIONS = "__pregel_subscriptions" RUNTIME_PLACEHOLDER = "__pregel_runtime_placeholder__" RESERVED = { INTERRUPT, ERROR, TASKS, + SUBSCRIPTIONS, CONFIG_KEY_SEND, CONFIG_KEY_READ, CONFIG_KEY_CHECKPOINTER, diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index 459f67b1a..7626c448f 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -30,6 +30,7 @@ from langgraph.constants import ( INTERRUPT, NS_SEP, RESERVED, + SUBSCRIPTIONS, TAG_HIDDEN, TASKS, Send, @@ -275,27 +276,85 @@ def prepare_next_tasks( checkpointer: Optional[BaseCheckpointSaver] = None, manager: Union[None, ParentRunManager, AsyncParentRunManager] = None, ) -> Union[list[PregelTask], list[PregelExecutableTask]]: + tasks: Union[list[PregelTask], list[PregelExecutableTask]] = [] + # Consume pending packets + for idx, _ in enumerate(checkpoint["pending_sends"]): + if task := prepare_single_task( + (TASKS, idx), + None, + checkpoint=checkpoint, + processes=processes, + channels=channels, + managed=managed, + config=config, + step=step, + for_execution=for_execution, + is_resuming=is_resuming, + checkpointer=checkpointer, + manager=manager, + current_len=len(tasks), + ): + tasks.append(task) + # Check if any processes should be run in next step + # If so, prepare the values to be passed to them + for name in processes: + if task := prepare_single_task( + (SUBSCRIPTIONS, name), + None, + checkpoint=checkpoint, + processes=processes, + channels=channels, + managed=managed, + config=config, + step=step, + for_execution=for_execution, + is_resuming=is_resuming, + checkpointer=checkpointer, + manager=manager, + current_len=len(tasks), + ): + tasks.append(task) + return tasks + + +def prepare_single_task( + task_path: tuple[str, Union[int, str]], + task_id_checksum: Optional[str], + *, + checkpoint: Checkpoint, + processes: Mapping[str, PregelNode], + channels: Mapping[str, BaseChannel], + managed: ManagedValueMapping, + config: RunnableConfig, + step: int, + for_execution: bool, + is_resuming: bool = False, + checkpointer: Optional[BaseCheckpointSaver] = None, + manager: Union[None, ParentRunManager, AsyncParentRunManager] = None, + current_len: int, +) -> Union[None, PregelTask, PregelExecutableTask]: checkpoint_id = UUID(checkpoint["id"]) configurable = config.get("configurable", {}) parent_ns = configurable.get("checkpoint_ns", "") - tasks: Union[list[PregelTask], list[PregelExecutableTask]] = [] - # Consume pending packets - for packet in checkpoint["pending_sends"]: + + if task_path[0] == TASKS: + idx = int(task_path[1]) + packet = checkpoint["pending_sends"][idx] if not isinstance(packet, Send): logger.warning( f"Ignoring invalid packet type {type(packet)} in pending sends" ) - continue + return if packet.node not in processes: logger.warning(f"Ignoring unknown node name {packet.node} in pending sends") - continue + return # create task id triggers = [TASKS] metadata = { "langgraph_step": step, "langgraph_node": packet.node, "langgraph_triggers": triggers, - "langgraph_task_idx": len(tasks), + "langgraph_path": task_path, } checkpoint_ns = ( f"{parent_ns}{NS_SEP}{packet.node}" if parent_ns else packet.node @@ -304,7 +363,17 @@ def prepare_next_tasks( uuid5( checkpoint_id, "".join( - (checkpoint_ns, str(step), packet.node, *triggers, str(len(tasks))) + (checkpoint_ns, str(step), packet.node, *triggers, TASKS, str(idx)) + ), + ) + ) + if task_id_checksum is not None: + assert task_id == task_id_checksum + legacy_id = str( + uuid5( + checkpoint_id, + "".join( + (checkpoint_ns, str(step), packet.node, *triggers, str(current_len)) ), ) ) @@ -314,19 +383,133 @@ def prepare_next_tasks( managed.replace_runtime_placeholders(step, packet.arg) writes = deque() task_checkpoint_ns = f"{checkpoint_ns}:{task_id}" - tasks.append( - PregelExecutableTask( - packet.node, - packet.arg, + return PregelExecutableTask( + packet.node, + packet.arg, + node, + writes, + patch_config( + merge_configs( + config, + processes[packet.node].config, + {"metadata": metadata}, + ), + run_name=packet.node, + callbacks=( + manager.get_child(f"graph:step:{step}") if manager else None + ), + configurable={ + CONFIG_KEY_TASK_ID: task_id, + # deque.extend is thread-safe + CONFIG_KEY_SEND: partial( + local_write, + step, + writes.extend, + processes, + channels, + managed, + ), + CONFIG_KEY_READ: partial( + local_read, + step, + checkpoint, + channels, + managed, + PregelTaskWrites(packet.node, writes, triggers), + config, + ), + CONFIG_KEY_CHECKPOINTER: ( + checkpointer + or configurable.get(CONFIG_KEY_CHECKPOINTER) + ), + CONFIG_KEY_CHECKPOINT_MAP: { + **configurable.get(CONFIG_KEY_CHECKPOINT_MAP, {}), + parent_ns: checkpoint["id"], + }, + CONFIG_KEY_RESUMING: is_resuming, + "checkpoint_id": None, + "checkpoint_ns": task_checkpoint_ns, + }, + ), + triggers, + proc.retry_policy, + None, + task_id, + legacy_id, + ) + + else: + return PregelTask(task_id, packet.node) + elif task_path[0] == SUBSCRIPTIONS: + name = str(task_path[1]) + proc = processes[name] + version_type = type(next(iter(checkpoint["channel_versions"].values()), None)) + null_version = version_type() + if null_version is None: + return + seen = checkpoint["versions_seen"].get(name, {}) + # If any of the channels read by this process were updated + if triggers := sorted( + chan + for chan in proc.triggers + if not isinstance( + read_channel(channels, chan, return_exception=True), EmptyChannelError + ) + and checkpoint["channel_versions"].get(chan, null_version) + > seen.get(chan, null_version) + ): + try: + val = next( + _proc_input( + step, proc, managed, channels, for_execution=for_execution + ) + ) + except StopIteration: + return + + # create task id + metadata = { + "langgraph_step": step, + "langgraph_node": name, + "langgraph_triggers": triggers, + "langgraph_path": task_path, + } + checkpoint_ns = f"{parent_ns}{NS_SEP}{name}" if parent_ns else name + task_id = str( + uuid5( + checkpoint_id, + "".join( + (checkpoint_ns, str(step), name, *triggers, SUBSCRIPTIONS, name) + ), + ) + ) + if task_id_checksum is not None: + assert task_id == task_id_checksum + legacy_id = str( + uuid5( + checkpoint_id, + "".join( + (checkpoint_ns, str(step), name, *triggers, str(current_len)) + ), + ) + ) + + if for_execution: + if node := proc.node: + writes = deque() + task_checkpoint_ns = f"{checkpoint_ns}:{task_id}" + return PregelExecutableTask( + name, + val, node, writes, patch_config( merge_configs( config, - processes[packet.node].config, + proc.config, {"metadata": metadata}, ), - run_name=packet.node, + run_name=name, callbacks=( manager.get_child(f"graph:step:{step}") if manager @@ -349,7 +532,7 @@ def prepare_next_tasks( checkpoint, channels, managed, - PregelTaskWrites(packet.node, writes, triggers), + PregelTaskWrites(name, writes, triggers), config, ), CONFIG_KEY_CHECKPOINTER: ( @@ -361,7 +544,6 @@ def prepare_next_tasks( parent_ns: checkpoint["id"], }, CONFIG_KEY_RESUMING: is_resuming, - "checkpoint_id": None, "checkpoint_ns": task_checkpoint_ns, }, ), @@ -369,119 +551,10 @@ def prepare_next_tasks( proc.retry_policy, None, task_id, - ) - ) - else: - tasks.append(PregelTask(task_id, packet.node)) - # Check if any processes should be run in next step - # If so, prepare the values to be passed to them - version_type = type(next(iter(checkpoint["channel_versions"].values()), None)) - null_version = version_type() - if null_version is None: - return tasks - for name, proc in processes.items(): - seen = checkpoint["versions_seen"].get(name, {}) - # If any of the channels read by this process were updated - if triggers := sorted( - chan - for chan in proc.triggers - if not isinstance( - read_channel(channels, chan, return_exception=True), EmptyChannelError - ) - and checkpoint["channel_versions"].get(chan, null_version) - > seen.get(chan, null_version) - ): - try: - val = next( - _proc_input( - step, proc, managed, channels, for_execution=for_execution - ) - ) - except StopIteration: - continue - - # create task id - metadata = { - "langgraph_step": step, - "langgraph_node": name, - "langgraph_triggers": triggers, - "langgraph_task_idx": len(tasks), - } - checkpoint_ns = f"{parent_ns}{NS_SEP}{name}" if parent_ns else name - task_id = str( - uuid5( - checkpoint_id, - "".join( - (checkpoint_ns, str(step), name, *triggers, str(len(tasks))) - ), - ) - ) - - if for_execution: - if node := proc.node: - writes = deque() - task_checkpoint_ns = f"{checkpoint_ns}:{task_id}" - tasks.append( - PregelExecutableTask( - name, - val, - node, - writes, - patch_config( - merge_configs( - config, - proc.config, - {"metadata": metadata}, - ), - run_name=name, - callbacks=( - manager.get_child(f"graph:step:{step}") - if manager - else None - ), - configurable={ - CONFIG_KEY_TASK_ID: task_id, - # deque.extend is thread-safe - CONFIG_KEY_SEND: partial( - local_write, - step, - writes.extend, - processes, - channels, - managed, - ), - CONFIG_KEY_READ: partial( - local_read, - step, - checkpoint, - channels, - managed, - PregelTaskWrites(name, writes, triggers), - config, - ), - CONFIG_KEY_CHECKPOINTER: ( - checkpointer - or configurable.get(CONFIG_KEY_CHECKPOINTER) - ), - CONFIG_KEY_CHECKPOINT_MAP: { - **configurable.get( - CONFIG_KEY_CHECKPOINT_MAP, {} - ), - parent_ns: checkpoint["id"], - }, - CONFIG_KEY_RESUMING: is_resuming, - "checkpoint_ns": task_checkpoint_ns, - }, - ), - triggers, - proc.retry_policy, - None, - task_id, - ) + legacy_id, ) else: - tasks.append(PregelTask(task_id, name)) - return tasks + return PregelTask(task_id, name) def _proc_input( diff --git a/libs/langgraph/langgraph/pregel/types.py b/libs/langgraph/langgraph/pregel/types.py index 6da463286..4c482f0ee 100644 --- a/libs/langgraph/langgraph/pregel/types.py +++ b/libs/langgraph/langgraph/pregel/types.py @@ -81,6 +81,7 @@ class PregelExecutableTask(NamedTuple): retry_policy: Optional[RetryPolicy] cache_policy: Optional[CachePolicy] id: str + legacy_id: str = "" class StateSnapshot(NamedTuple):