Split out prepare_single_task from prepare_next_tasks

This commit is contained in:
Nuno Campos
2024-09-04 18:09:42 -07:00
parent 6727277f8f
commit ca3715b809
3 changed files with 202 additions and 126 deletions
+2
View File
@@ -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,
+199 -126
View File
@@ -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(
+1
View File
@@ -81,6 +81,7 @@ class PregelExecutableTask(NamedTuple):
retry_policy: Optional[RetryPolicy]
cache_policy: Optional[CachePolicy]
id: str
legacy_id: str = ""
class StateSnapshot(NamedTuple):