Compare commits

...
Author SHA1 Message Date
Vadym BardaandGitHub adff439d4e langgraph: release 0.2.63 (#3064) 2025-01-16 09:19:25 -05:00
Vadym BardaandGitHub 89b0be3a7d checkpoint-postgres: release 2.0.13 (#3063) 2025-01-16 09:17:43 -05:00
Vadym BardaandGitHub 7ecda42b42 checkpoint-postgres: bring back missing migration (#3058) 2025-01-16 03:16:28 +00:00
Nuno CamposandGitHub bd99471705 tests: Add test for catching bunching in streaming (#3057) 2025-01-15 18:57:07 -08:00
Nuno Campos 0e4dbb4c62 Guard 2025-01-15 18:36:49 -08:00
Nuno Campos 145220f2a8 Fix 2025-01-15 18:29:34 -08:00
Eugene Yurtsev f9c25bba07 x 2025-01-15 21:15:38 -05:00
Eugene Yurtsev 7ecad39ecb x 2025-01-15 21:15:05 -05:00
Nuno CamposandGitHub 6bfc6be307 Add support for multiple subgraphs called in a single node (#3056) 2025-01-15 18:01:31 -08:00
Nuno Campos 6b9369876f Fix sync 2025-01-15 17:52:19 -08:00
Nuno Campos c26b0e78b6 Fix kafka lib 2025-01-15 17:15:58 -08:00
Nuno CamposandGitHub 1763dd69a7 Add checkpointer=True mode for subgraphs that want to keep state between turns (#3055) 2025-01-15 16:08:03 -08:00
Nuno Campos f4bd023ab1 Fix 2025-01-15 16:04:14 -08:00
Nuno Campos d402bf7379 Remove flag 2025-01-15 15:36:16 -08:00
Nuno Campos e8a73e1505 Remove flag 2025-01-15 15:36:10 -08:00
Nuno Campos d6492ef048 Add support for multiple subgraphs called in a single node 2025-01-15 15:36:10 -08:00
Nuno Campos 38d9b39f6e Add flag 2025-01-15 15:36:03 -08:00
Nuno Campos be8b4a1d7f Lint 2025-01-15 15:30:29 -08:00
Nuno Campos 71fbd6a8b4 Lint 2025-01-15 15:29:34 -08:00
Nuno Campos 5375af7827 Add checkpointer=True mode for subgraphs that want to keep state betweenn turns 2025-01-15 15:15:44 -08:00
Nuno CamposandGitHub dd7ac00953 Fix unexpected re-use of null resume value by subgraphs (#3054)
- Also stop exposing writes in config, in favor of scratchpad
2025-01-15 13:59:13 -08:00
Nuno Campos cd64075928 Lint 2025-01-15 13:49:33 -08:00
Nuno Campos 144ee31546 Lint 2025-01-15 13:43:14 -08:00
Nuno CamposandGitHub 8bb84c7096 Fix ignored goto when a mixed list of command and state updates is returned from a node (#3038) 2025-01-15 13:39:59 -08:00
Nuno Campos e83660885b Lint 2025-01-15 13:39:44 -08:00
Nuno Campos c2a57385c0 Update tests 2025-01-15 13:35:51 -08:00
Nuno Campos 3626478029 Fix unexpected re-use of null resume value by subgraphs
- Also stop exposing writes in config, in favor of scratchpad
2025-01-15 13:31:49 -08:00
Brace SproulandGitHub 13c9bfa282 feat(langgraph): Add interrupt schema to library (#2947) 2025-01-15 12:55:47 -08:00
Andrew NguonlyandGitHub b6fe3937fc docs: Add section about Persistence to Cloud SaaS concepts page (#3051) 2025-01-15 12:27:53 -08:00
Nuno Campos 5805e5709a Fix ignored goto when a mixed list of command and state updates is returned from a node 2025-01-15 11:42:28 -08:00
Nuno CamposandGitHub aab6fdf3f3 Fix Send order after interrupt/resume (#3037)
- order was incorrectly based on task id, instead of the correct task
path
- this requires storing task paths on checkpointers
- addition of task_path to put_writes is made backwards compatible by
checking signature on call, and treating it as an optional arg
2025-01-15 11:42:08 -08:00
Nuno CamposandGitHub be1d035aba tests: add test for multiple interrupts and tasks (#2941) 2025-01-15 11:41:30 -08:00
Nuno Campos 6e228f8a9c Lint 2025-01-15 11:32:42 -08:00
Nuno Campos 0adbd89d9a Bump checkpoint 2025-01-15 11:28:55 -08:00
Nuno Campos 47c37d140f Fix Send order after interrupt/resume
- order was incorrectly based on task id, instead of the correct task path
- this requires storing task paths on checkpointers
- addition of task_path to put_writes is made backwards compatible by checking signature on call, and treating it as an optinal arg
2025-01-15 11:28:26 -08:00
Nuno Campos 931d39124c Fix 2025-01-15 10:23:30 -08:00
Nuno Campos a7bb96da98 Fix 2025-01-15 10:21:40 -08:00
Nuno Campos 1e9a372dd7 Fix multiple interrupt/task test 2025-01-15 10:11:31 -08:00
Nuno Campos 17aebb6239 Fix flasy return from task 2025-01-15 09:14:36 -08:00
Eugene YurtsevandNuno Campos 29b70cbf39 x 2025-01-15 09:03:46 -08:00
Eugene YurtsevandNuno Campos 6b86fbb0a8 x 2025-01-15 09:03:46 -08:00
Nuno Campos 671f268651 Remove duckdb checkpointer and store
- duckdb is too buggy to be able to provide reliable checkpointer and store
2025-01-15 08:56:32 -08:00
bracesproul 0d50c62283 cr 2025-01-14 10:39:57 -08:00
bracesproul b52b32b38e format n lint 2025-01-13 13:14:36 -08:00
Brace SproulandGitHub 5aefa5dc8c Merge branch 'main' into brace/interrupt-schema 2025-01-13 13:08:37 -08:00
bracesproul cfb121ee8f cr 2025-01-10 10:45:45 -08:00
bracesproul d27beeed18 move to prebuilt 2025-01-10 10:45:22 -08:00
bracesproul 8534212a25 cr 2025-01-07 10:15:48 -08:00
bracesproul e39255a792 feat: Add interrupt schema to library 2025-01-07 10:00:38 -08:00
27 changed files with 1232 additions and 519 deletions
+13 -5
View File
@@ -6,21 +6,29 @@
## Overview ## Overview
LangGraph's Cloud SaaS is a managed service for deploying LangGraph APIs, regardless of its definition or dependencies. The service offers managed implementations of checkpointers and stores, allowing you to focus on building the right cognitive architecture for your use case. By handling scalable & secure infrastructure, LangGraph Cloud offers the fastest path to getting your LangGraph API deployed to production. LangGraph's Cloud SaaS is a managed service for deploying LangGraph Servers, regardless of its definition or dependencies. The service offers managed implementations of checkpointers and stores, allowing you to focus on building the right cognitive architecture for your use case. By handling scalable & secure infrastructure, LangGraph Cloud SaaS offers the fastest path to getting your LangGraph Server deployed to production.
## Deployment ## Deployment
A **deployment** is an instance of a LangGraph API. A single deployment can have many [revisions](#revision). When a deployment is created, all the necessary infrastructure (e.g. database, containers, secrets store) are automatically provisioned. See the [architecture diagram](#architecture) below for more details. A **deployment** is an instance of a LangGraph Server. A single deployment can have many [revisions](#revision). When a deployment is created, all the necessary infrastructure (e.g. database, containers, secrets store) are automatically provisioned. See the [architecture diagram](#architecture) below for more details.
See the [how-to guide](../cloud/deployment/cloud.md#create-new-deployment) for creating a new deployment. Resource Allocation:
## Resource Allocation
| **Deployment Type** | **CPU** | **Memory** | **Scaling** | | **Deployment Type** | **CPU** | **Memory** | **Scaling** |
|---------------------|---------|------------|---------------------| |---------------------|---------|------------|---------------------|
| Development | 1 CPU | 1 GB | Up to 1 container | | Development | 1 CPU | 1 GB | Up to 1 container |
| Production | 2 CPU | 2 GB | Up to 10 containers | | Production | 2 CPU | 2 GB | Up to 10 containers |
See the [how-to guide](../cloud/deployment/cloud.md#create-new-deployment) for creating a new deployment.
## Persistence
A dedicated database is automatically created for each deployment. The database serves as the [persistence layer](../concepts/persistence.md) for the deployment.
When defining a graph to be deployed to LangGraph Cloud SaaS, a [checkpointer](../concepts/persistence.md#checkpointer-libraries) should not be configured by the user. Instead, a checkpointer is automatically configured for the graph.
There is no direct access to the database. All access to the database occurs through the LangGraph Server APIs.
## Autoscaling ## Autoscaling
`Production` type deployments automatically scale up to 10 containers. Scaling is based on the current request load for a single container. Specifically, the autoscaling implementation scales the deployment so that each container is processing about 10 concurrent requests. For example... `Production` type deployments automatically scale up to 10 containers. Scaling is based on the current request load for a single container. Specifically, the autoscaling implementation scales the deployment so that each container is processing about 10 concurrent requests. For example...
@@ -57,6 +57,9 @@ MIGRATIONS = [
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx) PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
);""", );""",
"ALTER TABLE checkpoint_blobs ALTER COLUMN blob DROP not null;", "ALTER TABLE checkpoint_blobs ALTER COLUMN blob DROP not null;",
# NOTE: this is a no-op migration to ensure that the versions in the migrations table are correct.
# This is necessary due to an empty migration previously added to the list.
"SELECT 1;",
""" """
CREATE INDEX CONCURRENTLY IF NOT EXISTS checkpoints_thread_id_idx ON checkpoints(thread_id); CREATE INDEX CONCURRENTLY IF NOT EXISTS checkpoints_thread_id_idx ON checkpoints(thread_id);
""", """,
+1 -1
View File
@@ -1,6 +1,6 @@
[tool.poetry] [tool.poetry]
name = "langgraph-checkpoint-postgres" name = "langgraph-checkpoint-postgres"
version = "2.0.12" version = "2.0.13"
description = "Library with a Postgres implementation of LangGraph checkpoint saver." description = "Library with a Postgres implementation of LangGraph checkpoint saver."
authors = [] authors = []
license = "MIT" license = "MIT"
+1 -3
View File
@@ -75,9 +75,7 @@ CONFIG_KEY_CHECKPOINT_ID = sys.intern("checkpoint_id")
CONFIG_KEY_CHECKPOINT_NS = sys.intern("checkpoint_ns") CONFIG_KEY_CHECKPOINT_NS = sys.intern("checkpoint_ns")
# holds the current checkpoint_ns, "" for root graph # holds the current checkpoint_ns, "" for root graph
CONFIG_KEY_NODE_FINISHED = sys.intern("__pregel_node_finished") CONFIG_KEY_NODE_FINISHED = sys.intern("__pregel_node_finished")
# holds the value that "answers" an interrupt() call # holds a callback to be called when a node is finished
CONFIG_KEY_WRITES = sys.intern("__pregel_writes")
# read-only list of existing task writes
CONFIG_KEY_SCRATCHPAD = sys.intern("__pregel_scratchpad") CONFIG_KEY_SCRATCHPAD = sys.intern("__pregel_scratchpad")
# holds a mutable dict for temporary storage scoped to the current task # holds a mutable dict for temporary storage scoped to the current task
-15
View File
@@ -107,18 +107,3 @@ class CheckpointNotLatest(Exception):
"""Raised when the checkpoint is not the latest version (for distributed mode).""" """Raised when the checkpoint is not the latest version (for distributed mode)."""
pass pass
class MultipleSubgraphsError(Exception):
"""Raised when multiple subgraphs are called inside the same node.
Troubleshooting guides:
- [MULTIPLE_SUBGRAPHS](https://python.langchain.com/docs/troubleshooting/errors/MULTIPLE_SUBGRAPHS)
"""
pass
_SEEN_CHECKPOINT_NS: set[str] = set()
"""Used for subgraph detection."""
@@ -148,6 +148,7 @@ def entrypoint(
output_channels=END, output_channels=END,
stream_channels=END, stream_channels=END,
stream_mode=stream_mode, stream_mode=stream_mode,
stream_eager=True,
checkpointer=checkpointer, checkpointer=checkpointer,
store=store, store=store,
) )
+8 -16
View File
@@ -847,14 +847,10 @@ def _control_branch(value: Any) -> Sequence[Union[str, Send]]:
commands: list[Command] = [] commands: list[Command] = []
if isinstance(value, Command): if isinstance(value, Command):
commands.append(value) commands.append(value)
elif ( elif isinstance(value, (list, tuple)):
isinstance(value, (list, tuple)) for cmd in value:
and value if isinstance(cmd, Command):
and all(isinstance(i, Command) for i in value) commands.append(cmd)
):
commands.extend(value)
else:
return EMPTY_SEQ
rtn: list[Union[str, Send]] = [] rtn: list[Union[str, Send]] = []
for command in commands: for command in commands:
if command.graph == Command.PARENT: if command.graph == Command.PARENT:
@@ -874,14 +870,10 @@ async def _acontrol_branch(value: Any) -> Sequence[Union[str, Send]]:
commands: list[Command] = [] commands: list[Command] = []
if isinstance(value, Command): if isinstance(value, Command):
commands.append(value) commands.append(value)
elif ( elif isinstance(value, (list, tuple)):
isinstance(value, (list, tuple)) for cmd in value:
and value if isinstance(cmd, Command):
and all(isinstance(i, Command) for i in value) commands.append(cmd)
):
commands.extend(value)
else:
return EMPTY_SEQ
rtn: list[Union[str, Send]] = [] rtn: list[Union[str, Send]] = []
for command in commands: for command in commands:
if command.graph == Command.PARENT: if command.graph == Command.PARENT:
@@ -0,0 +1,94 @@
from typing import (
Literal,
Optional,
Union,
)
from typing_extensions import TypedDict
class HumanInterruptConfig(TypedDict):
"""Configuration that defines what actions are allowed for a human interrupt.
This controls the available interaction options when the graph is paused for human input.
Attributes:
allow_ignore: Whether the human can choose to ignore/skip the current step
allow_respond: Whether the human can provide a text response/feedback
allow_edit: Whether the human can edit the provided content/state
allow_accept: Whether the human can accept/approve the current state
"""
allow_ignore: bool
allow_respond: bool
allow_edit: bool
allow_accept: bool
class ActionRequest(TypedDict):
"""Represents a request for human action within the graph execution.
Contains the action type and any associated arguments needed for the action.
Attributes:
action: The type or name of action being requested (e.g., "Approve XYZ action")
args: Key-value pairs of arguments needed for the action
"""
action: str
args: dict
class HumanInterrupt(TypedDict):
"""Represents an interrupt triggered by the graph that requires human intervention.
This is passed to the `interrupt` function when execution is paused for human input.
Attributes:
action_request: The specific action being requested from the human
config: Configuration defining what actions are allowed
description: Optional detailed description of what input is needed
Example:
```python
# Extract a tool call from the state and create an interrupt request
request = HumanInterrupt(
action_request=ActionRequest(
action="run_command", # The action being requested
args={"command": "ls", "args": ["-l"]} # Arguments for the action
),
config=HumanInterruptConfig(
allow_ignore=True, # Allow skipping this step
allow_respond=True, # Allow text feedback
allow_edit=False, # Don't allow editing
allow_accept=True # Allow direct acceptance
),
description="Please review the command before execution"
)
# Send the interrupt request and get the response
response = interrupt([request])[0]
```
"""
action_request: ActionRequest
config: HumanInterruptConfig
description: Optional[str]
class HumanResponse(TypedDict):
"""The response provided by a human to an interrupt, which is returned when graph execution resumes.
Attributes:
type: The type of response:
- "accept": Approves the current state without changes
- "ignore": Skips/ignores the current step
- "response": Provides text feedback or instructions
- "edit": Modifies the current state/content
arg: The response payload:
- None: For ignore/accept actions
- str: For text responses
- ActionRequest: For edit actions with updated content
"""
type: Literal["accept", "ignore", "response", "edit"]
args: Union[None, str, ActionRequest]
+63 -50
View File
@@ -115,6 +115,7 @@ from langgraph.utils.config import (
patch_checkpoint_map, patch_checkpoint_map,
patch_config, patch_config,
patch_configurable, patch_configurable,
recast_checkpoint_ns,
) )
from langgraph.utils.fields import get_enhanced_type_hints from langgraph.utils.fields import get_enhanced_type_hints
from langgraph.utils.pydantic import create_model from langgraph.utils.pydantic import create_model
@@ -203,6 +204,10 @@ class Pregel(PregelProtocol):
stream_mode: StreamMode = "values" stream_mode: StreamMode = "values"
"""Mode to stream output, defaults to 'values'.""" """Mode to stream output, defaults to 'values'."""
stream_eager: bool = False
"""Whether to force emitting stream events eagerly, automatically turned on
for stream_mode "messages" and "custom"."""
output_channels: Union[str, Sequence[str]] output_channels: Union[str, Sequence[str]]
stream_channels: Optional[Union[str, Sequence[str]]] = None stream_channels: Optional[Union[str, Sequence[str]]] = None
@@ -242,6 +247,7 @@ class Pregel(PregelProtocol):
channels: Optional[dict[str, Union[BaseChannel, ManagedValueSpec]]], channels: Optional[dict[str, Union[BaseChannel, ManagedValueSpec]]],
auto_validate: bool = True, auto_validate: bool = True,
stream_mode: StreamMode = "values", stream_mode: StreamMode = "values",
stream_eager: bool = False,
output_channels: Union[str, Sequence[str]], output_channels: Union[str, Sequence[str]],
stream_channels: Optional[Union[str, Sequence[str]]] = None, stream_channels: Optional[Union[str, Sequence[str]]] = None,
interrupt_after_nodes: Union[All, Sequence[str]] = (), interrupt_after_nodes: Union[All, Sequence[str]] = (),
@@ -259,6 +265,7 @@ class Pregel(PregelProtocol):
self.nodes = nodes self.nodes = nodes
self.channels = channels or {} self.channels = channels or {}
self.stream_mode = stream_mode self.stream_mode = stream_mode
self.stream_eager = stream_eager
self.output_channels = output_channels self.output_channels = output_channels
self.stream_channels = stream_channels self.stream_channels = stream_channels
self.interrupt_after_nodes = interrupt_after_nodes self.interrupt_after_nodes = interrupt_after_nodes
@@ -494,7 +501,9 @@ class Pregel(PregelProtocol):
saved.metadata.get("step", -1) + 1, saved.metadata.get("step", -1) + 1,
for_execution=True, for_execution=True,
store=self.store, store=self.store,
checkpointer=self.checkpointer or None, checkpointer=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
manager=None, manager=None,
) )
# get the subgraphs # get the subgraphs
@@ -606,7 +615,9 @@ class Pregel(PregelProtocol):
saved.metadata.get("step", -1) + 1, saved.metadata.get("step", -1) + 1,
for_execution=True, for_execution=True,
store=self.store, store=self.store,
checkpointer=self.checkpointer or None, checkpointer=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
manager=None, manager=None,
) )
# get the subgraphs # get the subgraphs
@@ -690,19 +701,15 @@ class Pregel(PregelProtocol):
checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "") checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "")
) and CONFIG_KEY_CHECKPOINTER not in config[CONF]: ) and CONFIG_KEY_CHECKPOINTER not in config[CONF]:
# remove task_ids from checkpoint_ns # remove task_ids from checkpoint_ns
recast_checkpoint_ns = NS_SEP.join( recast = recast_checkpoint_ns(checkpoint_ns)
part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP)
)
# find the subgraph with the matching name # find the subgraph with the matching name
for _, pregel in self.get_subgraphs( for _, pregel in self.get_subgraphs(namespace=recast, recurse=True):
namespace=recast_checkpoint_ns, recurse=True
):
return pregel.get_state( return pregel.get_state(
patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}), patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}),
subgraphs=subgraphs, subgraphs=subgraphs,
) )
else: else:
raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") raise ValueError(f"Subgraph {recast} not found")
config = merge_configs(self.config, config) if self.config else config config = merge_configs(self.config, config) if self.config else config
saved = checkpointer.get_tuple(config) saved = checkpointer.get_tuple(config)
@@ -727,19 +734,15 @@ class Pregel(PregelProtocol):
checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "") checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "")
) and CONFIG_KEY_CHECKPOINTER not in config[CONF]: ) and CONFIG_KEY_CHECKPOINTER not in config[CONF]:
# remove task_ids from checkpoint_ns # remove task_ids from checkpoint_ns
recast_checkpoint_ns = NS_SEP.join( recast = recast_checkpoint_ns(checkpoint_ns)
part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP)
)
# find the subgraph with the matching name # find the subgraph with the matching name
async for _, pregel in self.aget_subgraphs( async for _, pregel in self.aget_subgraphs(namespace=recast, recurse=True):
namespace=recast_checkpoint_ns, recurse=True
):
return await pregel.aget_state( return await pregel.aget_state(
patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}), patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}),
subgraphs=subgraphs, subgraphs=subgraphs,
) )
else: else:
raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") raise ValueError(f"Subgraph {recast} not found")
config = merge_configs(self.config, config) if self.config else config config = merge_configs(self.config, config) if self.config else config
saved = await checkpointer.aget_tuple(config) saved = await checkpointer.aget_tuple(config)
@@ -770,13 +773,9 @@ class Pregel(PregelProtocol):
checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "") checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "")
) and CONFIG_KEY_CHECKPOINTER not in config[CONF]: ) and CONFIG_KEY_CHECKPOINTER not in config[CONF]:
# remove task_ids from checkpoint_ns # remove task_ids from checkpoint_ns
recast_checkpoint_ns = NS_SEP.join( recast = recast_checkpoint_ns(checkpoint_ns)
part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP)
)
# find the subgraph with the matching name # find the subgraph with the matching name
for _, pregel in self.get_subgraphs( for _, pregel in self.get_subgraphs(namespace=recast, recurse=True):
namespace=recast_checkpoint_ns, recurse=True
):
yield from pregel.get_state_history( yield from pregel.get_state_history(
patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}), patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}),
filter=filter, filter=filter,
@@ -785,7 +784,7 @@ class Pregel(PregelProtocol):
) )
return return
else: else:
raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") raise ValueError(f"Subgraph {recast} not found")
config = merge_configs( config = merge_configs(
self.config, self.config,
@@ -820,13 +819,9 @@ class Pregel(PregelProtocol):
checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "") checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "")
) and CONFIG_KEY_CHECKPOINTER not in config[CONF]: ) and CONFIG_KEY_CHECKPOINTER not in config[CONF]:
# remove task_ids from checkpoint_ns # remove task_ids from checkpoint_ns
recast_checkpoint_ns = NS_SEP.join( recast = recast_checkpoint_ns(checkpoint_ns)
part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP)
)
# find the subgraph with the matching name # find the subgraph with the matching name
async for _, pregel in self.aget_subgraphs( async for _, pregel in self.aget_subgraphs(namespace=recast, recurse=True):
namespace=recast_checkpoint_ns, recurse=True
):
async for state in pregel.aget_state_history( async for state in pregel.aget_state_history(
patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}), patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}),
filter=filter, filter=filter,
@@ -836,7 +831,7 @@ class Pregel(PregelProtocol):
yield state yield state
return return
else: else:
raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") raise ValueError(f"Subgraph {recast} not found")
config = merge_configs( config = merge_configs(
self.config, self.config,
@@ -875,20 +870,16 @@ class Pregel(PregelProtocol):
checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "") checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "")
) and CONFIG_KEY_CHECKPOINTER not in config[CONF]: ) and CONFIG_KEY_CHECKPOINTER not in config[CONF]:
# remove task_ids from checkpoint_ns # remove task_ids from checkpoint_ns
recast_checkpoint_ns = NS_SEP.join( recast = recast_checkpoint_ns(checkpoint_ns)
part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP)
)
# find the subgraph with the matching name # find the subgraph with the matching name
for _, pregel in self.get_subgraphs( for _, pregel in self.get_subgraphs(namespace=recast, recurse=True):
namespace=recast_checkpoint_ns, recurse=True
):
return pregel.update_state( return pregel.update_state(
patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}), patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}),
values, values,
as_node, as_node,
) )
else: else:
raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") raise ValueError(f"Subgraph {recast} not found")
# get last checkpoint # get last checkpoint
config = ensure_config(self.config, config) config = ensure_config(self.config, config)
@@ -926,7 +917,9 @@ class Pregel(PregelProtocol):
saved.metadata.get("step", -1) + 1, saved.metadata.get("step", -1) + 1,
for_execution=True, for_execution=True,
store=self.store, store=self.store,
checkpointer=self.checkpointer or None, checkpointer=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
manager=None, manager=None,
) )
# apply null writes # apply null writes
@@ -1020,7 +1013,9 @@ class Pregel(PregelProtocol):
saved.metadata.get("step", -1) + 1, saved.metadata.get("step", -1) + 1,
for_execution=True, for_execution=True,
store=self.store, store=self.store,
checkpointer=self.checkpointer or None, checkpointer=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
manager=None, manager=None,
) )
# apply null writes # apply null writes
@@ -1155,20 +1150,16 @@ class Pregel(PregelProtocol):
checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "") checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "")
) and CONFIG_KEY_CHECKPOINTER not in config[CONF]: ) and CONFIG_KEY_CHECKPOINTER not in config[CONF]:
# remove task_ids from checkpoint_ns # remove task_ids from checkpoint_ns
recast_checkpoint_ns = NS_SEP.join( recast = recast_checkpoint_ns(checkpoint_ns)
part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP)
)
# find the subgraph with the matching name # find the subgraph with the matching name
async for _, pregel in self.aget_subgraphs( async for _, pregel in self.aget_subgraphs(namespace=recast, recurse=True):
namespace=recast_checkpoint_ns, recurse=True
):
return await pregel.aupdate_state( return await pregel.aupdate_state(
patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}), patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}),
values, values,
as_node, as_node,
) )
else: else:
raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") raise ValueError(f"Subgraph {recast} not found")
# get last checkpoint # get last checkpoint
config = ensure_config(self.config, config) config = ensure_config(self.config, config)
@@ -1209,7 +1200,9 @@ class Pregel(PregelProtocol):
saved.metadata.get("step", -1) + 1, saved.metadata.get("step", -1) + 1,
for_execution=True, for_execution=True,
store=self.store, store=self.store,
checkpointer=self.checkpointer or None, checkpointer=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
manager=None, manager=None,
) )
# apply null writes # apply null writes
@@ -1303,7 +1296,9 @@ class Pregel(PregelProtocol):
saved.metadata.get("step", -1) + 1, saved.metadata.get("step", -1) + 1,
for_execution=True, for_execution=True,
store=self.store, store=self.store,
checkpointer=self.checkpointer or None, checkpointer=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
manager=None, manager=None,
) )
# apply null writes # apply null writes
@@ -1455,6 +1450,8 @@ class Pregel(PregelProtocol):
checkpointer: Optional[BaseCheckpointSaver] = None checkpointer: Optional[BaseCheckpointSaver] = None
elif CONFIG_KEY_CHECKPOINTER in config.get(CONF, {}): elif CONFIG_KEY_CHECKPOINTER in config.get(CONF, {}):
checkpointer = config[CONF][CONFIG_KEY_CHECKPOINTER] checkpointer = config[CONF][CONFIG_KEY_CHECKPOINTER]
elif self.checkpointer is True:
raise RuntimeError("checkpointer=True cannot be used for root graphs.")
else: else:
checkpointer = self.checkpointer checkpointer = self.checkpointer
if checkpointer and not config.get(CONF): if checkpointer and not config.get(CONF):
@@ -1598,6 +1595,12 @@ class Pregel(PregelProtocol):
interrupt_after=interrupt_after, interrupt_after=interrupt_after,
debug=debug, debug=debug,
) )
# set up subgraph checkpointing
if self.checkpointer is True:
ns = cast(str, config[CONF][CONFIG_KEY_CHECKPOINT_NS])
config[CONF][CONFIG_KEY_CHECKPOINT_NS] = NS_SEP.join(
part.split(NS_END)[0] for part in ns.split(NS_SEP)
)
# set up messages stream mode # set up messages stream mode
if "messages" in stream_modes: if "messages" in stream_modes:
run_manager.inheritable_handlers.append( run_manager.inheritable_handlers.append(
@@ -1634,7 +1637,12 @@ class Pregel(PregelProtocol):
if subgraphs: if subgraphs:
loop.config[CONF][CONFIG_KEY_STREAM] = loop.stream loop.config[CONF][CONFIG_KEY_STREAM] = loop.stream
# enable concurrent streaming # enable concurrent streaming
if subgraphs or "messages" in stream_modes or "custom" in stream_modes: if (
self.stream_eager
or subgraphs
or "messages" in stream_modes
or "custom" in stream_modes
):
# we are careful to have a single waiter live at any one time # we are careful to have a single waiter live at any one time
# because on exit we increment semaphore count by exactly 1 # because on exit we increment semaphore count by exactly 1
waiter: Optional[concurrent.futures.Future] = None waiter: Optional[concurrent.futures.Future] = None
@@ -1864,7 +1872,12 @@ class Pregel(PregelProtocol):
stream_put, stream_modes stream_put, stream_modes
) )
# enable concurrent streaming # enable concurrent streaming
if subgraphs or "messages" in stream_modes or "custom" in stream_modes: if (
self.stream_eager
or subgraphs
or "messages" in stream_modes
or "custom" in stream_modes
):
def get_waiter() -> asyncio.Task[None]: def get_waiter() -> asyncio.Task[None]:
return aioloop.create_task(stream.wait()) return aioloop.create_task(stream.wait())
+41 -28
View File
@@ -42,10 +42,10 @@ from langgraph.constants import (
CONFIG_KEY_SEND, CONFIG_KEY_SEND,
CONFIG_KEY_STORE, CONFIG_KEY_STORE,
CONFIG_KEY_TASK_ID, CONFIG_KEY_TASK_ID,
CONFIG_KEY_WRITES,
EMPTY_SEQ, EMPTY_SEQ,
ERROR, ERROR,
INTERRUPT, INTERRUPT,
MISSING,
NO_WRITES, NO_WRITES,
NS_END, NS_END,
NS_SEP, NS_SEP,
@@ -71,6 +71,7 @@ from langgraph.types import (
All, All,
LoopProtocol, LoopProtocol,
PregelExecutableTask, PregelExecutableTask,
PregelScratchpad,
PregelTask, PregelTask,
RetryPolicy, RetryPolicy,
) )
@@ -236,7 +237,7 @@ def apply_writes(
# sort tasks on path, to ensure deterministic order for update application # sort tasks on path, to ensure deterministic order for update application
# any path parts after the 3rd are ignored for sorting # any path parts after the 3rd are ignored for sorting
# (we use them for eg. task ids which aren't good for sorting) # (we use them for eg. task ids which aren't good for sorting)
tasks = sorted(tasks, key=lambda t: _tuple_str(t.path[:3])) tasks = sorted(tasks, key=lambda t: task_path_str(t.path[:3]))
# if no task has triggers this is applying writes from the null task only # if no task has triggers this is applying writes from the null task only
# so we don't do anything other than update the channels written to # so we don't do anything other than update the channels written to
bump_step = any(t.triggers for t in tasks) bump_step = any(t.triggers for t in tasks)
@@ -450,7 +451,7 @@ def prepare_single_task(
str(step), str(step),
name, name,
PUSH, PUSH,
_tuple_str(task_path[1]), task_path_str(task_path[1]),
str(task_path[2]), str(task_path[2]),
) )
task_checkpoint_ns = f"{checkpoint_ns}:{task_id}" task_checkpoint_ns = f"{checkpoint_ns}:{task_id}"
@@ -502,13 +503,10 @@ def prepare_single_task(
}, },
CONFIG_KEY_CHECKPOINT_ID: None, CONFIG_KEY_CHECKPOINT_ID: None,
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns, CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
CONFIG_KEY_WRITES: [ CONFIG_KEY_SCRATCHPAD: _scratchpad(
w pending_writes,
for w in pending_writes task_id,
+ configurable.get(CONFIG_KEY_WRITES, []) ),
if w[0] in (NULL_TASK_ID, task_id)
],
CONFIG_KEY_SCRATCHPAD: {},
}, },
), ),
triggers, triggers,
@@ -614,13 +612,10 @@ def prepare_single_task(
}, },
CONFIG_KEY_CHECKPOINT_ID: None, CONFIG_KEY_CHECKPOINT_ID: None,
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns, CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
CONFIG_KEY_WRITES: [ CONFIG_KEY_SCRATCHPAD: _scratchpad(
w pending_writes,
for w in pending_writes task_id,
+ configurable.get(CONFIG_KEY_WRITES, []) ),
if w[0] in (NULL_TASK_ID, task_id)
],
CONFIG_KEY_SCRATCHPAD: {},
}, },
), ),
triggers, triggers,
@@ -685,7 +680,7 @@ def prepare_single_task(
"langgraph_checkpoint_ns": task_checkpoint_ns, "langgraph_checkpoint_ns": task_checkpoint_ns,
} }
if task_id_checksum is not None: if task_id_checksum is not None:
assert task_id == task_id_checksum assert task_id == task_id_checksum, f"{task_id} != {task_id_checksum}"
if for_execution: if for_execution:
if node := proc.node: if node := proc.node:
if proc.metadata: if proc.metadata:
@@ -738,13 +733,10 @@ def prepare_single_task(
}, },
CONFIG_KEY_CHECKPOINT_ID: None, CONFIG_KEY_CHECKPOINT_ID: None,
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns, CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
CONFIG_KEY_WRITES: [ CONFIG_KEY_SCRATCHPAD: _scratchpad(
w pending_writes,
for w in pending_writes task_id,
+ configurable.get(CONFIG_KEY_WRITES, []) ),
if w[0] in (NULL_TASK_ID, task_id)
],
CONFIG_KEY_SCRATCHPAD: {},
}, },
), ),
triggers, triggers,
@@ -758,6 +750,27 @@ def prepare_single_task(
return PregelTask(task_id, name, task_path[:3]) return PregelTask(task_id, name, task_path[:3])
def _scratchpad(
pending_writes: Sequence[PendingWrite],
task_id: str,
) -> PregelScratchpad:
return PregelScratchpad(
# call
call_counter=0,
# interrupt
interrupt_counter=-1,
resume=next(
(w[2] for w in pending_writes if w[0] == task_id and w[1] == RESUME), []
),
null_resume=next(
(w[2] for w in pending_writes if w[0] == NULL_TASK_ID and w[1] == RESUME),
MISSING,
),
# subgraph
subgraph_counter=0,
)
def _proc_input( def _proc_input(
proc: PregelNode, proc: PregelNode,
managed: ManagedValueMapping, managed: ManagedValueMapping,
@@ -813,10 +826,10 @@ def _uuid5_str(namespace: bytes, *parts: str) -> str:
return f"{hex[:8]}-{hex[8:12]}-{hex[12:16]}-{hex[16:20]}-{hex[20:32]}" return f"{hex[:8]}-{hex[8:12]}-{hex[12:16]}-{hex[16:20]}-{hex[20:32]}"
def _tuple_str(tup: Union[str, int, tuple]) -> str: def task_path_str(tup: Union[str, int, tuple]) -> str:
"""Generate a string representation of a tuple.""" """Generate a string representation of the task path."""
return ( return (
f"~{', '.join(_tuple_str(x) for x in tup)}" f"~{', '.join(task_path_str(x) for x in tup)}"
if isinstance(tup, (tuple, list)) if isinstance(tup, (tuple, list))
else f"{tup:010d}" else f"{tup:010d}"
if isinstance(tup, int) if isinstance(tup, int)
+3 -1
View File
@@ -10,6 +10,7 @@ from langgraph.constants import (
EMPTY_SEQ, EMPTY_SEQ,
ERROR, ERROR,
INTERRUPT, INTERRUPT,
MISSING,
NULL_TASK_ID, NULL_TASK_ID,
RESUME, RESUME,
RETURN, RETURN,
@@ -173,7 +174,8 @@ def map_output_updates(
return return
updated: list[tuple[str, Any]] = [] updated: list[tuple[str, Any]] = []
for task, writes in output_tasks: for task, writes in output_tasks:
if rtn := next((value for chan, value in writes if chan == RETURN), None): rtn = next((value for chan, value in writes if chan == RETURN), MISSING)
if rtn is not MISSING:
updated.append((task.name, rtn)) updated.append((task.name, rtn))
elif isinstance(output_channels, str): elif isinstance(output_channels, str):
updated.extend( updated.extend(
+57 -20
View File
@@ -2,6 +2,7 @@ import asyncio
import concurrent.futures import concurrent.futures
from collections import defaultdict, deque from collections import defaultdict, deque
from contextlib import AsyncExitStack, ExitStack from contextlib import AsyncExitStack, ExitStack
from inspect import signature
from types import TracebackType from types import TracebackType
from typing import ( from typing import (
Any, Any,
@@ -46,12 +47,14 @@ from langgraph.constants import (
CONFIG_KEY_DELEGATE, CONFIG_KEY_DELEGATE,
CONFIG_KEY_ENSURE_LATEST, CONFIG_KEY_ENSURE_LATEST,
CONFIG_KEY_RESUMING, CONFIG_KEY_RESUMING,
CONFIG_KEY_SCRATCHPAD,
CONFIG_KEY_STREAM, CONFIG_KEY_STREAM,
CONFIG_KEY_TASK_ID, CONFIG_KEY_TASK_ID,
EMPTY_SEQ, EMPTY_SEQ,
ERROR, ERROR,
INPUT, INPUT,
INTERRUPT, INTERRUPT,
MISSING,
NS_SEP, NS_SEP,
NULL_TASK_ID, NULL_TASK_ID,
PUSH, PUSH,
@@ -60,12 +63,10 @@ from langgraph.constants import (
TAG_HIDDEN, TAG_HIDDEN,
) )
from langgraph.errors import ( from langgraph.errors import (
_SEEN_CHECKPOINT_NS,
CheckpointNotLatest, CheckpointNotLatest,
EmptyInputError, EmptyInputError,
GraphDelegate, GraphDelegate,
GraphInterrupt, GraphInterrupt,
MultipleSubgraphsError,
) )
from langgraph.managed.base import ( from langgraph.managed.base import (
ManagedValueMapping, ManagedValueMapping,
@@ -81,6 +82,7 @@ from langgraph.pregel.algo import (
prepare_next_tasks, prepare_next_tasks,
prepare_single_task, prepare_single_task,
should_interrupt, should_interrupt,
task_path_str,
) )
from langgraph.pregel.debug import ( from langgraph.pregel.debug import (
map_debug_checkpoint, map_debug_checkpoint,
@@ -112,6 +114,7 @@ from langgraph.types import (
Command, Command,
LoopProtocol, LoopProtocol,
PregelExecutableTask, PregelExecutableTask,
PregelScratchpad,
StreamChunk, StreamChunk,
StreamProtocol, StreamProtocol,
) )
@@ -151,6 +154,7 @@ class PregelLoop(LoopProtocol):
checkpointer_put_writes: Optional[ checkpointer_put_writes: Optional[
Callable[[RunnableConfig, Sequence[tuple[str, Any]], str], Any] Callable[[RunnableConfig, Sequence[tuple[str, Any]], str], Any]
] ]
checkpointer_put_writes_accepts_task_path: bool
_checkpointer_put_after_previous: Optional[ _checkpointer_put_after_previous: Optional[
Callable[ Callable[
[ [
@@ -198,7 +202,6 @@ class PregelLoop(LoopProtocol):
interrupt_after: Union[All, Sequence[str]] = EMPTY_SEQ, interrupt_after: Union[All, Sequence[str]] = EMPTY_SEQ,
interrupt_before: Union[All, Sequence[str]] = EMPTY_SEQ, interrupt_before: Union[All, Sequence[str]] = EMPTY_SEQ,
manager: Union[None, AsyncParentRunManager, ParentRunManager] = None, manager: Union[None, AsyncParentRunManager, ParentRunManager] = None,
check_subgraphs: bool = True,
debug: bool = False, debug: bool = False,
) -> None: ) -> None:
super().__init__( super().__init__(
@@ -225,20 +228,26 @@ class PregelLoop(LoopProtocol):
self.debug = debug self.debug = debug
if self.stream is not None and CONFIG_KEY_STREAM in config[CONF]: if self.stream is not None and CONFIG_KEY_STREAM in config[CONF]:
self.stream = DuplexStream(self.stream, config[CONF][CONFIG_KEY_STREAM]) self.stream = DuplexStream(self.stream, config[CONF][CONFIG_KEY_STREAM])
scratchpad: Optional[PregelScratchpad] = config[CONF].get(CONFIG_KEY_SCRATCHPAD)
if not self.config[CONF].get(CONFIG_KEY_DELEGATE) and scratchpad is not None:
if scratchpad["subgraph_counter"]:
self.config = patch_configurable(
self.config,
{
CONFIG_KEY_CHECKPOINT_NS: NS_SEP.join(
(
config[CONF][CONFIG_KEY_CHECKPOINT_NS],
str(scratchpad["subgraph_counter"]),
)
)
},
)
scratchpad["subgraph_counter"] += 1
if not self.is_nested and config[CONF].get(CONFIG_KEY_CHECKPOINT_NS): if not self.is_nested and config[CONF].get(CONFIG_KEY_CHECKPOINT_NS):
self.config = patch_configurable( self.config = patch_configurable(
self.config, self.config,
{CONFIG_KEY_CHECKPOINT_NS: "", CONFIG_KEY_CHECKPOINT_ID: None}, {CONFIG_KEY_CHECKPOINT_NS: "", CONFIG_KEY_CHECKPOINT_ID: None},
) )
if check_subgraphs and self.is_nested and self.checkpointer is not None:
if self.config[CONF][CONFIG_KEY_CHECKPOINT_NS] in _SEEN_CHECKPOINT_NS:
raise MultipleSubgraphsError(
"Multiple subgraphs called inside the same node\n\n"
"Troubleshooting URL: https://python.langchain.com/docs"
"/troubleshooting/errors/MULTIPLE_SUBGRAPHS/"
)
else:
_SEEN_CHECKPOINT_NS.add(self.config[CONF][CONFIG_KEY_CHECKPOINT_NS])
if ( if (
CONFIG_KEY_CHECKPOINT_MAP in self.config[CONF] CONFIG_KEY_CHECKPOINT_MAP in self.config[CONF]
and self.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS) and self.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS)
@@ -288,9 +297,7 @@ class PregelLoop(LoopProtocol):
else: else:
self.checkpoint_pending_writes.append((task_id, c, v)) self.checkpoint_pending_writes.append((task_id, c, v))
if self.checkpointer_put_writes is not None: if self.checkpointer_put_writes is not None:
self.submit( config = patch_configurable(
self.checkpointer_put_writes,
patch_configurable(
self.checkpoint_config, self.checkpoint_config,
{ {
CONFIG_KEY_CHECKPOINT_NS: self.config[CONF].get( CONFIG_KEY_CHECKPOINT_NS: self.config[CONF].get(
@@ -298,7 +305,23 @@ class PregelLoop(LoopProtocol):
), ),
CONFIG_KEY_CHECKPOINT_ID: self.checkpoint["id"], CONFIG_KEY_CHECKPOINT_ID: self.checkpoint["id"],
}, },
), )
if self.checkpointer_put_writes_accepts_task_path:
if hasattr(self, "tasks"):
task = self.tasks.get(task_id)
else:
task = None
self.submit(
self.checkpointer_put_writes,
config,
writes,
task_id,
task_path_str(task.path) if task else "",
)
else:
self.submit(
self.checkpointer_put_writes,
config,
writes, writes,
task_id, task_id,
) )
@@ -539,8 +562,16 @@ class PregelLoop(LoopProtocol):
) )
) )
# take resume value from parent
if scratchpad := configurable.get(CONFIG_KEY_SCRATCHPAD):
if scratchpad["null_resume"] is not MISSING:
self.put_writes(NULL_TASK_ID, [(RESUME, scratchpad["null_resume"])])
# map command to writes # map command to writes
if isinstance(self.input, Command): if isinstance(self.input, Command):
if self.input.resume is not None and not self.checkpointer:
raise RuntimeError(
"Cannot use Command(resume=...) without checkpointer"
)
writes: defaultdict[str, list[tuple[str, Any]]] = defaultdict(list) writes: defaultdict[str, list[tuple[str, Any]]] = defaultdict(list)
# group writes by task ID # group writes by task ID
for tid, c, v in map_command(self.input, self.checkpoint_pending_writes): for tid, c, v in map_command(self.input, self.checkpoint_pending_writes):
@@ -790,7 +821,6 @@ class SyncPregelLoop(PregelLoop, ContextManager):
interrupt_before: Union[All, Sequence[str]] = EMPTY_SEQ, interrupt_before: Union[All, Sequence[str]] = EMPTY_SEQ,
output_keys: Union[str, Sequence[str]] = EMPTY_SEQ, output_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ, stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
check_subgraphs: bool = True,
debug: bool = False, debug: bool = False,
) -> None: ) -> None:
super().__init__( super().__init__(
@@ -805,7 +835,6 @@ class SyncPregelLoop(PregelLoop, ContextManager):
stream_keys=stream_keys, stream_keys=stream_keys,
interrupt_after=interrupt_after, interrupt_after=interrupt_after,
interrupt_before=interrupt_before, interrupt_before=interrupt_before,
check_subgraphs=check_subgraphs,
manager=manager, manager=manager,
debug=debug, debug=debug,
) )
@@ -813,10 +842,15 @@ class SyncPregelLoop(PregelLoop, ContextManager):
if checkpointer: if checkpointer:
self.checkpointer_get_next_version = checkpointer.get_next_version self.checkpointer_get_next_version = checkpointer.get_next_version
self.checkpointer_put_writes = checkpointer.put_writes self.checkpointer_put_writes = checkpointer.put_writes
self.checkpointer_put_writes_accepts_task_path = (
signature(checkpointer.put_writes).parameters.get("task_path")
is not None
)
else: else:
self.checkpointer_get_next_version = increment self.checkpointer_get_next_version = increment
self._checkpointer_put_after_previous = None # type: ignore[assignment] self._checkpointer_put_after_previous = None # type: ignore[assignment]
self.checkpointer_put_writes = None self.checkpointer_put_writes = None
self.checkpointer_put_writes_accepts_task_path = False
def _checkpointer_put_after_previous( def _checkpointer_put_after_previous(
self, self,
@@ -922,7 +956,6 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
manager: Union[None, AsyncParentRunManager, ParentRunManager] = None, manager: Union[None, AsyncParentRunManager, ParentRunManager] = None,
output_keys: Union[str, Sequence[str]] = EMPTY_SEQ, output_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ, stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
check_subgraphs: bool = True,
debug: bool = False, debug: bool = False,
) -> None: ) -> None:
super().__init__( super().__init__(
@@ -937,7 +970,6 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
stream_keys=stream_keys, stream_keys=stream_keys,
interrupt_after=interrupt_after, interrupt_after=interrupt_after,
interrupt_before=interrupt_before, interrupt_before=interrupt_before,
check_subgraphs=check_subgraphs,
manager=manager, manager=manager,
debug=debug, debug=debug,
) )
@@ -945,10 +977,15 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
if checkpointer: if checkpointer:
self.checkpointer_get_next_version = checkpointer.get_next_version self.checkpointer_get_next_version = checkpointer.get_next_version
self.checkpointer_put_writes = checkpointer.aput_writes self.checkpointer_put_writes = checkpointer.aput_writes
self.checkpointer_put_writes_accepts_task_path = (
signature(checkpointer.aput_writes).parameters.get("task_path")
is not None
)
else: else:
self.checkpointer_get_next_version = increment self.checkpointer_get_next_version = increment
self._checkpointer_put_after_previous = None # type: ignore[assignment] self._checkpointer_put_after_previous = None # type: ignore[assignment]
self.checkpointer_put_writes = None self.checkpointer_put_writes = None
self.checkpointer_put_writes_accepts_task_path = False
async def _checkpointer_put_after_previous( async def _checkpointer_put_after_previous(
self, self,
+9 -17
View File
@@ -12,7 +12,7 @@ from langgraph.constants import (
CONFIG_KEY_RESUMING, CONFIG_KEY_RESUMING,
NS_SEP, NS_SEP,
) )
from langgraph.errors import _SEEN_CHECKPOINT_NS, GraphBubbleUp, ParentCommand from langgraph.errors import GraphBubbleUp, ParentCommand
from langgraph.types import Command, PregelExecutableTask, RetryPolicy from langgraph.types import Command, PregelExecutableTask, RetryPolicy
from langgraph.utils.config import patch_configurable from langgraph.utils.config import patch_configurable
@@ -48,7 +48,10 @@ def run_with_retry(
break break
elif cmd.graph == Command.PARENT: elif cmd.graph == Command.PARENT:
# this command is for the parent graph, assign it to the parent # this command is for the parent graph, assign it to the parent
parent_ns = NS_SEP.join(ns.split(NS_SEP)[:-1]) parts = ns.split(NS_SEP)
if parts[-1].isdigit():
parts.pop()
parent_ns = NS_SEP.join(parts[:-1])
exc.args = (replace(cmd, graph=parent_ns),) exc.args = (replace(cmd, graph=parent_ns),)
# bubble up # bubble up
raise raise
@@ -96,13 +99,6 @@ def run_with_retry(
) )
# signal subgraphs to resume (if available) # signal subgraphs to resume (if available)
config = patch_configurable(config, {CONFIG_KEY_RESUMING: True}) config = patch_configurable(config, {CONFIG_KEY_RESUMING: True})
# clear checkpoint_ns seen (for subgraph detection)
if checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS):
_SEEN_CHECKPOINT_NS.discard(checkpoint_ns)
finally:
# clear checkpoint_ns seen (for subgraph detection)
if checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS):
_SEEN_CHECKPOINT_NS.discard(checkpoint_ns)
async def arun_with_retry( async def arun_with_retry(
@@ -140,7 +136,10 @@ async def arun_with_retry(
break break
elif cmd.graph == Command.PARENT: elif cmd.graph == Command.PARENT:
# this command is for the parent graph, assign it to the parent # this command is for the parent graph, assign it to the parent
parent_ns = NS_SEP.join(ns.split(NS_SEP)[:-1]) parts = ns.split(NS_SEP)
if parts[-1].isdigit():
parts.pop()
parent_ns = NS_SEP.join(parts[:-1])
exc.args = (replace(cmd, graph=parent_ns),) exc.args = (replace(cmd, graph=parent_ns),)
# bubble up # bubble up
raise raise
@@ -188,10 +187,3 @@ async def arun_with_retry(
) )
# signal subgraphs to resume (if available) # signal subgraphs to resume (if available)
config = patch_configurable(config, {CONFIG_KEY_RESUMING: True}) config = patch_configurable(config, {CONFIG_KEY_RESUMING: True})
# clear checkpoint_ns seen (for subgraph detection)
if checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS):
_SEEN_CHECKPOINT_NS.discard(checkpoint_ns)
finally:
# clear checkpoint_ns seen (for subgraph detection)
if checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS):
_SEEN_CHECKPOINT_NS.discard(checkpoint_ns)
+42 -40
View File
@@ -1,6 +1,5 @@
import asyncio import asyncio
import concurrent.futures import concurrent.futures
import threading
import time import time
from functools import partial from functools import partial
from typing import ( from typing import (
@@ -22,9 +21,11 @@ from langchain_core.callbacks import Callbacks
from langgraph.constants import ( from langgraph.constants import (
CONF, CONF,
CONFIG_KEY_CALL, CONFIG_KEY_CALL,
CONFIG_KEY_SCRATCHPAD,
CONFIG_KEY_SEND, CONFIG_KEY_SEND,
ERROR, ERROR,
INTERRUPT, INTERRUPT,
MISSING,
NO_WRITES, NO_WRITES,
PUSH, PUSH,
RESUME, RESUME,
@@ -70,8 +71,6 @@ class PregelRunner:
retry_policy: Optional[RetryPolicy] = None, retry_policy: Optional[RetryPolicy] = None,
get_waiter: Optional[Callable[[], concurrent.futures.Future[None]]] = None, get_waiter: Optional[Callable[[], concurrent.futures.Future[None]]] = None,
) -> Iterator[None]: ) -> Iterator[None]:
locks: dict[str, threading.Lock] = {}
def writer( def writer(
task: PregelExecutableTask, task: PregelExecutableTask,
writes: Sequence[tuple[str, Any]], writes: Sequence[tuple[str, Any]],
@@ -81,24 +80,19 @@ class PregelRunner:
if all(w[0] != PUSH for w in writes): if all(w[0] != PUSH for w in writes):
return task.config[CONF][CONFIG_KEY_SEND](writes) return task.config[CONF][CONFIG_KEY_SEND](writes)
if task.id not in locks:
locks[task.id] = threading.Lock()
with locks[task.id]:
prev_length = len(task.writes)
# delegate to the underlying writer
task.config[CONF][CONFIG_KEY_SEND](writes)
# confirm no other concurrent writes were added
assert len(task.writes) == prev_length + len(writes)
# schedule PUSH tasks, collect futures # schedule PUSH tasks, collect futures
scratchpad = task.config[CONF][CONFIG_KEY_SCRATCHPAD]
scratchpad.setdefault("call_counter", 0)
rtn: dict[int, Optional[concurrent.futures.Future]] = {} rtn: dict[int, Optional[concurrent.futures.Future]] = {}
for idx, w in enumerate(writes, start=prev_length): for idx, w in enumerate(writes):
# bail if not a PUSH write # bail if not a PUSH write
if w[0] != PUSH: if w[0] != PUSH:
continue continue
# schedule the next task, if the callback returns one # schedule the next task, if the callback returns one
if next_task := self.schedule_task( wcall = calls[idx] if calls else None
task, idx, calls[idx - prev_length] if calls else None cnt = scratchpad["call_counter"]
): scratchpad["call_counter"] += 1
if next_task := self.schedule_task(task, cnt, wcall):
if fut := next( if fut := next(
( (
f f
@@ -109,13 +103,18 @@ class PregelRunner:
): ):
# if the parent task was retried, # if the parent task was retried,
# the next task might already be running # the next task might already be running
rtn[idx - prev_length] = fut rtn[idx] = fut
elif next_task.writes: elif next_task.writes:
# if it already ran, return the result # if it already ran, return the result
fut = concurrent.futures.Future() fut = concurrent.futures.Future()
if val := next(v for c, v in next_task.writes if c == RETURN): ret = next(
fut.set_result(val) (v for c, v in next_task.writes if c == RETURN), MISSING
elif exc := next(v for c, v in next_task.writes if c == ERROR): )
if ret is not MISSING:
fut.set_result(ret)
elif exc := next(
(v for c, v in next_task.writes if c == ERROR), None
):
fut.set_exception( fut.set_exception(
exc exc
if isinstance(exc, BaseException) if isinstance(exc, BaseException)
@@ -123,7 +122,7 @@ class PregelRunner:
) )
else: else:
fut.set_result(None) fut.set_result(None)
rtn[idx - prev_length] = fut rtn[idx] = fut
else: else:
# schedule the next task # schedule the next task
fut = self.submit( fut = self.submit(
@@ -141,7 +140,7 @@ class PregelRunner:
) )
fut.add_done_callback(partial(self.commit, next_task)) fut.add_done_callback(partial(self.commit, next_task))
futures[fut] = next_task futures[fut] = next_task
rtn[idx - prev_length] = fut rtn[idx] = fut
return [rtn.get(i) for i in range(len(writes))] return [rtn.get(i) for i in range(len(writes))]
def call( def call(
@@ -189,6 +188,8 @@ class PregelRunner:
raise raise
if not futures: # maybe `t` schuduled another task if not futures: # maybe `t` schuduled another task
return return
else:
tasks = () # don't reschedule this task
# add waiter task if requested # add waiter task if requested
if get_waiter is not None: if get_waiter is not None:
futures[get_waiter()] = None futures[get_waiter()] = None
@@ -255,8 +256,6 @@ class PregelRunner:
retry_policy: Optional[RetryPolicy] = None, retry_policy: Optional[RetryPolicy] = None,
get_waiter: Optional[Callable[[], asyncio.Future[None]]] = None, get_waiter: Optional[Callable[[], asyncio.Future[None]]] = None,
) -> AsyncIterator[None]: ) -> AsyncIterator[None]:
locks: dict[str, threading.Lock] = {}
def writer( def writer(
task: PregelExecutableTask, task: PregelExecutableTask,
writes: Sequence[tuple[str, Any]], writes: Sequence[tuple[str, Any]],
@@ -266,23 +265,19 @@ class PregelRunner:
if all(w[0] != PUSH for w in writes): if all(w[0] != PUSH for w in writes):
return task.config[CONF][CONFIG_KEY_SEND](writes) return task.config[CONF][CONFIG_KEY_SEND](writes)
if task.id not in locks:
locks[task.id] = threading.Lock()
with locks[task.id]:
prev_length = len(task.writes)
# delegate to the underlying writer
task.config[CONF][CONFIG_KEY_SEND](writes)
# confirm no other concurrent writes were added
assert len(task.writes) == prev_length + len(writes)
# schedule PUSH tasks, collect futures # schedule PUSH tasks, collect futures
scratchpad = task.config[CONF][CONFIG_KEY_SCRATCHPAD]
scratchpad.setdefault("call_counter", 0)
rtn: dict[int, Optional[asyncio.Future]] = {} rtn: dict[int, Optional[asyncio.Future]] = {}
for idx, w in enumerate(writes, start=prev_length): for idx, w in enumerate(writes):
# bail if not a PUSH write # bail if not a PUSH write
if w[0] != PUSH: if w[0] != PUSH:
continue continue
# schedule the next task, if the callback returns one # schedule the next task, if the callback returns one
wcall = calls[idx - prev_length] if calls is not None else None wcall = calls[idx] if calls is not None else None
if next_task := self.schedule_task(task, idx, wcall): cnt = scratchpad["call_counter"]
scratchpad["call_counter"] += 1
if next_task := self.schedule_task(task, cnt, wcall):
# if the parent task was retried, # if the parent task was retried,
# the next task might already be running # the next task might already be running
if fut := next( if fut := next(
@@ -295,13 +290,18 @@ class PregelRunner:
): ):
# if the parent task was retried, # if the parent task was retried,
# the next task might already be running # the next task might already be running
rtn[idx - prev_length] = fut rtn[idx] = fut
elif next_task.writes: elif next_task.writes:
# if it already ran, return the result # if it already ran, return the result
fut = asyncio.Future() fut = asyncio.Future()
if val := next(v for c, v in next_task.writes if c == RETURN): ret = next(
fut.set_result(val) (v for c, v in next_task.writes if c == RETURN), MISSING
elif exc := next(v for c, v in next_task.writes if c == ERROR): )
if ret is not MISSING:
fut.set_result(ret)
elif exc := next(
(v for c, v in next_task.writes if c == ERROR), None
):
fut.set_exception( fut.set_exception(
exc exc
if isinstance(exc, BaseException) if isinstance(exc, BaseException)
@@ -309,7 +309,7 @@ class PregelRunner:
) )
else: else:
fut.set_result(None) fut.set_result(None)
rtn[idx - prev_length] = fut rtn[idx] = fut
else: else:
# schedule the next task # schedule the next task
fut = cast( fut = cast(
@@ -333,7 +333,7 @@ class PregelRunner:
) )
fut.add_done_callback(partial(self.commit, next_task)) fut.add_done_callback(partial(self.commit, next_task))
futures[fut] = next_task futures[fut] = next_task
rtn[idx - prev_length] = fut rtn[idx] = fut
return [rtn.get(i) for i in range(len(writes))] return [rtn.get(i) for i in range(len(writes))]
def call( def call(
@@ -388,6 +388,8 @@ class PregelRunner:
raise raise
if not futures: # maybe `t` schuduled another task if not futures: # maybe `t` schuduled another task
return return
else:
tasks = () # don't reschedule this task
# add waiter task if requested # add waiter task if requested
if get_waiter is not None: if get_waiter is not None:
futures[get_waiter()] = None futures[get_waiter()] = None
+17 -25
View File
@@ -21,11 +21,7 @@ from typing import (
from langchain_core.runnables import Runnable, RunnableConfig from langchain_core.runnables import Runnable, RunnableConfig
from typing_extensions import Self, TypedDict from typing_extensions import Self, TypedDict
from langgraph.checkpoint.base import ( from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata
BaseCheckpointSaver,
CheckpointMetadata,
PendingWrite,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from langgraph.store.base import BaseStore from langgraph.store.base import BaseStore
@@ -42,9 +38,11 @@ except ImportError:
All = Literal["*"] All = Literal["*"]
"""Special value to indicate that graph should interrupt on all nodes.""" """Special value to indicate that graph should interrupt on all nodes."""
Checkpointer = Union[None, Literal[False], BaseCheckpointSaver] Checkpointer = Union[None, bool, BaseCheckpointSaver]
"""Type of the checkpointer to use for a subgraph. False disables checkpointing, """Type of the checkpointer to use for a subgraph.
even if the parent graph has a checkpointer. None inherits checkpointer.""" - True enables persistent checkpointing for this subgraph.
- False disables checkpointing, even if the parent graph has a checkpointer.
- None inherits checkpointer from the parent graph."""
StreamMode = Literal["values", "updates", "debug", "messages", "custom"] StreamMode = Literal["values", "updates", "debug", "messages", "custom"]
"""How the stream method should emit outputs. """How the stream method should emit outputs.
@@ -341,10 +339,15 @@ class LoopProtocol:
self.stop = stop self.stop = stop
class PregelScratchpad(TypedDict, total=False): class PregelScratchpad(TypedDict):
# call
call_counter: int
# interrupt
interrupt_counter: int interrupt_counter: int
used_null_resume: bool
resume: list[Any] resume: list[Any]
null_resume: Any
# subgraph
subgraph_counter: int
def interrupt(value: Any) -> Any: def interrupt(value: Any) -> Any:
@@ -446,10 +449,8 @@ def interrupt(value: Any) -> Any:
CONFIG_KEY_CHECKPOINT_NS, CONFIG_KEY_CHECKPOINT_NS,
CONFIG_KEY_SCRATCHPAD, CONFIG_KEY_SCRATCHPAD,
CONFIG_KEY_SEND, CONFIG_KEY_SEND,
CONFIG_KEY_TASK_ID, MISSING,
CONFIG_KEY_WRITES,
NS_SEP, NS_SEP,
NULL_TASK_ID,
RESUME, RESUME,
) )
from langgraph.errors import GraphInterrupt from langgraph.errors import GraphInterrupt
@@ -458,26 +459,17 @@ def interrupt(value: Any) -> Any:
conf = get_config()["configurable"] conf = get_config()["configurable"]
# track interrupt index # track interrupt index
scratchpad: PregelScratchpad = conf[CONFIG_KEY_SCRATCHPAD] scratchpad: PregelScratchpad = conf[CONFIG_KEY_SCRATCHPAD]
if "interrupt_counter" not in scratchpad:
scratchpad["interrupt_counter"] = 0
else:
scratchpad["interrupt_counter"] += 1 scratchpad["interrupt_counter"] += 1
idx = scratchpad["interrupt_counter"] idx = scratchpad["interrupt_counter"]
# find previous resume values # find previous resume values
task_id = conf[CONFIG_KEY_TASK_ID]
writes: list[PendingWrite] = conf[CONFIG_KEY_WRITES]
scratchpad.setdefault(
"resume", next((w[2] for w in writes if w[0] == task_id and w[1] == RESUME), [])
)
if scratchpad["resume"]: if scratchpad["resume"]:
if idx < len(scratchpad["resume"]): if idx < len(scratchpad["resume"]):
return scratchpad["resume"][idx] return scratchpad["resume"][idx]
# find current resume value # find current resume value
if not scratchpad.get("used_null_resume"): if scratchpad["null_resume"] is not MISSING:
scratchpad["used_null_resume"] = True
for tid, c, v in sorted(writes, key=lambda x: x[0], reverse=True):
if tid == NULL_TASK_ID and c == RESUME:
assert len(scratchpad["resume"]) == idx, (scratchpad["resume"], idx) assert len(scratchpad["resume"]) == idx, (scratchpad["resume"], idx)
v = scratchpad["null_resume"]
scratchpad["null_resume"] = MISSING
scratchpad["resume"].append(v) scratchpad["resume"].append(v)
conf[CONFIG_KEY_SEND]([(RESUME, scratchpad["resume"])]) conf[CONFIG_KEY_SEND]([(RESUME, scratchpad["resume"])])
return v return v
+16
View File
@@ -23,6 +23,22 @@ from langgraph.constants import (
CONFIG_KEY_CHECKPOINT_ID, CONFIG_KEY_CHECKPOINT_ID,
CONFIG_KEY_CHECKPOINT_MAP, CONFIG_KEY_CHECKPOINT_MAP,
CONFIG_KEY_CHECKPOINT_NS, CONFIG_KEY_CHECKPOINT_NS,
NS_END,
NS_SEP,
)
def recast_checkpoint_ns(ns: str) -> str:
"""Remove task IDs from checkpoint namespace.
Args:
ns (str): The checkpoint namespace with task IDs.
Returns:
str: The checkpoint namespace without task IDs.
"""
return NS_SEP.join(
part.split(NS_END)[0] for part in ns.split(NS_SEP) if not part.isdigit()
) )
+6 -6
View File
@@ -1348,7 +1348,7 @@ typing-extensions = ">=4.7"
[[package]] [[package]]
name = "langgraph-checkpoint" name = "langgraph-checkpoint"
version = "2.0.9" version = "2.0.10"
description = "Library with base interfaces for LangGraph checkpoint savers." description = "Library with base interfaces for LangGraph checkpoint savers."
optional = false optional = false
python-versions = "^3.9.0,<4.0" python-versions = "^3.9.0,<4.0"
@@ -1366,7 +1366,7 @@ url = "../checkpoint"
[[package]] [[package]]
name = "langgraph-checkpoint-postgres" name = "langgraph-checkpoint-postgres"
version = "2.0.11" version = "2.0.12"
description = "Library with a Postgres implementation of LangGraph checkpoint saver." description = "Library with a Postgres implementation of LangGraph checkpoint saver."
optional = false optional = false
python-versions = "^3.9.0,<4.0" python-versions = "^3.9.0,<4.0"
@@ -1375,7 +1375,7 @@ files = []
develop = true develop = true
[package.dependencies] [package.dependencies]
langgraph-checkpoint = "^2.0.7" langgraph-checkpoint = "^2.0.10"
orjson = ">=3.10.1" orjson = ">=3.10.1"
psycopg = "^3.2.0" psycopg = "^3.2.0"
psycopg-pool = "^3.2.0" psycopg-pool = "^3.2.0"
@@ -1386,7 +1386,7 @@ url = "../checkpoint-postgres"
[[package]] [[package]]
name = "langgraph-checkpoint-sqlite" name = "langgraph-checkpoint-sqlite"
version = "2.0.2" version = "2.0.3"
description = "Library with a SQLite implementation of LangGraph checkpoint saver." description = "Library with a SQLite implementation of LangGraph checkpoint saver."
optional = false optional = false
python-versions = "^3.9.0" python-versions = "^3.9.0"
@@ -1396,7 +1396,7 @@ develop = true
[package.dependencies] [package.dependencies]
aiosqlite = "^0.20.0" aiosqlite = "^0.20.0"
langgraph-checkpoint = "^2.0.2" langgraph-checkpoint = "^2.0.10"
[package.source] [package.source]
type = "directory" type = "directory"
@@ -3491,4 +3491,4 @@ type = ["pytest-mypy"]
[metadata] [metadata]
lock-version = "2.1" lock-version = "2.1"
python-versions = ">=3.9.0,<4.0" python-versions = ">=3.9.0,<4.0"
content-hash = "356f7e84cf1375119bd3a118ecf4e9476d95913736f8c426a76d55717eb39161" content-hash = "caf943b02b6913c05d15c37fda6d216669f789e2a059b7e8e2490b2bdcd23e0e"
+2 -2
View File
@@ -1,6 +1,6 @@
[tool.poetry] [tool.poetry]
name = "langgraph" name = "langgraph"
version = "0.2.62" version = "0.2.63"
description = "Building stateful, multi-actor applications with LLMs" description = "Building stateful, multi-actor applications with LLMs"
authors = [] authors = []
license = "MIT" license = "MIT"
@@ -10,7 +10,7 @@ repository = "https://www.github.com/langchain-ai/langgraph"
[tool.poetry.dependencies] [tool.poetry.dependencies]
python = ">=3.9.0,<4.0" python = ">=3.9.0,<4.0"
langchain-core = ">=0.2.43,<0.4.0,!=0.3.0,!=0.3.1,!=0.3.2,!=0.3.3,!=0.3.4,!=0.3.5,!=0.3.6,!=0.3.7,!=0.3.8,!=0.3.9,!=0.3.10,!=0.3.11,!=0.3.12,!=0.3.13,!=0.3.14,!=0.3.15,!=0.3.16,!=0.3.17,!=0.3.18,!=0.3.19,!=0.3.20,!=0.3.21,!=0.3.22" langchain-core = ">=0.2.43,<0.4.0,!=0.3.0,!=0.3.1,!=0.3.2,!=0.3.3,!=0.3.4,!=0.3.5,!=0.3.6,!=0.3.7,!=0.3.8,!=0.3.9,!=0.3.10,!=0.3.11,!=0.3.12,!=0.3.13,!=0.3.14,!=0.3.15,!=0.3.16,!=0.3.17,!=0.3.18,!=0.3.19,!=0.3.20,!=0.3.21,!=0.3.22"
langgraph-checkpoint = "^2.0.4" langgraph-checkpoint = "^2.0.10"
langgraph-sdk = "^0.1.42" langgraph-sdk = "^0.1.42"
[tool.poetry.group.dev.dependencies] [tool.poetry.group.dev.dependencies]
+6 -6
View File
@@ -1,6 +1,6 @@
from langgraph.checkpoint.base import empty_checkpoint from langgraph.checkpoint.base import empty_checkpoint
from langgraph.constants import PULL, PUSH from langgraph.constants import PULL, PUSH
from langgraph.pregel.algo import _tuple_str, prepare_next_tasks from langgraph.pregel.algo import prepare_next_tasks, task_path_str
from langgraph.pregel.manager import ChannelsManager from langgraph.pregel.manager import ChannelsManager
@@ -49,16 +49,16 @@ def test_tuple_str() -> None:
push_path_b = (PUSH, push_path_a, 1) push_path_b = (PUSH, push_path_a, 1)
push_path_c = (PUSH, push_path_b, 3) push_path_c = (PUSH, push_path_b, 3)
assert _tuple_str(push_path_a) == f"~{PUSH}, 0000000002" assert task_path_str(push_path_a) == f"~{PUSH}, 0000000002"
assert _tuple_str(push_path_b) == f"~{PUSH}, ~{PUSH}, 0000000002, 0000000001" assert task_path_str(push_path_b) == f"~{PUSH}, ~{PUSH}, 0000000002, 0000000001"
assert ( assert (
_tuple_str(push_path_c) task_path_str(push_path_c)
== f"~{PUSH}, ~{PUSH}, ~{PUSH}, 0000000002, 0000000001, 0000000003" == f"~{PUSH}, ~{PUSH}, ~{PUSH}, 0000000002, 0000000001, 0000000003"
) )
assert _tuple_str(pull_path_a) == f"~{PULL}, abc" assert task_path_str(pull_path_a) == f"~{PULL}, abc"
path_list = [push_path_b, push_path_a, pull_path_a, push_path_c] path_list = [push_path_b, push_path_a, pull_path_a, push_path_c]
assert sorted(map(_tuple_str, path_list)) == [ assert sorted(map(task_path_str, path_list)) == [
f"~{PULL}, abc", f"~{PULL}, abc",
f"~{PUSH}, 0000000002", f"~{PUSH}, 0000000002",
f"~{PUSH}, ~{PUSH}, 0000000002, 0000000001", f"~{PUSH}, ~{PUSH}, 0000000002, 0000000001",
+136 -84
View File
@@ -7255,7 +7255,6 @@ def test_branch_then(
) )
@pytest.mark.skip("TODO: re-enable in next PR")
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_send_dedupe_on_resume( def test_send_dedupe_on_resume(
request: pytest.FixtureRequest, checkpointer_name: str request: pytest.FixtureRequest, checkpointer_name: str
@@ -7314,30 +7313,35 @@ def test_send_dedupe_on_resume(
assert graph.invoke(["0"], thread1, debug=1) == [ assert graph.invoke(["0"], thread1, debug=1) == [
"0", "0",
"1", "1",
"3.1",
"2|Command(goto=Send(node='2', arg=3))", "2|Command(goto=Send(node='2', arg=3))",
"2|Command(goto=Send(node='flaky', arg=4))", "2|Command(goto=Send(node='flaky', arg=4))",
"3",
"2|3", "2|3",
] ]
assert builder.nodes["2"].runnable.func.ticks == 3 assert builder.nodes["2"].runnable.func.ticks == 3
assert builder.nodes["flaky"].runnable.func.ticks == 1 assert builder.nodes["flaky"].runnable.func.ticks == 1
# check state # check state
state = graph.get_state(thread1) state = graph.get_state(thread1)
if "shallow" in checkpointer_name:
pytest.xfail("TODO: shallow checkpointer reports wrong next set")
assert state.next == ("flaky",) assert state.next == ("flaky",)
# check history # check history
if "shallow" not in checkpointer_name: if "shallow" not in checkpointer_name:
history = [c for c in graph.get_state_history(thread1)] history = [c for c in graph.get_state_history(thread1)]
assert len(history) == 2 assert len(history) == 4
# resume execution # resume execution
assert graph.invoke(None, thread1, debug=1) == [ assert graph.invoke(None, thread1, debug=1) == [
"0", "0",
"1", "1",
"3.1",
"2|Command(goto=Send(node='2', arg=3))", "2|Command(goto=Send(node='2', arg=3))",
"2|Command(goto=Send(node='flaky', arg=4))", "2|Command(goto=Send(node='flaky', arg=4))",
"3",
"2|3", "2|3",
"flaky|4", "flaky|4",
"3", "3",
"3.1",
] ]
# node "2" doesn't get called again, as we recover writes saved before # node "2" doesn't get called again, as we recover writes saved before
assert builder.nodes["2"].runnable.func.ticks == 3 assert builder.nodes["2"].runnable.func.ticks == 3
@@ -7353,12 +7357,13 @@ def test_send_dedupe_on_resume(
values=[ values=[
"0", "0",
"1", "1",
"3.1",
"2|Command(goto=Send(node='2', arg=3))", "2|Command(goto=Send(node='2', arg=3))",
"2|Command(goto=Send(node='flaky', arg=4))", "2|Command(goto=Send(node='flaky', arg=4))",
"3",
"2|3", "2|3",
"flaky|4", "flaky|4",
"3", "3",
"3.1",
], ],
next=(), next=(),
config={ config={
@@ -7370,35 +7375,33 @@ def test_send_dedupe_on_resume(
}, },
metadata={ metadata={
"source": "loop", "source": "loop",
"writes": {"3": ["3"], "3.1": ["3.1"]}, "writes": {"3": ["3"]},
"thread_id": "1", "thread_id": "1",
"step": 2, "step": 4,
"parents": {}, "parents": {},
}, },
created_at=AnyStr(), created_at=AnyStr(),
parent_config=( parent_config={
None
if "shallow" in checkpointer_name
else {
"configurable": { "configurable": {
"thread_id": "1", "thread_id": "1",
"checkpoint_ns": "", "checkpoint_ns": "",
"checkpoint_id": AnyStr(), "checkpoint_id": AnyStr(),
} }
} },
),
tasks=(), tasks=(),
), ),
StateSnapshot( StateSnapshot(
values=[ values=[
"0", "0",
"1", "1",
"3.1",
"2|Command(goto=Send(node='2', arg=3))", "2|Command(goto=Send(node='2', arg=3))",
"2|Command(goto=Send(node='flaky', arg=4))", "2|Command(goto=Send(node='flaky', arg=4))",
"3",
"2|3", "2|3",
"flaky|4", "flaky|4",
], ],
next=("3", "3.1"), next=("3",),
config={ config={
"configurable": { "configurable": {
"thread_id": "1", "thread_id": "1",
@@ -7408,17 +7411,9 @@ def test_send_dedupe_on_resume(
}, },
metadata={ metadata={
"source": "loop", "source": "loop",
"writes": { "writes": {"2": ["2|3"], "3": ["3"], "flaky": ["flaky|4"]},
"1": ["1"],
"2": [
["2|Command(goto=Send(node='2', arg=3))"],
["2|Command(goto=Send(node='flaky', arg=4))"],
["2|3"],
],
"flaky": ["flaky|4"],
},
"thread_id": "1", "thread_id": "1",
"step": 1, "step": 3,
"parents": {}, "parents": {},
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -7439,6 +7434,123 @@ def test_send_dedupe_on_resume(
state=None, state=None,
result=["3"], result=["3"],
), ),
),
),
StateSnapshot(
values=[
"0",
"1",
"3.1",
"2|Command(goto=Send(node='2', arg=3))",
"2|Command(goto=Send(node='flaky', arg=4))",
],
next=("2", "flaky", "3"),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {
"2": [
["2|Command(goto=Send(node='2', arg=3))"],
["2|Command(goto=Send(node='flaky', arg=4))"],
],
"3.1": ["3.1"],
},
"thread_id": "1",
"step": 2,
"parents": {},
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
},
tasks=(
PregelTask(
id=AnyStr(),
name="2",
path=("__pregel_push", 0),
error=None,
interrupts=(),
state=None,
result=["2|3"],
),
PregelTask(
id=AnyStr(),
name="flaky",
path=("__pregel_push", 1),
error=None,
interrupts=(
Interrupt(
value="Bahh", resumable=False, ns=None, when="during"
),
),
state=None,
result=["flaky|4"],
),
PregelTask(
id=AnyStr(),
name="3",
path=("__pregel_pull", "3"),
error=None,
interrupts=(),
state=None,
result=["3"],
),
),
),
StateSnapshot(
values=["0", "1"],
next=("2", "2", "3.1"),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"1": ["1"]},
"thread_id": "1",
"step": 1,
"parents": {},
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
},
tasks=(
PregelTask(
id=AnyStr(),
name="2",
path=("__pregel_push", 0),
error=None,
interrupts=(),
state=None,
result=["2|Command(goto=Send(node='2', arg=3))"],
),
PregelTask(
id=AnyStr(),
name="2",
path=("__pregel_push", 1),
error=None,
interrupts=(),
state=None,
result=["2|Command(goto=Send(node='flaky', arg=4))"],
),
PregelTask( PregelTask(
id=AnyStr(), id=AnyStr(),
name="3.1", name="3.1",
@@ -7452,7 +7564,7 @@ def test_send_dedupe_on_resume(
), ),
StateSnapshot( StateSnapshot(
values=["0"], values=["0"],
next=("1", "2", "2", "2", "flaky"), next=("1",),
config={ config={
"configurable": { "configurable": {
"thread_id": "1", "thread_id": "1",
@@ -7485,66 +7597,6 @@ def test_send_dedupe_on_resume(
state=None, state=None,
result=["1"], result=["1"],
), ),
PregelTask(
id=AnyStr(),
name="2",
path=(
"__pregel_push",
("__pregel_pull", "1"),
2,
),
error=None,
interrupts=(),
state=None,
result=["2|Command(goto=Send(node='2', arg=3))"],
),
PregelTask(
id=AnyStr(),
name="2",
path=(
"__pregel_push",
("__pregel_pull", "1"),
3,
),
error=None,
interrupts=(),
state=None,
result=["2|Command(goto=Send(node='flaky', arg=4))"],
),
PregelTask(
id=AnyStr(),
name="2",
path=(
"__pregel_push",
(
"__pregel_push",
("__pregel_pull", "1"),
2,
),
2,
),
error=None,
interrupts=(),
state=None,
result=["2|3"],
),
PregelTask(
id=AnyStr(),
name="flaky",
path=(
"__pregel_push",
(
"__pregel_push",
("__pregel_pull", "1"),
3,
),
2,
),
error=None,
interrupts=(Interrupt(value="Bahh", when="during"),),
state=None,
result=["flaky|4"],
),
), ),
), ),
StateSnapshot( StateSnapshot(
+284 -18
View File
@@ -51,7 +51,7 @@ from langgraph.checkpoint.base import (
) )
from langgraph.checkpoint.memory import MemorySaver from langgraph.checkpoint.memory import MemorySaver
from langgraph.constants import CONFIG_KEY_NODE_FINISHED, ERROR, PULL, START from langgraph.constants import CONFIG_KEY_NODE_FINISHED, ERROR, PULL, START
from langgraph.errors import InvalidUpdateError, MultipleSubgraphsError from langgraph.errors import InvalidUpdateError
from langgraph.func import entrypoint, task from langgraph.func import entrypoint, task
from langgraph.graph import END, Graph, StateGraph from langgraph.graph import END, Graph, StateGraph
from langgraph.graph.message import MessageGraph, MessagesState, add_messages from langgraph.graph.message import MessageGraph, MessagesState, add_messages
@@ -1745,9 +1745,8 @@ def test_invoke_join_then_call_other_pregel(
# add checkpointer # add checkpointer
app.checkpointer = checkpointer app.checkpointer = checkpointer
# subgraph is called twice in the same node, through .map(), so raises # subgraph is called twice in the same node, but that works
with pytest.raises(MultipleSubgraphsError): assert app.invoke([2, 3], {"configurable": {"thread_id": "1"}}) == 27
app.invoke([2, 3], {"configurable": {"thread_id": "1"}})
# set inner graph checkpointer NeverCheckpoint # set inner graph checkpointer NeverCheckpoint
inner_app.checkpointer = False inner_app.checkpointer = False
@@ -2167,10 +2166,10 @@ def test_in_one_fan_out_state_graph_waiting_edge(
@workflow.add_node @workflow.add_node
def rewrite_query(data: State) -> State: def rewrite_query(data: State) -> State:
return {"query": f'query: {data["query"]}'} return {"query": f"query: {data['query']}"}
def analyzer_one(data: State) -> State: def analyzer_one(data: State) -> State:
return {"query": f'analyzed: {data["query"]}'} return {"query": f"analyzed: {data['query']}"}
def retriever_one(data: State) -> State: def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]} return {"docs": ["doc1", "doc2"]}
@@ -2307,10 +2306,10 @@ def test_in_one_fan_out_state_graph_waiting_edge_via_branch(
docs: Annotated[list[str], sorted_add] docs: Annotated[list[str], sorted_add]
def rewrite_query(data: State) -> State: def rewrite_query(data: State) -> State:
return {"query": f'query: {data["query"]}'} return {"query": f"query: {data['query']}"}
def analyzer_one(data: State) -> State: def analyzer_one(data: State) -> State:
return {"query": f'analyzed: {data["query"]}'} return {"query": f"analyzed: {data['query']}"}
def retriever_one(data: State) -> State: def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]} return {"docs": ["doc1", "doc2"]}
@@ -2740,11 +2739,11 @@ def test_in_one_fan_out_state_graph_waiting_edge_plus_regular(
docs: Annotated[list[str], sorted_add] docs: Annotated[list[str], sorted_add]
def rewrite_query(data: State) -> State: def rewrite_query(data: State) -> State:
return {"query": f'query: {data["query"]}'} return {"query": f"query: {data['query']}"}
def analyzer_one(data: State) -> State: def analyzer_one(data: State) -> State:
time.sleep(0.1) time.sleep(0.1)
return {"query": f'analyzed: {data["query"]}'} return {"query": f"analyzed: {data['query']}"}
def retriever_one(data: State) -> State: def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]} return {"docs": ["doc1", "doc2"]}
@@ -2830,10 +2829,10 @@ def test_in_one_fan_out_state_graph_waiting_edge_multiple() -> None:
docs: Annotated[list[str], sorted_add] docs: Annotated[list[str], sorted_add]
def rewrite_query(data: State) -> State: def rewrite_query(data: State) -> State:
return {"query": f'query: {data["query"]}'} return {"query": f"query: {data['query']}"}
def analyzer_one(data: State) -> State: def analyzer_one(data: State) -> State:
return {"query": f'analyzed: {data["query"]}'} return {"query": f"analyzed: {data['query']}"}
def retriever_one(data: State) -> State: def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]} return {"docs": ["doc1", "doc2"]}
@@ -2903,10 +2902,10 @@ def test_callable_in_conditional_edges_with_no_path_map() -> None:
query: str query: str
def rewrite(data: State) -> State: def rewrite(data: State) -> State:
return {"query": f'query: {data["query"]}'} return {"query": f"query: {data['query']}"}
def analyze(data: State) -> State: def analyze(data: State) -> State:
return {"query": f'analyzed: {data["query"]}'} return {"query": f"analyzed: {data['query']}"}
class ChooseAnalyzer: class ChooseAnalyzer:
def __call__(self, data: State) -> str: def __call__(self, data: State) -> str:
@@ -2929,10 +2928,10 @@ def test_function_in_conditional_edges_with_no_path_map() -> None:
query: str query: str
def rewrite(data: State) -> State: def rewrite(data: State) -> State:
return {"query": f'query: {data["query"]}'} return {"query": f"query: {data['query']}"}
def analyze(data: State) -> State: def analyze(data: State) -> State:
return {"query": f'analyzed: {data["query"]}'} return {"query": f"analyzed: {data['query']}"}
def choose_analyzer(data: State) -> str: def choose_analyzer(data: State) -> str:
return "analyzer" return "analyzer"
@@ -2965,13 +2964,13 @@ def test_in_one_fan_out_state_graph_waiting_edge_multiple_cond_edge() -> None:
docs: Annotated[list[str], sorted_add] docs: Annotated[list[str], sorted_add]
def rewrite_query(data: State) -> State: def rewrite_query(data: State) -> State:
return {"query": f'query: {data["query"]}'} return {"query": f"query: {data['query']}"}
def retriever_picker(data: State) -> list[str]: def retriever_picker(data: State) -> list[str]:
return ["analyzer_one", "retriever_two"] return ["analyzer_one", "retriever_two"]
def analyzer_one(data: State) -> State: def analyzer_one(data: State) -> State:
return {"query": f'analyzed: {data["query"]}'} return {"query": f"analyzed: {data['query']}"}
def retriever_one(data: State) -> State: def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]} return {"docs": ["doc1", "doc2"]}
@@ -3189,6 +3188,66 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None:
] ]
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_subgraph_checkpoint_true(
request: pytest.FixtureRequest, checkpointer_name: str
) -> None:
checkpointer = request.getfixturevalue("checkpointer_" + checkpointer_name)
class InnerState(TypedDict):
my_key: Annotated[str, operator.add]
my_other_key: str
def inner_1(state: InnerState):
return {"my_key": " got here", "my_other_key": state["my_key"]}
def inner_2(state: InnerState):
return {"my_key": " and there"}
inner = StateGraph(InnerState)
inner.add_node("inner_1", inner_1)
inner.add_node("inner_2", inner_2)
inner.add_edge("inner_1", "inner_2")
inner.set_entry_point("inner_1")
inner.set_finish_point("inner_2")
class State(TypedDict):
my_key: str
graph = StateGraph(State)
graph.add_node("inner", inner.compile(checkpointer=True))
graph.add_edge(START, "inner")
graph.add_conditional_edges(
"inner", lambda s: "inner" if s["my_key"].count("there") < 2 else END
)
app = graph.compile(checkpointer=checkpointer)
config = {"configurable": {"thread_id": "2"}}
assert [c for c in app.stream({"my_key": ""}, config, subgraphs=True)] == [
(("inner",), {"inner_1": {"my_key": " got here", "my_other_key": ""}}),
(("inner",), {"inner_2": {"my_key": " and there"}}),
((), {"inner": {"my_key": " got here and there"}}),
(
("inner",),
{
"inner_1": {
"my_key": " got here",
"my_other_key": " got here and there got here and there",
}
},
),
(("inner",), {"inner_2": {"my_key": " and there"}}),
(
(),
{
"inner": {
"my_key": " got here and there got here and there got here and there"
}
},
),
]
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_stream_subgraphs_during_execution( def test_stream_subgraphs_during_execution(
request: pytest.FixtureRequest, checkpointer_name: str request: pytest.FixtureRequest, checkpointer_name: str
@@ -5260,3 +5319,210 @@ def test_multiple_updates() -> None:
{"node_a": [{"foo": "a1"}, {"foo": "a2"}]}, {"node_a": [{"foo": "a1"}, {"foo": "a2"}]},
{"node_b": {"foo": "b"}}, {"node_b": {"foo": "b"}},
] ]
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_falsy_return_from_task(request: pytest.FixtureRequest, checkpointer_name: str):
"""Test with a falsy return from a task."""
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
@task
def falsy_task() -> bool:
return False
@entrypoint(checkpointer=checkpointer)
def graph(state: dict) -> dict:
"""React tool."""
falsy_task().result()
interrupt("test")
configurable = {"configurable": {"thread_id": str(uuid.uuid4())}}
graph.invoke({"a": 5}, configurable)
graph.invoke(Command(resume="123"), configurable)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_multiple_interrupts_imperative(
request: pytest.FixtureRequest, checkpointer_name: str
):
"""Test multiple interrupts with an imperative API."""
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
counter = 0
@task
def double(x: int) -> int:
"""Increment the counter."""
nonlocal counter
counter += 1
return 2 * x
@entrypoint(checkpointer=checkpointer)
def graph(state: dict) -> dict:
"""React tool."""
values = []
for idx in [1, 2, 3]:
values.extend([double(idx).result(), interrupt({"a": "boo"})])
return {"values": values}
configurable = {"configurable": {"thread_id": str(uuid.uuid4())}}
graph.invoke({}, configurable)
graph.invoke(Command(resume="a"), configurable)
graph.invoke(Command(resume="b"), configurable)
result = graph.invoke(Command(resume="c"), configurable)
# `double` value should be cached appropriately when used w/ `interrupt`
assert result == {
"values": [2, "a", 4, "b", 6, "c"],
}
assert counter == 3
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_double_interrupt_subgraph(
request: pytest.FixtureRequest, checkpointer_name: str
) -> None:
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
class AgentState(TypedDict):
input: str
def node_1(state: AgentState):
result = interrupt("interrupt node 1")
return {"input": result}
def node_2(state: AgentState):
result = interrupt("interrupt node 2")
return {"input": result}
subgraph_builder = (
StateGraph(AgentState)
.add_node("node_1", node_1)
.add_node("node_2", node_2)
.add_edge(START, "node_1")
.add_edge("node_1", "node_2")
.add_edge("node_2", END)
)
# invoke the sub graph
subgraph = subgraph_builder.compile(checkpointer=checkpointer)
thread = {"configurable": {"thread_id": str(uuid.uuid4())}}
assert [c for c in subgraph.stream({"input": "test"}, thread)] == [
{
"__interrupt__": (
Interrupt(
value="interrupt node 1",
resumable=True,
ns=[AnyStr("node_1:")],
when="during",
),
)
},
]
# resume from the first interrupt
assert [c for c in subgraph.stream(Command(resume="123"), thread)] == [
{
"node_1": {"input": "123"},
},
{
"__interrupt__": (
Interrupt(
value="interrupt node 2",
resumable=True,
ns=[AnyStr("node_2:")],
when="during",
),
)
},
]
# resume from the second interrupt
assert [c for c in subgraph.stream(Command(resume="123"), thread)] == [
{
"node_2": {"input": "123"},
},
]
subgraph = subgraph_builder.compile()
def invoke_sub_agent(state: AgentState):
return subgraph.invoke(state)
parent_agent = (
StateGraph(AgentState)
.add_node("invoke_sub_agent", invoke_sub_agent)
.add_edge(START, "invoke_sub_agent")
.add_edge("invoke_sub_agent", END)
.compile(checkpointer=checkpointer)
)
assert [c for c in parent_agent.stream({"input": "test"}, thread)] == [
{
"__interrupt__": (
Interrupt(
value="interrupt node 1",
resumable=True,
ns=[AnyStr("invoke_sub_agent:"), AnyStr("node_1:")],
when="during",
),
)
},
]
# resume from the first interrupt
assert [c for c in parent_agent.stream(Command(resume=True), thread)] == [
{
"__interrupt__": (
Interrupt(
value="interrupt node 2",
resumable=True,
ns=[AnyStr("invoke_sub_agent:"), AnyStr("node_2:")],
when="during",
),
)
}
]
# resume from 2nd interrupt
assert [c for c in parent_agent.stream(Command(resume=True), thread)] == [
{
"invoke_sub_agent": {"input": True},
},
]
def test_sync_streaming_with_functional_api() -> None:
"""Test streaming with functional API.
This test verifies that we're able to stream results as they're being generated
rather than have all the results arrive at once after the graph has completed.
The time of arrival between the two updates corresponding to the two `slow` tasks
should be greater than the time delay between the two tasks.
"""
time_delay = 0.01
@task()
def slow() -> dict:
time.sleep(time_delay) # Simulate a delay of 10 ms
return {"tic": time.time()}
@entrypoint()
def graph(inputs: dict) -> list:
first = slow().result()
second = slow().result()
return [first, second]
arrival_times = []
for chunk in graph.stream({}):
if "slow" not in chunk: # We'll just look at the updates from `slow`
continue
arrival_times.append(time.time())
assert len(arrival_times) == 2
delta = arrival_times[1] - arrival_times[0]
# Delta cannot be less than 10 ms if it is streaming as results are generated.
assert delta > time_delay
+246 -68
View File
@@ -48,7 +48,7 @@ from langgraph.checkpoint.base import (
) )
from langgraph.checkpoint.memory import MemorySaver from langgraph.checkpoint.memory import MemorySaver
from langgraph.constants import CONFIG_KEY_NODE_FINISHED, ERROR, PULL, PUSH, START from langgraph.constants import CONFIG_KEY_NODE_FINISHED, ERROR, PULL, PUSH, START
from langgraph.errors import InvalidUpdateError, MultipleSubgraphsError, NodeInterrupt from langgraph.errors import InvalidUpdateError, NodeInterrupt
from langgraph.func import entrypoint, task from langgraph.func import entrypoint, task
from langgraph.graph import END, Graph, StateGraph from langgraph.graph import END, Graph, StateGraph
from langgraph.graph.message import MessagesState, add_messages from langgraph.graph.message import MessagesState, add_messages
@@ -89,6 +89,11 @@ logger = logging.getLogger(__name__)
pytestmark = pytest.mark.anyio pytestmark = pytest.mark.anyio
NEEDS_CONTEXTVARS = pytest.mark.skipif(
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
async def test_checkpoint_errors() -> None: async def test_checkpoint_errors() -> None:
class FaultyGetCheckpointer(MemorySaver): class FaultyGetCheckpointer(MemorySaver):
@@ -501,10 +506,7 @@ async def test_node_cancellation_on_other_node_exception_two() -> None:
await graph.ainvoke(1) await graph.ainvoke(1)
@pytest.mark.skipif( @NEEDS_CONTEXTVARS
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_dynamic_interrupt(checkpointer_name: str) -> None: async def test_dynamic_interrupt(checkpointer_name: str) -> None:
class State(TypedDict): class State(TypedDict):
@@ -678,10 +680,7 @@ async def test_dynamic_interrupt(checkpointer_name: str) -> None:
) )
@pytest.mark.skipif( @NEEDS_CONTEXTVARS
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_dynamic_interrupt_subgraph(checkpointer_name: str) -> None: async def test_dynamic_interrupt_subgraph(checkpointer_name: str) -> None:
class SubgraphState(TypedDict): class SubgraphState(TypedDict):
@@ -872,10 +871,7 @@ async def test_dynamic_interrupt_subgraph(checkpointer_name: str) -> None:
) )
@pytest.mark.skipif( @NEEDS_CONTEXTVARS
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_copy_checkpoint(checkpointer_name: str) -> None: async def test_copy_checkpoint(checkpointer_name: str) -> None:
class State(TypedDict): class State(TypedDict):
@@ -1079,10 +1075,7 @@ async def test_copy_checkpoint(checkpointer_name: str) -> None:
) )
@pytest.mark.skipif( @NEEDS_CONTEXTVARS
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_node_not_cancelled_on_other_node_interrupted( async def test_node_not_cancelled_on_other_node_interrupted(
checkpointer_name: str, checkpointer_name: str,
@@ -2373,7 +2366,6 @@ async def test_concurrent_emit_sends() -> None:
) )
@pytest.mark.skip("TODO: re-enable in next PR")
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_send_sequences(checkpointer_name: str) -> None: async def test_send_sequences(checkpointer_name: str) -> None:
class Node: class Node:
@@ -2443,10 +2435,7 @@ async def test_send_sequences(checkpointer_name: str) -> None:
] ]
@pytest.mark.skipif( @NEEDS_CONTEXTVARS
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_imp_task(checkpointer_name: str) -> None: async def test_imp_task(checkpointer_name: str) -> None:
async with awith_checkpointer(checkpointer_name) as checkpointer: async with awith_checkpointer(checkpointer_name) as checkpointer:
@@ -2494,10 +2483,7 @@ async def test_imp_task(checkpointer_name: str) -> None:
assert mapper_calls == 2 assert mapper_calls == 2
@pytest.mark.skipif( @NEEDS_CONTEXTVARS
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_imp_task_cancel(checkpointer_name: str) -> None: async def test_imp_task_cancel(checkpointer_name: str) -> None:
async with awith_checkpointer(checkpointer_name) as checkpointer: async with awith_checkpointer(checkpointer_name) as checkpointer:
@@ -2548,10 +2534,7 @@ async def test_imp_task_cancel(checkpointer_name: str) -> None:
assert mapper_cancels == 2 assert mapper_cancels == 2
@pytest.mark.skipif( @NEEDS_CONTEXTVARS
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_imp_sync_from_async(checkpointer_name: str) -> None: async def test_imp_sync_from_async(checkpointer_name: str) -> None:
async with awith_checkpointer(checkpointer_name) as checkpointer: async with awith_checkpointer(checkpointer_name) as checkpointer:
@@ -2584,10 +2567,7 @@ async def test_imp_sync_from_async(checkpointer_name: str) -> None:
] ]
@pytest.mark.skipif( @NEEDS_CONTEXTVARS
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_imp_stream_order(checkpointer_name: str) -> None: async def test_imp_stream_order(checkpointer_name: str) -> None:
async with awith_checkpointer(checkpointer_name) as checkpointer: async with awith_checkpointer(checkpointer_name) as checkpointer:
@@ -2621,7 +2601,6 @@ async def test_imp_stream_order(checkpointer_name: str) -> None:
] ]
@pytest.mark.skip("TODO: re-enable in next PR")
@pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_ASYNC)
async def test_send_dedupe_on_resume(checkpointer_name: str) -> None: async def test_send_dedupe_on_resume(checkpointer_name: str) -> None:
class InterruptOnce: class InterruptOnce:
@@ -2685,7 +2664,6 @@ async def test_send_dedupe_on_resume(checkpointer_name: str) -> None:
] ]
assert builder.nodes["2"].runnable.func.ticks == 3 assert builder.nodes["2"].runnable.func.ticks == 3
assert builder.nodes["flaky"].runnable.func.ticks == 1 assert builder.nodes["flaky"].runnable.func.ticks == 1
print((await graph.aget_state(thread1)).tasks)
# resume execution # resume execution
assert await graph.ainvoke(None, thread1, debug=1) == [ assert await graph.ainvoke(None, thread1, debug=1) == [
"0", "0",
@@ -2694,8 +2672,8 @@ async def test_send_dedupe_on_resume(checkpointer_name: str) -> None:
"2|Command(goto=Send(node='2', arg=3))", "2|Command(goto=Send(node='2', arg=3))",
"2|Command(goto=Send(node='flaky', arg=4))", "2|Command(goto=Send(node='flaky', arg=4))",
"3", "3",
"flaky|4",
"2|3", "2|3",
"flaky|4",
"3", "3",
] ]
# node "2" doesn't get called again, as we recover writes saved before # node "2" doesn't get called again, as we recover writes saved before
@@ -2713,8 +2691,8 @@ async def test_send_dedupe_on_resume(checkpointer_name: str) -> None:
"2|Command(goto=Send(node='2', arg=3))", "2|Command(goto=Send(node='2', arg=3))",
"2|Command(goto=Send(node='flaky', arg=4))", "2|Command(goto=Send(node='flaky', arg=4))",
"3", "3",
"flaky|4",
"2|3", "2|3",
"flaky|4",
"3", "3",
], ],
next=(), next=(),
@@ -2750,8 +2728,8 @@ async def test_send_dedupe_on_resume(checkpointer_name: str) -> None:
"2|Command(goto=Send(node='2', arg=3))", "2|Command(goto=Send(node='2', arg=3))",
"2|Command(goto=Send(node='flaky', arg=4))", "2|Command(goto=Send(node='flaky', arg=4))",
"3", "3",
"flaky|4",
"2|3", "2|3",
"flaky|4",
], ],
next=("3",), next=("3",),
config={ config={
@@ -4071,9 +4049,8 @@ async def test_invoke_join_then_call_other_pregel(
async with awith_checkpointer(checkpointer_name) as checkpointer: async with awith_checkpointer(checkpointer_name) as checkpointer:
# add checkpointer # add checkpointer
app.checkpointer = checkpointer app.checkpointer = checkpointer
# subgraph is called twice in the same node, through .map(), so raises # subgraph is called twice, and that works
with pytest.raises(MultipleSubgraphsError): assert await app.ainvoke([2, 3], {"configurable": {"thread_id": "1"}}) == 27
await app.ainvoke([2, 3], {"configurable": {"thread_id": "1"}})
# set inner graph checkpointer NeverCheckpoint # set inner graph checkpointer NeverCheckpoint
inner_app.checkpointer = False inner_app.checkpointer = False
@@ -4297,10 +4274,10 @@ async def test_in_one_fan_out_state_graph_waiting_edge(checkpointer_name: str) -
docs: Annotated[list[str], sorted_add] docs: Annotated[list[str], sorted_add]
async def rewrite_query(data: State) -> State: async def rewrite_query(data: State) -> State:
return {"query": f'query: {data["query"]}'} return {"query": f"query: {data['query']}"}
async def analyzer_one(data: State) -> State: async def analyzer_one(data: State) -> State:
return {"query": f'analyzed: {data["query"]}'} return {"query": f"analyzed: {data['query']}"}
async def retriever_one(data: State) -> State: async def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]} return {"docs": ["doc1", "doc2"]}
@@ -4387,10 +4364,10 @@ async def test_in_one_fan_out_state_graph_waiting_edge_via_branch(
docs: Annotated[list[str], sorted_add] docs: Annotated[list[str], sorted_add]
async def rewrite_query(data: State) -> State: async def rewrite_query(data: State) -> State:
return {"query": f'query: {data["query"]}'} return {"query": f"query: {data['query']}"}
async def analyzer_one(data: State) -> State: async def analyzer_one(data: State) -> State:
return {"query": f'analyzed: {data["query"]}'} return {"query": f"analyzed: {data['query']}"}
async def retriever_one(data: State) -> State: async def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]} return {"docs": ["doc1", "doc2"]}
@@ -4804,11 +4781,11 @@ async def test_in_one_fan_out_state_graph_waiting_edge_plus_regular(
docs: Annotated[list[str], sorted_add] docs: Annotated[list[str], sorted_add]
async def rewrite_query(data: State) -> State: async def rewrite_query(data: State) -> State:
return {"query": f'query: {data["query"]}'} return {"query": f"query: {data['query']}"}
async def analyzer_one(data: State) -> State: async def analyzer_one(data: State) -> State:
await asyncio.sleep(0.1) await asyncio.sleep(0.1)
return {"query": f'analyzed: {data["query"]}'} return {"query": f"analyzed: {data['query']}"}
async def retriever_one(data: State) -> State: async def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]} return {"docs": ["doc1", "doc2"]}
@@ -4898,10 +4875,10 @@ async def test_in_one_fan_out_state_graph_waiting_edge_multiple() -> None:
docs: Annotated[list[str], sorted_add] docs: Annotated[list[str], sorted_add]
async def rewrite_query(data: State) -> State: async def rewrite_query(data: State) -> State:
return {"query": f'query: {data["query"]}'} return {"query": f"query: {data['query']}"}
async def analyzer_one(data: State) -> State: async def analyzer_one(data: State) -> State:
return {"query": f'analyzed: {data["query"]}'} return {"query": f"analyzed: {data['query']}"}
async def retriever_one(data: State) -> State: async def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]} return {"docs": ["doc1", "doc2"]}
@@ -4982,13 +4959,13 @@ async def test_in_one_fan_out_state_graph_waiting_edge_multiple_cond_edge() -> N
docs: Annotated[list[str], sorted_add] docs: Annotated[list[str], sorted_add]
async def rewrite_query(data: State) -> State: async def rewrite_query(data: State) -> State:
return {"query": f'query: {data["query"]}'} return {"query": f"query: {data['query']}"}
async def retriever_picker(data: State) -> list[str]: async def retriever_picker(data: State) -> list[str]:
return ["analyzer_one", "retriever_two"] return ["analyzer_one", "retriever_two"]
async def analyzer_one(data: State) -> State: async def analyzer_one(data: State) -> State:
return {"query": f'analyzed: {data["query"]}'} return {"query": f"analyzed: {data['query']}"}
async def retriever_one(data: State) -> State: async def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]} return {"docs": ["doc1", "doc2"]}
@@ -6120,10 +6097,7 @@ async def test_parent_command(checkpointer_name: str) -> None:
) )
@pytest.mark.skipif( @NEEDS_CONTEXTVARS
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_interrupt_subgraph(checkpointer_name: str): async def test_interrupt_subgraph(checkpointer_name: str):
class State(TypedDict): class State(TypedDict):
@@ -6156,10 +6130,7 @@ async def test_interrupt_subgraph(checkpointer_name: str):
assert await graph.ainvoke(Command(resume="bar"), thread1) assert await graph.ainvoke(Command(resume="bar"), thread1)
@pytest.mark.skipif( @NEEDS_CONTEXTVARS
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_interrupt_multiple(checkpointer_name: str): async def test_interrupt_multiple(checkpointer_name: str):
class State(TypedDict): class State(TypedDict):
@@ -6223,10 +6194,7 @@ async def test_interrupt_multiple(checkpointer_name: str):
] ]
@pytest.mark.skipif( @NEEDS_CONTEXTVARS
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_interrupt_loop(checkpointer_name: str): async def test_interrupt_loop(checkpointer_name: str):
class State(TypedDict): class State(TypedDict):
@@ -6511,10 +6479,7 @@ async def test_parallel_node_execution():
assert duration < 3.0 assert duration < 3.0
@pytest.mark.skipif( @NEEDS_CONTEXTVARS
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_multiple_interrupt_state_persistence(checkpointer_name: str) -> None: async def test_multiple_interrupt_state_persistence(checkpointer_name: str) -> None:
"""Test that state is preserved correctly across multiple interrupts.""" """Test that state is preserved correctly across multiple interrupts."""
@@ -6693,3 +6658,216 @@ async def test_multiple_updates() -> None:
{"node_a": [{"foo": "a1"}, {"foo": "a2"}]}, {"node_a": [{"foo": "a1"}, {"foo": "a2"}]},
{"node_b": {"foo": "b"}}, {"node_b": {"foo": "b"}},
] ]
@NEEDS_CONTEXTVARS
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_falsy_return_from_task(checkpointer_name: str) -> None:
"""Test with a falsy return from a task."""
@task
async def falsy_task() -> bool:
return False
async with awith_checkpointer(checkpointer_name) as checkpointer:
@entrypoint(checkpointer=checkpointer)
async def graph(state: dict) -> dict:
"""React tool."""
await falsy_task()
interrupt("test")
configurable = {"configurable": {"thread_id": str(uuid.uuid4())}}
await graph.ainvoke({"a": 5}, configurable)
await graph.ainvoke(Command(resume="123"), configurable)
@NEEDS_CONTEXTVARS
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_multiple_interrupts_imperative(checkpointer_name: str) -> None:
"""Test multiple interrupts with an imperative API."""
from langgraph.func import entrypoint, task
counter = 0
@task
async def double(x: int) -> int:
"""Increment the counter."""
nonlocal counter
counter += 1
return 2 * x
async with awith_checkpointer(checkpointer_name) as checkpointer:
@entrypoint(checkpointer=checkpointer)
async def graph(state: dict) -> dict:
"""React tool."""
values = []
for idx in [1, 2, 3]:
values.extend([await double(idx), interrupt({"a": "boo"})])
return {"values": values}
configurable = {"configurable": {"thread_id": str(uuid.uuid4())}}
await graph.ainvoke({}, configurable)
await graph.ainvoke(Command(resume="a"), configurable)
await graph.ainvoke(Command(resume="b"), configurable)
result = await graph.ainvoke(Command(resume="c"), configurable)
# `double` value should be cached appropriately when used w/ `interrupt`
assert result == {
"values": [2, "a", 4, "b", 6, "c"],
}
assert counter == 3
@NEEDS_CONTEXTVARS
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_double_interrupt_subgraph(checkpointer_name: str) -> None:
class AgentState(TypedDict):
input: str
def node_1(state: AgentState):
result = interrupt("interrupt node 1")
return {"input": result}
def node_2(state: AgentState):
result = interrupt("interrupt node 2")
return {"input": result}
subgraph_builder = (
StateGraph(AgentState)
.add_node("node_1", node_1)
.add_node("node_2", node_2)
.add_edge(START, "node_1")
.add_edge("node_1", "node_2")
.add_edge("node_2", END)
)
async with awith_checkpointer(checkpointer_name) as checkpointer:
# invoke the sub graph
subgraph = subgraph_builder.compile(checkpointer=checkpointer)
thread = {"configurable": {"thread_id": str(uuid.uuid4())}}
assert [c async for c in subgraph.astream({"input": "test"}, thread)] == [
{
"__interrupt__": (
Interrupt(
value="interrupt node 1",
resumable=True,
ns=[AnyStr("node_1:")],
when="during",
),
)
},
]
# resume from the first interrupt
assert [c async for c in subgraph.astream(Command(resume="123"), thread)] == [
{
"node_1": {"input": "123"},
},
{
"__interrupt__": (
Interrupt(
value="interrupt node 2",
resumable=True,
ns=[AnyStr("node_2:")],
when="during",
),
)
},
]
# resume from the second interrupt
assert [c async for c in subgraph.astream(Command(resume="123"), thread)] == [
{
"node_2": {"input": "123"},
},
]
subgraph = subgraph_builder.compile()
def invoke_sub_agent(state: AgentState):
return subgraph.invoke(state)
parent_agent = (
StateGraph(AgentState)
.add_node("invoke_sub_agent", invoke_sub_agent)
.add_edge(START, "invoke_sub_agent")
.add_edge("invoke_sub_agent", END)
.compile(checkpointer=checkpointer)
)
assert [c async for c in parent_agent.astream({"input": "test"}, thread)] == [
{
"__interrupt__": (
Interrupt(
value="interrupt node 1",
resumable=True,
ns=[AnyStr("invoke_sub_agent:"), AnyStr("node_1:")],
when="during",
),
)
},
]
# resume from the first interrupt
assert [
c async for c in parent_agent.astream(Command(resume=True), thread)
] == [
{
"__interrupt__": (
Interrupt(
value="interrupt node 2",
resumable=True,
ns=[AnyStr("invoke_sub_agent:"), AnyStr("node_2:")],
when="during",
),
)
}
]
# resume from 2nd interrupt
assert [
c async for c in parent_agent.astream(Command(resume=True), thread)
] == [
{
"invoke_sub_agent": {"input": True},
},
]
@NEEDS_CONTEXTVARS
async def test_async_streaming_with_functional_api() -> None:
"""Test streaming with functional API.
This test verifies that we're able to stream results as they're being generated
rather than have all the results arrive at once after the graph has completed.
The time of arrival between the two updates corresponding to the two `slow` tasks
should be greater than the time delay between the two tasks.
"""
time_delay = 0.01
@task()
async def slow() -> dict:
await asyncio.sleep(time_delay) # Simulate a delay of 10 ms
return {"tic": asyncio.get_running_loop().time()}
@entrypoint()
async def graph(inputs: dict) -> list:
first = await slow()
second = await slow()
return [first, second]
arrival_times = []
async for chunk in graph.astream({}):
if "slow" not in chunk: # We'll just look at the updates from `slow`
continue
arrival_times.append(asyncio.get_running_loop().time())
assert len(arrival_times) == 2
delta = arrival_times[1] - arrival_times[0]
# Delta cannot be less than 10 ms if it is streaming as results are generated.
assert delta > time_delay
@@ -1,5 +1,6 @@
import asyncio import asyncio
import concurrent.futures import concurrent.futures
from collections.abc import Sequence
from contextlib import ( from contextlib import (
AbstractAsyncContextManager, AbstractAsyncContextManager,
AbstractContextManager, AbstractContextManager,
@@ -7,7 +8,7 @@ from contextlib import (
ExitStack, ExitStack,
) )
from functools import partial from functools import partial
from typing import Any, Optional, Sequence from typing import Any, Optional
from uuid import UUID from uuid import UUID
import orjson import orjson
@@ -15,7 +16,7 @@ from langchain_core.runnables import RunnableConfig
from typing_extensions import Self from typing_extensions import Self
import langgraph.scheduler.kafka.serde as serde import langgraph.scheduler.kafka.serde as serde
from langgraph.constants import CONFIG_KEY_DELEGATE, ERROR, NS_END, NS_SEP from langgraph.constants import CONFIG_KEY_DELEGATE, ERROR
from langgraph.errors import CheckpointNotLatest, GraphDelegate, TaskNotFound from langgraph.errors import CheckpointNotLatest, GraphDelegate, TaskNotFound
from langgraph.pregel import Pregel from langgraph.pregel import Pregel
from langgraph.pregel.algo import prepare_single_task from langgraph.pregel.algo import prepare_single_task
@@ -39,7 +40,7 @@ from langgraph.scheduler.kafka.types import (
Topics, Topics,
) )
from langgraph.types import LoopProtocol, PregelExecutableTask, RetryPolicy from langgraph.types import LoopProtocol, PregelExecutableTask, RetryPolicy
from langgraph.utils.config import patch_configurable from langgraph.utils.config import patch_configurable, recast_checkpoint_ns
class AsyncKafkaExecutor(AbstractAsyncContextManager): class AsyncKafkaExecutor(AbstractAsyncContextManager):
@@ -165,14 +166,12 @@ class AsyncKafkaExecutor(AbstractAsyncContextManager):
# find graph # find graph
if checkpoint_ns := msg["config"]["configurable"].get("checkpoint_ns"): if checkpoint_ns := msg["config"]["configurable"].get("checkpoint_ns"):
# remove task_ids from checkpoint_ns # remove task_ids from checkpoint_ns
recast_checkpoint_ns = NS_SEP.join( recast = recast_checkpoint_ns(checkpoint_ns)
part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP)
)
# find the subgraph with the matching name # find the subgraph with the matching name
if recast_checkpoint_ns in self.subgraphs: if recast in self.subgraphs:
graph = self.subgraphs[recast_checkpoint_ns] graph = self.subgraphs[recast]
else: else:
raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") raise ValueError(f"Subgraph {recast} not found")
else: else:
graph = self.graph graph = self.graph
# process message # process message
@@ -183,7 +182,8 @@ class AsyncKafkaExecutor(AbstractAsyncContextManager):
raise RuntimeError("Checkpoint not found") raise RuntimeError("Checkpoint not found")
if saved.checkpoint["id"] != msg["config"]["configurable"]["checkpoint_id"]: if saved.checkpoint["id"] != msg["config"]["configurable"]["checkpoint_id"]:
raise CheckpointNotLatest() raise CheckpointNotLatest()
async with AsyncChannelsManager( async with (
AsyncChannelsManager(
graph.channels, graph.channels,
saved.checkpoint, saved.checkpoint,
LoopProtocol( LoopProtocol(
@@ -192,7 +192,9 @@ class AsyncKafkaExecutor(AbstractAsyncContextManager):
step=saved.metadata["step"] + 1, step=saved.metadata["step"] + 1,
stop=saved.metadata["step"] + 2, stop=saved.metadata["step"] + 2,
), ),
) as (channels, managed), AsyncBackgroundExecutor(msg["config"]) as submit: ) as (channels, managed),
AsyncBackgroundExecutor(msg["config"]) as submit,
):
if task := await asyncio.to_thread( if task := await asyncio.to_thread(
prepare_single_task, prepare_single_task,
msg["task"]["path"], msg["task"]["path"],
@@ -378,14 +380,12 @@ class KafkaExecutor(AbstractContextManager):
# find graph # find graph
if checkpoint_ns := msg["config"]["configurable"].get("checkpoint_ns"): if checkpoint_ns := msg["config"]["configurable"].get("checkpoint_ns"):
# remove task_ids from checkpoint_ns # remove task_ids from checkpoint_ns
recast_checkpoint_ns = NS_SEP.join( recast = recast_checkpoint_ns(checkpoint_ns)
part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP)
)
# find the subgraph with the matching name # find the subgraph with the matching name
if recast_checkpoint_ns in self.subgraphs: if recast in self.subgraphs:
graph = self.subgraphs[recast_checkpoint_ns] graph = self.subgraphs[recast]
else: else:
raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") raise ValueError(f"Subgraph {recast} not found")
else: else:
graph = self.graph graph = self.graph
# process message # process message
@@ -396,7 +396,8 @@ class KafkaExecutor(AbstractContextManager):
raise RuntimeError("Checkpoint not found") raise RuntimeError("Checkpoint not found")
if saved.checkpoint["id"] != msg["config"]["configurable"]["checkpoint_id"]: if saved.checkpoint["id"] != msg["config"]["configurable"]["checkpoint_id"]:
raise CheckpointNotLatest() raise CheckpointNotLatest()
with ChannelsManager( with (
ChannelsManager(
graph.channels, graph.channels,
saved.checkpoint, saved.checkpoint,
LoopProtocol( LoopProtocol(
@@ -405,7 +406,9 @@ class KafkaExecutor(AbstractContextManager):
step=saved.metadata["step"] + 1, step=saved.metadata["step"] + 1,
stop=saved.metadata["step"] + 2, stop=saved.metadata["step"] + 2,
), ),
) as (channels, managed), BackgroundExecutor({}) as submit: ) as (channels, managed),
BackgroundExecutor({}) as submit,
):
if task := prepare_single_task( if task := prepare_single_task(
msg["task"]["path"], msg["task"]["path"],
msg["task"]["id"], msg["task"]["id"],
@@ -13,11 +13,11 @@ from typing_extensions import Self
import langgraph.scheduler.kafka.serde as serde import langgraph.scheduler.kafka.serde as serde
from langgraph.constants import ( from langgraph.constants import (
CONF,
CONFIG_KEY_DEDUPE_TASKS, CONFIG_KEY_DEDUPE_TASKS,
CONFIG_KEY_ENSURE_LATEST, CONFIG_KEY_ENSURE_LATEST,
CONFIG_KEY_SCRATCHPAD,
INTERRUPT, INTERRUPT,
NS_END,
NS_SEP,
SCHEDULED, SCHEDULED,
) )
from langgraph.errors import CheckpointNotLatest, GraphInterrupt from langgraph.errors import CheckpointNotLatest, GraphInterrupt
@@ -37,7 +37,7 @@ from langgraph.scheduler.kafka.types import (
Topics, Topics,
) )
from langgraph.types import RetryPolicy from langgraph.types import RetryPolicy
from langgraph.utils.config import patch_configurable from langgraph.utils.config import patch_configurable, recast_checkpoint_ns
class AsyncKafkaOrchestrator(AbstractAsyncContextManager): class AsyncKafkaOrchestrator(AbstractAsyncContextManager):
@@ -140,14 +140,12 @@ class AsyncKafkaOrchestrator(AbstractAsyncContextManager):
# find graph # find graph
if checkpoint_ns := msg["config"]["configurable"].get("checkpoint_ns"): if checkpoint_ns := msg["config"]["configurable"].get("checkpoint_ns"):
# remove task_ids from checkpoint_ns # remove task_ids from checkpoint_ns
recast_checkpoint_ns = NS_SEP.join( recast = recast_checkpoint_ns(checkpoint_ns)
part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP)
)
# find the subgraph with the matching name # find the subgraph with the matching name
if recast_checkpoint_ns in self.subgraphs: if recast in self.subgraphs:
graph = self.subgraphs[recast_checkpoint_ns] graph = self.subgraphs[recast]
else: else:
raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") raise ValueError(f"Subgraph {recast} not found")
else: else:
graph = self.graph graph = self.graph
# process message # process message
@@ -163,7 +161,6 @@ class AsyncKafkaOrchestrator(AbstractAsyncContextManager):
stream_keys=graph.stream_channels, stream_keys=graph.stream_channels,
interrupt_after=graph.interrupt_after_nodes, interrupt_after=graph.interrupt_after_nodes,
interrupt_before=graph.interrupt_before_nodes, interrupt_before=graph.interrupt_before_nodes,
check_subgraphs=False,
) as loop: ) as loop:
if loop.tick(input_keys=graph.input_channels): if loop.tick(input_keys=graph.input_channels):
# wait for checkpoint to be saved # wait for checkpoint to be saved
@@ -173,6 +170,16 @@ class AsyncKafkaOrchestrator(AbstractAsyncContextManager):
if new_tasks := [ if new_tasks := [
t for t in loop.tasks.values() if not t.scheduled and not t.writes t for t in loop.tasks.values() if not t.scheduled and not t.writes
]: ]:
config = patch_configurable(
loop.config,
{
**loop.checkpoint_config["configurable"],
CONFIG_KEY_DEDUPE_TASKS: True,
CONFIG_KEY_ENSURE_LATEST: True,
},
)
if CONFIG_KEY_SCRATCHPAD in config[CONF]:
config[CONF][CONFIG_KEY_SCRATCHPAD]["subgraph_counter"] = 0
# send messages to executor # send messages to executor
futures = await asyncio.gather( futures = await asyncio.gather(
*( *(
@@ -180,16 +187,7 @@ class AsyncKafkaOrchestrator(AbstractAsyncContextManager):
self.topics.executor, self.topics.executor,
value=serde.dumps( value=serde.dumps(
MessageToExecutor( MessageToExecutor(
config=patch_configurable( config=config,
loop.config,
{
**loop.checkpoint_config[
"configurable"
],
CONFIG_KEY_DEDUPE_TASKS: True,
CONFIG_KEY_ENSURE_LATEST: True,
},
),
task=ExecutorTask(id=task.id, path=task.path), task=ExecutorTask(id=task.id, path=task.path),
finally_send=msg.get("finally_send"), finally_send=msg.get("finally_send"),
) )
@@ -330,14 +328,12 @@ class KafkaOrchestrator(AbstractContextManager):
# find graph # find graph
if checkpoint_ns := msg["config"]["configurable"].get("checkpoint_ns"): if checkpoint_ns := msg["config"]["configurable"].get("checkpoint_ns"):
# remove task_ids from checkpoint_ns # remove task_ids from checkpoint_ns
recast_checkpoint_ns = NS_SEP.join( recast = recast_checkpoint_ns(checkpoint_ns)
part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP)
)
# find the subgraph with the matching name # find the subgraph with the matching name
if recast_checkpoint_ns in self.subgraphs: if recast in self.subgraphs:
graph = self.subgraphs[recast_checkpoint_ns] graph = self.subgraphs[recast]
else: else:
raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") raise ValueError(f"Subgraph {recast} not found")
else: else:
graph = self.graph graph = self.graph
# process message # process message
@@ -353,7 +349,6 @@ class KafkaOrchestrator(AbstractContextManager):
stream_keys=graph.stream_channels, stream_keys=graph.stream_channels,
interrupt_after=graph.interrupt_after_nodes, interrupt_after=graph.interrupt_after_nodes,
interrupt_before=graph.interrupt_before_nodes, interrupt_before=graph.interrupt_before_nodes,
check_subgraphs=False,
) as loop: ) as loop:
if loop.tick(input_keys=graph.input_channels): if loop.tick(input_keys=graph.input_channels):
# wait for checkpoint to be saved # wait for checkpoint to be saved
@@ -363,12 +358,6 @@ class KafkaOrchestrator(AbstractContextManager):
if new_tasks := [ if new_tasks := [
t for t in loop.tasks.values() if not t.scheduled and not t.writes t for t in loop.tasks.values() if not t.scheduled and not t.writes
]: ]:
# send messages to executor
futures = [
self.producer.send(
self.topics.executor,
value=serde.dumps(
MessageToExecutor(
config = patch_configurable( config = patch_configurable(
loop.config, loop.config,
{ {
@@ -376,7 +365,16 @@ class KafkaOrchestrator(AbstractContextManager):
CONFIG_KEY_DEDUPE_TASKS: True, CONFIG_KEY_DEDUPE_TASKS: True,
CONFIG_KEY_ENSURE_LATEST: True, CONFIG_KEY_ENSURE_LATEST: True,
}, },
), )
if CONFIG_KEY_SCRATCHPAD in config[CONF]:
config[CONF][CONFIG_KEY_SCRATCHPAD]["subgraph_counter"] = 0
# send messages to executor
futures = [
self.producer.send(
self.topics.executor,
value=serde.dumps(
MessageToExecutor(
config=config,
task=ExecutorTask(id=task.id, path=task.path), task=ExecutorTask(id=task.id, path=task.path),
finally_send=msg.get("finally_send"), finally_send=msg.get("finally_send"),
) )
+8
View File
@@ -53,3 +53,11 @@ class AnyList(list):
return False return False
else: else:
return True return True
class AnyInt(int):
def __init__(self) -> None:
super().__init__()
def __eq__(self, other: object) -> bool:
return isinstance(other, int)
+43 -13
View File
@@ -15,7 +15,7 @@ from langgraph.graph.state import StateGraph
from langgraph.pregel import Pregel from langgraph.pregel import Pregel
from langgraph.scheduler.kafka import serde from langgraph.scheduler.kafka import serde
from langgraph.scheduler.kafka.types import MessageToOrchestrator, Topics from langgraph.scheduler.kafka.types import MessageToOrchestrator, Topics
from tests.any import AnyDict, AnyList from tests.any import AnyDict, AnyInt
from tests.drain import drain_topics_async from tests.drain import drain_topics_async
from tests.messages import _AnyIdAIMessage, _AnyIdHumanMessage from tests.messages import _AnyIdAIMessage, _AnyIdHumanMessage
@@ -197,8 +197,13 @@ async def test_subgraph_w_interrupt(
"__pregel_resuming": False, "__pregel_resuming": False,
"__pregel_store": None, "__pregel_store": None,
"__pregel_task_id": history[0].tasks[0].id, "__pregel_task_id": history[0].tasks[0].id,
"__pregel_scratchpad": {}, "__pregel_scratchpad": {
"__pregel_writes": AnyList(), "subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"checkpoint_id": None, "checkpoint_id": None,
"checkpoint_map": { "checkpoint_map": {
"": history[0].config["configurable"]["checkpoint_id"] "": history[0].config["configurable"]["checkpoint_id"]
@@ -264,8 +269,13 @@ async def test_subgraph_w_interrupt(
"__pregel_resuming": False, "__pregel_resuming": False,
"__pregel_store": None, "__pregel_store": None,
"__pregel_task_id": history[0].tasks[0].id, "__pregel_task_id": history[0].tasks[0].id,
"__pregel_scratchpad": {}, "__pregel_scratchpad": {
"__pregel_writes": AnyList(), "subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"checkpoint_id": c.config["configurable"]["checkpoint_id"], "checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": { "checkpoint_map": {
"": history[0].config["configurable"]["checkpoint_id"] "": history[0].config["configurable"]["checkpoint_id"]
@@ -361,8 +371,13 @@ async def test_subgraph_w_interrupt(
"__pregel_resuming": False, "__pregel_resuming": False,
"__pregel_store": None, "__pregel_store": None,
"__pregel_task_id": history[0].tasks[0].id, "__pregel_task_id": history[0].tasks[0].id,
"__pregel_scratchpad": {}, "__pregel_scratchpad": {
"__pregel_writes": AnyList(), "subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"checkpoint_id": c.config["configurable"]["checkpoint_id"], "checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": { "checkpoint_map": {
"": history[0].config["configurable"]["checkpoint_id"] "": history[0].config["configurable"]["checkpoint_id"]
@@ -468,8 +483,13 @@ async def test_subgraph_w_interrupt(
"__pregel_resuming": True, "__pregel_resuming": True,
"__pregel_store": None, "__pregel_store": None,
"__pregel_task_id": history[1].tasks[0].id, "__pregel_task_id": history[1].tasks[0].id,
"__pregel_scratchpad": {}, "__pregel_scratchpad": {
"__pregel_writes": AnyList(), "subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"checkpoint_id": None, "checkpoint_id": None,
"checkpoint_map": { "checkpoint_map": {
"": history[1].config["configurable"]["checkpoint_id"] "": history[1].config["configurable"]["checkpoint_id"]
@@ -530,8 +550,13 @@ async def test_subgraph_w_interrupt(
"__pregel_resuming": True, "__pregel_resuming": True,
"__pregel_store": None, "__pregel_store": None,
"__pregel_task_id": history[1].tasks[0].id, "__pregel_task_id": history[1].tasks[0].id,
"__pregel_scratchpad": {}, "__pregel_scratchpad": {
"__pregel_writes": AnyList(), "subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"checkpoint_id": c.config["configurable"]["checkpoint_id"], "checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": { "checkpoint_map": {
"": history[1].config["configurable"]["checkpoint_id"] "": history[1].config["configurable"]["checkpoint_id"]
@@ -648,8 +673,13 @@ async def test_subgraph_w_interrupt(
"__pregel_resuming": True, "__pregel_resuming": True,
"__pregel_store": None, "__pregel_store": None,
"__pregel_task_id": history[1].tasks[0].id, "__pregel_task_id": history[1].tasks[0].id,
"__pregel_scratchpad": {}, "__pregel_scratchpad": {
"__pregel_writes": AnyList(), "subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"checkpoint_id": c.config["configurable"]["checkpoint_id"], "checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": { "checkpoint_map": {
"": history[1].config["configurable"]["checkpoint_id"] "": history[1].config["configurable"]["checkpoint_id"]
@@ -15,7 +15,7 @@ from langgraph.pregel import Pregel
from langgraph.scheduler.kafka import serde from langgraph.scheduler.kafka import serde
from langgraph.scheduler.kafka.default_sync import DefaultProducer from langgraph.scheduler.kafka.default_sync import DefaultProducer
from langgraph.scheduler.kafka.types import MessageToOrchestrator, Topics from langgraph.scheduler.kafka.types import MessageToOrchestrator, Topics
from tests.any import AnyDict, AnyList from tests.any import AnyDict, AnyInt
from tests.drain import drain_topics from tests.drain import drain_topics
from tests.messages import _AnyIdAIMessage, _AnyIdHumanMessage from tests.messages import _AnyIdAIMessage, _AnyIdHumanMessage
@@ -196,8 +196,13 @@ def test_subgraph_w_interrupt(
"__pregel_resuming": False, "__pregel_resuming": False,
"__pregel_store": None, "__pregel_store": None,
"__pregel_task_id": history[0].tasks[0].id, "__pregel_task_id": history[0].tasks[0].id,
"__pregel_scratchpad": {}, "__pregel_scratchpad": {
"__pregel_writes": AnyList(), "subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"checkpoint_id": None, "checkpoint_id": None,
"checkpoint_map": { "checkpoint_map": {
"": history[0].config["configurable"]["checkpoint_id"] "": history[0].config["configurable"]["checkpoint_id"]
@@ -263,8 +268,13 @@ def test_subgraph_w_interrupt(
"__pregel_dedupe_tasks": True, "__pregel_dedupe_tasks": True,
"__pregel_resuming": False, "__pregel_resuming": False,
"__pregel_task_id": history[0].tasks[0].id, "__pregel_task_id": history[0].tasks[0].id,
"__pregel_scratchpad": {}, "__pregel_scratchpad": {
"__pregel_writes": AnyList(), "subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"checkpoint_id": c.config["configurable"]["checkpoint_id"], "checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": { "checkpoint_map": {
"": history[0].config["configurable"]["checkpoint_id"] "": history[0].config["configurable"]["checkpoint_id"]
@@ -360,8 +370,13 @@ def test_subgraph_w_interrupt(
"__pregel_store": None, "__pregel_store": None,
"__pregel_resuming": False, "__pregel_resuming": False,
"__pregel_task_id": history[0].tasks[0].id, "__pregel_task_id": history[0].tasks[0].id,
"__pregel_scratchpad": {}, "__pregel_scratchpad": {
"__pregel_writes": AnyList(), "subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"checkpoint_id": c.config["configurable"]["checkpoint_id"], "checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": { "checkpoint_map": {
"": history[0].config["configurable"]["checkpoint_id"] "": history[0].config["configurable"]["checkpoint_id"]
@@ -466,8 +481,13 @@ def test_subgraph_w_interrupt(
"__pregel_store": None, "__pregel_store": None,
"__pregel_resuming": True, "__pregel_resuming": True,
"__pregel_task_id": history[1].tasks[0].id, "__pregel_task_id": history[1].tasks[0].id,
"__pregel_scratchpad": {}, "__pregel_scratchpad": {
"__pregel_writes": AnyList(), "subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"checkpoint_id": None, "checkpoint_id": None,
"checkpoint_map": { "checkpoint_map": {
"": history[1].config["configurable"]["checkpoint_id"] "": history[1].config["configurable"]["checkpoint_id"]
@@ -528,8 +548,13 @@ def test_subgraph_w_interrupt(
"__pregel_store": None, "__pregel_store": None,
"__pregel_resuming": True, "__pregel_resuming": True,
"__pregel_task_id": history[1].tasks[0].id, "__pregel_task_id": history[1].tasks[0].id,
"__pregel_scratchpad": {}, "__pregel_scratchpad": {
"__pregel_writes": AnyList(), "subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"checkpoint_id": c.config["configurable"]["checkpoint_id"], "checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": { "checkpoint_map": {
"": history[1].config["configurable"]["checkpoint_id"] "": history[1].config["configurable"]["checkpoint_id"]
@@ -646,8 +671,13 @@ def test_subgraph_w_interrupt(
"__pregel_resuming": True, "__pregel_resuming": True,
"__pregel_store": None, "__pregel_store": None,
"__pregel_task_id": history[1].tasks[0].id, "__pregel_task_id": history[1].tasks[0].id,
"__pregel_scratchpad": {}, "__pregel_scratchpad": {
"__pregel_writes": AnyList(), "subgraph_counter": AnyInt(),
"call_counter": 0,
"interrupt_counter": -1,
"null_resume": None,
"resume": [],
},
"checkpoint_id": c.config["configurable"]["checkpoint_id"], "checkpoint_id": c.config["configurable"]["checkpoint_id"],
"checkpoint_map": { "checkpoint_map": {
"": history[1].config["configurable"]["checkpoint_id"] "": history[1].config["configurable"]["checkpoint_id"]