diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 687f59161..b79bb00d5 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -1404,97 +1404,9 @@ class Pregel(PregelProtocol): next_tasks[tid].writes.append((k, v)) if tasks := [t for t in next_tasks.values() if t.writes]: apply_writes(checkpoint, channels, tasks, None) - valid_updates: list[tuple[str, Optional[dict[str, Any]]]] = [] - if len(updates) == 1: - values, as_node = updates[0] - - next_checkpoint = create_checkpoint(checkpoint, None, step) - # copy checkpoint - next_config = checkpointer.put( - checkpoint_config, - next_checkpoint, - { - **checkpoint_metadata, - "source": "update", - "step": step + 1, - "writes": {}, - "parents": saved.metadata.get("parents", {}) if saved else {}, - }, - {}, - ) - return patch_checkpoint_map( - next_config, saved.metadata if saved else None - ) - # no values, copy checkpoint - if values is None and as_node == "__copy__": - if len(updates) > 1: - raise InvalidUpdateError( - "Cannot copy checkpoint with multiple updates" - ) - - next_checkpoint = create_checkpoint(checkpoint, None, step) - # copy checkpoint - next_config = checkpointer.put( - saved.parent_config or saved.config if saved else checkpoint_config, - next_checkpoint, - { - **checkpoint_metadata, - "source": "fork", - "step": step + 1, - "parents": saved.metadata.get("parents", {}) if saved else {}, - }, - {}, - ) - return patch_checkpoint_map( - next_config, saved.metadata if saved else None - ) - # apply pending writes, if not on specific checkpoint - if ( - CONFIG_KEY_CHECKPOINT_ID not in config[CONF] - and saved is not None - and saved.pending_writes - ): - # tasks for this checkpoint - next_tasks = prepare_next_tasks( - checkpoint, - saved.pending_writes, - self.nodes, - channels, - managed, - saved.config, - saved.metadata.get("step", -1) + 1, - for_execution=True, - store=self.store, - checkpointer=( - self.checkpointer - if isinstance(self.checkpointer, BaseCheckpointSaver) - else None - ), - manager=None, - ) - # apply null writes - if null_writes := [ - w[1:] for w in saved.pending_writes or [] if w[0] == NULL_TASK_ID - ]: - apply_writes( - saved.checkpoint, - channels, - [PregelTaskWrites((), INPUT, null_writes, [])], - None, - ) - # apply writes - for tid, k, v in saved.pending_writes: - if k in (ERROR, INTERRUPT, SCHEDULED): - continue - if tid not in next_tasks: - continue - next_tasks[tid].writes.append((k, v)) - if tasks := [t for t in next_tasks.values() if t.writes]: - apply_writes(checkpoint, channels, tasks, None) valid_updates: list[tuple[str, Optional[dict[str, Any]]]] = [] if len(updates) == 1: values, as_node = updates[0] - # find last node that updated the state, if not provided if as_node is None and not any( v @@ -1554,9 +1466,7 @@ class Pregel(PregelProtocol): task_id = str(uuid5(UUID(checkpoint["id"]), INTERRUPT)) run_tasks.append(task) run_task_ids.append(task_id) - run = RunnableSequence(*writers) if len(writers) > 1 else writers[0] - # execute task run.invoke( values, @@ -1582,24 +1492,19 @@ class Pregel(PregelProtocol): }, ), ) - # save task writes for task_id, task in zip(run_task_ids, run_tasks): channel_writes = [w for w in task.writes if w[0] != PUSH] - # channel writes are saved to current checkpoint if saved and channel_writes: checkpointer.put_writes( checkpoint_config, channel_writes, task_id ) - # apply to checkpoint and save mv_writes = apply_writes( checkpoint, channels, run_tasks, checkpointer.get_next_version ) - assert not mv_writes, "Can't write to SharedValues from update_state" - checkpoint = create_checkpoint(checkpoint, channels, step + 1) next_config = checkpointer.put( checkpoint_config, @@ -1617,7 +1522,6 @@ class Pregel(PregelProtocol): checkpoint_previous_versions, checkpoint["channel_versions"] ), ) - for task_id, task in zip(run_task_ids, run_tasks): # save push writes if push_writes := [w for w in task.writes if w[0] == PUSH]: @@ -1714,14 +1618,12 @@ class Pregel(PregelProtocol): managed, ): values, as_node = updates[0] - # no values, just clear all tasks if values is None and as_node == END: if len(updates) > 1: raise InvalidUpdateError( "Cannot apply multiple updates when clearing state" ) - if saved is not None: # tasks for this checkpoint next_tasks = prepare_next_tasks( @@ -1875,74 +1777,9 @@ class Pregel(PregelProtocol): next_tasks[tid].writes.append((k, v)) if tasks := [t for t in next_tasks.values() if t.writes]: apply_writes(checkpoint, channels, tasks, None) - valid_updates: list[tuple[str, Optional[dict[str, Any]]]] = [] - - if len(updates) == 1: - values, as_node = updates[0] - - next_checkpoint = create_checkpoint(checkpoint, None, step) - # copy checkpoint - next_config = await checkpointer.aput( - saved.parent_config or saved.config if saved else checkpoint_config, - next_checkpoint, - { - **checkpoint_metadata, - "source": "fork", - "step": step + 1, - "parents": saved.metadata.get("parents", {}) if saved else {}, - }, - {}, - ) - return patch_checkpoint_map( - next_config, saved.metadata if saved else None - ) - # apply pending writes, if not on specific checkpoint - if ( - CONFIG_KEY_CHECKPOINT_ID not in config[CONF] - and saved is not None - and saved.pending_writes - ): - # tasks for this checkpoint - next_tasks = prepare_next_tasks( - checkpoint, - saved.pending_writes, - self.nodes, - channels, - managed, - saved.config, - saved.metadata.get("step", -1) + 1, - for_execution=True, - store=self.store, - checkpointer=( - self.checkpointer - if isinstance(self.checkpointer, BaseCheckpointSaver) - else None - ), - manager=None, - ) - # apply null writes - if null_writes := [ - w[1:] for w in saved.pending_writes or [] if w[0] == NULL_TASK_ID - ]: - apply_writes( - saved.checkpoint, - channels, - [PregelTaskWrites((), INPUT, null_writes, [])], - None, - ) - for tid, k, v in saved.pending_writes: - if k in (ERROR, INTERRUPT, SCHEDULED): - continue - if tid not in next_tasks: - continue - next_tasks[tid].writes.append((k, v)) - if tasks := [t for t in next_tasks.values() if t.writes]: - apply_writes(checkpoint, channels, tasks, None) valid_updates: list[tuple[str, Optional[dict[str, Any]]]] = [] - if len(updates) == 1: values, as_node = updates[0] - # find last node that updated the state, if not provided if as_node is None and not saved: if ( @@ -1965,10 +1802,8 @@ class Pregel(PregelProtocol): as_node = last_seen_by_node[-1][1] if as_node is None: raise InvalidUpdateError("Ambiguous update, specify as_node") - if as_node not in self.nodes: raise InvalidUpdateError(f"Node {as_node} does not exist") - valid_updates.append((as_node, values)) else: for values, as_node in updates: @@ -1976,12 +1811,6 @@ class Pregel(PregelProtocol): raise InvalidUpdateError( "as_node is required when applying multiple updates" ) - # if two nodes updated the state at the same time, it's ambiguous - if last_seen_by_node: - if len(last_seen_by_node) == 1: - as_node = last_seen_by_node[0][1] - elif last_seen_by_node[-1][0] != last_seen_by_node[-2][0]: - as_node = last_seen_by_node[-1][1] if as_node is None: raise InvalidUpdateError("Ambiguous update, specify as_node") @@ -1989,103 +1818,80 @@ class Pregel(PregelProtocol): raise InvalidUpdateError(f"Node {as_node} does not exist") valid_updates.append((as_node, values)) - else: - for values, as_node in updates: - if as_node is None: - raise InvalidUpdateError( - "as_node is required when applying multiple updates" - ) - - if as_node not in self.nodes: - raise InvalidUpdateError(f"Node {as_node} does not exist") - - valid_updates.append((as_node, values)) run_tasks: list[PregelTaskWrites] = [] run_task_ids: list[str] = [] - for as_node, values in valid_updates: - # create task to run all writers of the chosen node - writers = self.nodes[as_node].flat_writers - if not writers: - raise InvalidUpdateError(f"Node {as_node} has no writers") - writes: deque[tuple[str, Any]] = deque() - task = PregelTaskWrites((), as_node, writes, [INTERRUPT]) - task_id = str(uuid5(UUID(checkpoint["id"]), INTERRUPT)) - run_tasks.append(task) - run_task_ids.append(task_id) - - run = RunnableSequence(*writers) if len(writers) > 1 else writers[0] - - # execute task - await run.ainvoke( - values, - patch_config( - config, - run_name=self.name + "UpdateState", - configurable={ - # deque.extend is thread-safe - CONFIG_KEY_SEND: partial( - local_write, - writes.extend, - self.nodes.keys(), - ), - CONFIG_KEY_READ: partial( - local_read, - step + 1, - checkpoint, - channels, - managed, - task, - config, - ), - }, - ), - ) - - # save task writes - for task_id, task in zip(run_task_ids, run_tasks): - # channel writes are saved to current checkpoint - channel_writes = [w for w in task.writes if w[0] != PUSH] - if saved and channel_writes: - await checkpointer.aput_writes( - checkpoint_config, channel_writes, task_id - ) - - # apply to checkpoint and save - mv_writes = apply_writes( - checkpoint, channels, run_tasks, checkpointer.get_next_version - ) - assert not mv_writes, "Can't write to SharedValues from update_state" - checkpoint = create_checkpoint(checkpoint, channels, step + 1) - # save checkpoint, after applying writes - next_config = await checkpointer.aput( - checkpoint_config, - checkpoint, - { - **checkpoint_metadata, - "source": "update", - "step": step + 1, - "writes": { - as_node: values for as_node, values in valid_updates + for as_node, values in valid_updates: + # create task to run all writers of the chosen node + writers = self.nodes[as_node].flat_writers + if not writers: + raise InvalidUpdateError(f"Node {as_node} has no writers") + writes: deque[tuple[str, Any]] = deque() + task = PregelTaskWrites((), as_node, writes, [INTERRUPT]) + task_id = str(uuid5(UUID(checkpoint["id"]), INTERRUPT)) + run_tasks.append(task) + run_task_ids.append(task_id) + run = RunnableSequence(*writers) if len(writers) > 1 else writers[0] + # execute task + await run.ainvoke( + values, + patch_config( + config, + run_name=self.name + "UpdateState", + configurable={ + # deque.extend is thread-safe + CONFIG_KEY_SEND: partial( + local_write, + writes.extend, + self.nodes.keys(), + ), + CONFIG_KEY_READ: partial( + local_read, + step + 1, + checkpoint, + channels, + managed, + task, + config, + ), }, - "parents": saved.metadata.get("parents", {}) if saved else {}, - }, - get_new_channel_versions( - checkpoint_previous_versions, checkpoint["channel_versions"] ), ) - - for task_id, task in zip(run_task_ids, run_tasks): - # save push writes - if push_writes := [w for w in task.writes if w[0] == PUSH]: - await checkpointer.aput_writes( - next_config, push_writes, task_id - ) - - return patch_checkpoint_map( - next_config, saved.metadata if saved else None - ) + # save task writes + for task_id, task in zip(run_task_ids, run_tasks): + # channel writes are saved to current checkpoint + channel_writes = [w for w in task.writes if w[0] != PUSH] + if saved and channel_writes: + await checkpointer.aput_writes( + checkpoint_config, channel_writes, task_id + ) + # apply to checkpoint and save + mv_writes = apply_writes( + checkpoint, channels, run_tasks, checkpointer.get_next_version + ) + assert not mv_writes, "Can't write to SharedValues from update_state" + checkpoint = create_checkpoint(checkpoint, channels, step + 1) + # save checkpoint, after applying writes + next_config = await checkpointer.aput( + checkpoint_config, + checkpoint, + { + **checkpoint_metadata, + "source": "update", + "step": step + 1, + "writes": {as_node: values for as_node, values in valid_updates}, + "parents": saved.metadata.get("parents", {}) if saved else {}, + }, + get_new_channel_versions( + checkpoint_previous_versions, checkpoint["channel_versions"] + ), + ) + for task_id, task in zip(run_task_ids, run_tasks): + # save push writes + if push_writes := [w for w in task.writes if w[0] == PUSH]: + await checkpointer.aput_writes(next_config, push_writes, task_id) + return patch_checkpoint_map(next_config, saved.metadata if saved else None) current_config = config for superstep in supersteps: