Compare commits

..
Author SHA1 Message Date
Caspar Broekhuizen 4b6ac3e3d7 style(langgraph): fml 2025-10-16 11:57:54 -07:00
Caspar Broekhuizen 5eae68324a fix(langgraph): remove unnecessary code handling None invoke case 2025-10-16 11:56:16 -07:00
Caspar Broekhuizen dd02a773ab fix(langgraph): do NOT re-execute nodes on invoke(None, ...). fix tests 2025-10-16 11:33:43 -07:00
Caspar Broekhuizen 8b42793d30 style(langgraph): remove prints 2025-10-09 14:37:13 -07:00
Caspar Broekhuizen 7dba4f6791 style(langgraph): format lint 2025-10-09 14:33:30 -07:00
Caspar Broekhuizen b4549b436f fix(langgraph): don't save null writes to checkpoint 2025-10-09 14:22:53 -07:00
Caspar Broekhuizen 015563bd47 fix(langgraph): fix duplicate interrupt writes when resuming with None 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen 5e60b7eaeb fix(langgraph): add missing context var check 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen 8611ff7e98 fix(langgraph): add missing context check 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen c97f818c5c fix(langgraph): add context var check for async test 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen 330df89868 style(langgraph): fix spelling error 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen f2c3b3cc42 style(langgraph): lint 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen 26b7a6da77 fix(langgraph): cleanup rebase error 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen 9893a1602a refactor(langgraph): move helper 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen f06402ca80 style(langgraph): refactor and fml 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen 1b83cc280d fix(langgraph): add optimization support for node with multiple interrupts 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen eefe1f4d16 style(langgraph): rename vars and add comments for clarity 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen 9b48311b42 test(langgraph): add xfail test that node with multiple interrupts should not execute until both have been resumed 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen 6a58e0cd6a refactor(langgraph): clean up optimization logic 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen 29f1ae79ec fix(langgraph): fix interrupt optimization for AsyncPregelLoop 2025-10-08 15:43:31 -07:00
Caspar Broekhuizen 8420e966c4 test(langgraph): add async interrupt test. still failing test_interrupt_with_send_payloads_sequential_resume_async 2025-10-08 15:43:31 -07:00
Eugene YurtsevandCaspar Broekhuizen ab704272b8 x 2025-10-08 15:43:31 -07:00
Eugene YurtsevandCaspar Broekhuizen d23914adcc x 2025-10-08 15:43:31 -07:00
Eugene YurtsevandCaspar Broekhuizen e8cc79e3f7 x 2025-10-08 15:43:31 -07:00
Eugene YurtsevandCaspar Broekhuizen 711d81bc38 Test with multiple interrupts 2025-10-08 15:43:31 -07:00
Eugene YurtsevandCaspar Broekhuizen 40f0f72870 x 2025-10-08 15:43:31 -07:00
Caspar BroekhuizenandGitHub 420550501f fix(langgraph): revert selective interrupt task scheduling (#6252)
Reverts langchain-ai/langgraph#6158
2025-10-08 12:34:01 -07:00
Sam CrowderandGitHub a0599139b8 fix(cli): rename studio to debugger (#6246)
begin process of renaming Studio to Debugger

keep --studio-url around for now as an option as well
2025-10-08 09:40:26 -07:00
Parker J. RuleandGitHub c2f359f708 chore(cli): bump to 0.4.3 (#6251)
Releases #6193 (and a few other minor changes).
2025-10-08 15:31:54 +00:00
Caspar BroekhuizenandGitHub 7f78a011fd chore(langgraph): bump version (#6245)
- Bump langgraph version from 0.6.8 to 0.6.9
2025-10-07 13:45:25 -07:00
Caspar BroekhuizenandGitHub 7d166bfb9f chore(checkpoint): bump patch version (#6244)
- Bump `langgraph-checkpoint` to 2.1.2
- Bump `langgraph-checkpoint-postgres` to 2.0.25 and raise
`langgraph-checkpoint` dep lower bound to 2.1.2
2025-10-07 10:41:24 -07:00
6cc8899818 fix(langgraph): selective interrupt task scheduling (#6158)
### Description

Prevents interrupt tasks from executing when the resume value has not
yet been specified.

Implemented for sync and async Pregel loop

If a task execution is skipped, the skipped interrupt is still included
in the graph result for consistency:
``` python
result = graph.invoke(...)
interrupts = result.get("__interrupt__", [])   # [interrupt_1, interrupt_2]

partial_result = graph.invoke(Command(resume=interrupt_1_resume_map), ...)
remaining_interrupts = partial_result.get("__interrupt__", [])  # [interrupt_2]
```

### Tests

- `test_interrupt_with_send_payloads`: test for a single resume map that
resumes all interrupts at once
- `test_interrupt_with_send_payloads_sequential_resume`: test for two
resume maps delivered in sequence
- `test_node_with_multiple_interrupts_requires_full_resume` test
optimization for multiple interrupts within a single node

Solves https://github.com/langchain-ai/langgraph/issues/6208

---------

Co-authored-by: Eugene Yurtsev <eyurtsev@gmail.com>
2025-10-06 13:11:52 -07:00
1ba96f49bf fix(checkpoint): handle metadata.writes when serializing old checkpoints with Jsonb (#6236)
Issue

Support for `Checkpoint.metadata.writes` was dropped in `langgraph`
v0.5.x.

In `langgraph-checkpoint-postgres` v2.0.23, metadata was serialized with
`BasePostgresSaver._dump_metadata` -> `JsonPlusSerializer.dumps` which
handles `pydantic.BaseModel`.

In v2.0.23, metadata is serialized with `psycopg.types.json.Jsonb`,
which raises `TypeError: Object of type AIMessage is not JSON
serializable` when trying to serialize `writes`.

Solution

- Add `BaseCheckpointSaver.get_serializable_checkpoint_metadata` which
pops the `writes` key.
- Log deprecation warning when strange version combinations are used 

Solves https://github.com/langchain-ai/langgraph/issues/5769

---------

Co-authored-by: Alex Kondratev <56111142+soapun@users.noreply.github.com>
2025-10-06 11:27:34 -07:00
Caspar BroekhuizenandGitHub b0958115c1 fix(langgraph): task result from stream mode debug / tasks should match format from get_state_history / get_state (#6233)
Overview

Python port of https://github.com/langchain-ai/langgraphjs/pull/1551

Introduces `map_task_result_writes` to standardize task result format
across `get_state_history` and `map_task_result_writes` response
structures.

Solves https://github.com/langchain-ai/langgraph/issues/6073
2025-10-03 09:06:58 -07:00
Mason DaughertyandGitHub 04fb14d3ae fix(langgraph): don't use rst code blocks in docstrings (#6231) 2025-10-01 00:08:07 +00:00
Mason DaughertyandGitHub efb0e8c176 docs(langgraph): standardize version-added admonitions (#6230) 2025-09-30 18:39:12 -04:00
Caspar BroekhuizenandGitHub 0584eaa5c4 fix(langgraph): fix supersteps not populating task.result field (#6195)
### Description

Fix `bulk_update_state` and `abulk_update_state` so history populates
`tasks[*].result` when creating state via supersteps.

There was a branch in these functions that I'm guessing was meant to be
triggered when a `StateUpdate.as_node` was the name of a real node (not
`"__input__"` or `"__copy__"`), but was never being triggered because of
a condition `CONFIG_KEY_CHECKPOINT_ID not in config[CONF]`:
```python
# apply pending writes, if not on specific checkpoint
if (
    CONFIG_KEY_CHECKPOINT_ID not in config[CONF]
    and saved is not None
    and saved.pending_writes
):
    next_tasks = prepare_next_tasks(...)
```

From what I can tell, in the bulk-update flow every superstep carries a
`checkpoint_id`, so the condition was always false. That skipped
`prepare_next_tasks(...)` and prevented us from discovering the task IDs
that we would need to attach the task result. So, I removed this check.

I also replaced the `pending_writes` check with a more lenient one (just
check it is not None to satisfy type checkers). I found that
`saved.pending_writes` was sometimes just `[]`, and in this case we
would skip `prepare_next_tasks(...)` and never attach the task result.

Now for each task discovered in `prepare_next_tasks(...)`, I collect the
task IDs and reuse them when running all writers of the chosen node
(applying the updates).

### Tests

- `test_supersteps_populate_task_results` for `PregelLoop` and
`AsyncPregelLoop`
 
These tests build a single node graph and compare history from two
threads: one uses `.invoke` and the other is build from supersteps. Both
tests fail on main and pass with this PR.

### Issue

Solves https://github.com/langchain-ai/langgraph/issues/6206
2025-09-30 12:52:59 -07:00
Caspar BroekhuizenandGitHub 0c73af5624 fix(langgraph): revert -- reuse cached writes on nested resume to prevent task re-execution (#6227)
Reverts langchain-ai/langgraph#6161
2025-09-30 11:28:41 -07:00
32 changed files with 1158 additions and 256 deletions
+3 -3
View File
@@ -277,9 +277,9 @@ def my_function(arg1: int, arg2: str) -> float:
Examples:
This is a section for examples of how to use the function.
.. code-block:: python
my_function(1, "hello")
```python
my_function(1, "hello")
\```
Args:
arg1: This is a description of arg1. We do not need to specify the type since
+1 -1
View File
@@ -33,7 +33,7 @@ LangGraph provides three ways to manage context, which combines the mutability a
**Static runtime context** represents immutable data like user metadata, tools, and database connections that are passed to an application at the start of a run via the `context` argument to `invoke`/`stream`. This data does not change during execution.
!!! version-added "New in LangGraph v0.6: `context` replaces `config['configurable']`"
!!! version-added "Added in version 0.6.0: `context` replaces `config['configurable']`"
Runtime context is now passed to the `context` argument of `invoke`/`stream`,
which replaces the previous pattern of passing application configuration to `config['configurable']`.
+5 -1
View File
@@ -211,7 +211,7 @@ output = agent.invoke(
print(output["messages"][-1].text())
```
!!! version-added "New in LangGraph v0.6"
!!! version-added "Added in version 0.6.0"
:::
@@ -351,11 +351,13 @@ If your desired LLM isn't officially supported by LangChain, consider these opti
:::python
1. **Implement a custom LangChain chat model**: Create a model conforming to the [LangChain chat model interface](https://python.langchain.com/docs/how_to/custom_chat_model/). This enables full compatibility with LangGraph's agents and workflows but requires understanding of the LangChain framework.
:::
:::js
1. **Implement a custom LangChain chat model**: Create a model conforming to the [LangChain chat model interface](https://js.langchain.com/docs/how_to/custom_chat/). This enables full compatibility with LangGraph's agents and workflows but requires understanding of the LangChain framework.
:::
2. **Direct invocation with custom streaming**: Use your model directly by [adding custom streaming logic](../how-tos/streaming.md#use-with-any-llm) with `StreamWriter`.
@@ -371,6 +373,7 @@ If your desired LLM isn't officially supported by LangChain, consider these opti
- [Force model to call a specific tool](https://python.langchain.com/docs/how_to/tool_choice/)
- [All chat model how-to guides](https://python.langchain.com/docs/how_to/#chat-models)
- [Chat model integrations](https://python.langchain.com/docs/integrations/chat/)
:::
:::js
@@ -381,4 +384,5 @@ If your desired LLM isn't officially supported by LangChain, consider these opti
- [Force model to call a specific tool](https://js.langchain.com/docs/how_to/tool_choice/)
- [All chat model how-to guides](https://js.langchain.com/docs/how_to/#chat-models)
- [Chat model integrations](https://js.langchain.com/docs/integrations/chat/)
:::
+1 -1
View File
@@ -244,7 +244,7 @@ output = agent.invoke(
print(output["messages"][-1].text())
```
!!! version-added "New in langgraph>=0.6"
!!! version-added "Added in version 0.6.0"
:::
@@ -14,7 +14,7 @@ from langgraph.checkpoint.base import (
CheckpointMetadata,
CheckpointTuple,
get_checkpoint_id,
get_checkpoint_metadata,
get_serializable_checkpoint_metadata,
)
from langgraph.checkpoint.serde.base import SerializerProtocol
from psycopg import Capabilities, Connection, Cursor, Pipeline
@@ -325,7 +325,7 @@ class PostgresSaver(BasePostgresSaver):
checkpoint["id"],
checkpoint_id,
Jsonb(copy),
Jsonb(get_checkpoint_metadata(config, metadata)),
Jsonb(get_serializable_checkpoint_metadata(config, metadata)),
),
)
return next_config
@@ -14,7 +14,7 @@ from langgraph.checkpoint.base import (
CheckpointMetadata,
CheckpointTuple,
get_checkpoint_id,
get_checkpoint_metadata,
get_serializable_checkpoint_metadata,
)
from langgraph.checkpoint.serde.base import SerializerProtocol
from psycopg import AsyncConnection, AsyncCursor, AsyncPipeline, Capabilities
@@ -283,7 +283,7 @@ class AsyncPostgresSaver(BasePostgresSaver):
checkpoint["id"],
checkpoint_id,
Jsonb(copy),
Jsonb(get_checkpoint_metadata(config, metadata)),
Jsonb(get_serializable_checkpoint_metadata(config, metadata)),
),
)
return next_config
@@ -1,7 +1,9 @@
from __future__ import annotations
import random
import warnings
from collections.abc import Sequence
from importlib.metadata import version as get_version
from typing import Any, Optional, cast
from langchain_core.runnables import RunnableConfig
@@ -16,6 +18,18 @@ from psycopg.types.json import Jsonb
MetadataInput = Optional[dict[str, Any]]
try:
major, minor = get_version("langgraph").split(".")[:2]
if int(major) == 0 and int(minor) < 5:
warnings.warn(
"You're using incompatible versions of langgraph and checkpoint-postgres. Please upgrade langgraph to avoid unexpected behavior.",
DeprecationWarning,
stacklevel=2,
)
except Exception:
# skip version check if running from source
pass
"""
To add a new migration, add a new string to the MIGRATIONS list.
The position of the migration in the list is the version number.
@@ -12,7 +12,7 @@ from langgraph.checkpoint.base import (
Checkpoint,
CheckpointMetadata,
CheckpointTuple,
get_checkpoint_metadata,
get_serializable_checkpoint_metadata,
)
from langgraph.checkpoint.serde.base import SerializerProtocol
from langgraph.checkpoint.serde.types import TASKS
@@ -441,7 +441,7 @@ class ShallowPostgresSaver(BasePostgresSaver):
thread_id,
checkpoint_ns,
Jsonb(copy),
Jsonb(get_checkpoint_metadata(config, metadata)),
Jsonb(get_serializable_checkpoint_metadata(config, metadata)),
),
)
return next_config
@@ -774,7 +774,7 @@ class AsyncShallowPostgresSaver(BasePostgresSaver):
thread_id,
checkpoint_ns,
Jsonb(copy),
Jsonb(get_checkpoint_metadata(config, metadata)),
Jsonb(get_serializable_checkpoint_metadata(config, metadata)),
),
)
return next_config
+2 -2
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph-checkpoint-postgres"
version = "2.0.24"
version = "2.0.25"
description = "Library with a Postgres implementation of LangGraph checkpoint saver."
authors = []
requires-python = ">=3.9"
@@ -12,7 +12,7 @@ readme = "README.md"
license = "MIT"
license-files = ['LICENSE']
dependencies = [
"langgraph-checkpoint>=2.0.21,<3.0.0",
"langgraph-checkpoint>=2.1.2,<3.0.0",
"orjson>=3.10.1",
"psycopg>=3.2.0",
"psycopg-pool>=3.2.0",
@@ -187,13 +187,11 @@ def test_data():
metadata_1: CheckpointMetadata = {
"source": "input",
"step": 2,
"writes": {},
"score": 1,
}
metadata_2: CheckpointMetadata = {
"source": "loop",
"step": 1,
"writes": {"foo": "bar"},
"score": None,
}
metadata_3: CheckpointMetadata = {}
@@ -220,7 +218,6 @@ async def test_combined_metadata(saver_name: str, test_data) -> None:
metadata: CheckpointMetadata = {
"source": "loop",
"step": 1,
"writes": {"foo": "bar"},
"score": None,
}
await saver.aput(config, chkpnt, metadata, {})
@@ -246,7 +243,6 @@ async def test_asearch(saver_name: str, test_data) -> None:
query_1 = {"source": "input"} # search by 1 key
query_2 = {
"step": 1,
"writes": {"foo": "bar"},
} # search by multiple keys
query_3: dict[str, Any] = {} # search by no keys, return all checkpoints
query_4 = {"source": "update", "step": 1} # no match
@@ -169,13 +169,11 @@ def test_data():
metadata_1: CheckpointMetadata = {
"source": "input",
"step": 2,
"writes": {},
"score": 1,
}
metadata_2: CheckpointMetadata = {
"source": "loop",
"step": 1,
"writes": {"foo": "bar"},
"score": None,
}
metadata_3: CheckpointMetadata = {}
@@ -202,7 +200,6 @@ def test_combined_metadata(saver_name: str, test_data) -> None:
metadata: CheckpointMetadata = {
"source": "loop",
"step": 1,
"writes": {"foo": "bar"},
"score": None,
}
saver.put(config, chkpnt, metadata, {})
@@ -228,7 +225,6 @@ def test_search(saver_name: str, test_data) -> None:
query_1 = {"source": "input"} # search by 1 key
query_2 = {
"step": 1,
"writes": {"foo": "bar"},
} # search by multiple keys
query_3: dict[str, Any] = {} # search by no keys, return all checkpoints
query_4 = {"source": "update", "step": 1} # no match
+2 -2
View File
@@ -245,7 +245,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "2.1.1"
version = "2.1.2"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -276,7 +276,7 @@ dev = [
[[package]]
name = "langgraph-checkpoint-postgres"
version = "2.0.24"
version = "2.0.25"
source = { editable = "." }
dependencies = [
{ name = "langgraph-checkpoint" },
+1 -1
View File
@@ -257,7 +257,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "2.1.1"
version = "2.1.2"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -404,6 +404,16 @@ def get_checkpoint_metadata(
return metadata
def get_serializable_checkpoint_metadata(
config: RunnableConfig, metadata: CheckpointMetadata
) -> CheckpointMetadata:
"""Get checkpoint metadata in a backwards-compatible manner."""
checkpoint_metadata = get_checkpoint_metadata(config, metadata)
if "writes" in checkpoint_metadata:
checkpoint_metadata.pop("writes")
return checkpoint_metadata
"""
Mapping from error type to error index.
Regular writes just map to their index in the list of writes being saved.
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph-checkpoint"
version = "2.1.1"
version = "2.1.2"
description = "Library with base interfaces for LangGraph checkpoint savers."
authors = []
requires-python = ">=3.9"
+1 -1
View File
@@ -273,7 +273,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "2.1.1"
version = "2.1.2"
source = { editable = "." }
dependencies = [
{ name = "langchain-core" },
+1 -1
View File
@@ -1 +1 @@
__version__ = "0.4.2"
__version__ = "0.4.3"
+18 -3
View File
@@ -271,7 +271,7 @@ For production use, requires a license key in env var LANGGRAPH_CLOUD_LICENSE_KE
f"""Ready!
- API: http://localhost:{port}
- Docs: http://localhost:{port}/docs
- LangGraph Studio: {debugger_origin}/studio/?baseUrl={debugger_base_url_query}
- LangSmith Debugger: {debugger_origin}/studio/?baseUrl={debugger_base_url_query}
"""
)
sys.stdout.flush()
@@ -652,11 +652,17 @@ def dockerfile(
help="Wait for a debugger client to connect to the debug port before starting the server",
default=False,
)
@click.option(
"--debugger-url",
type=str,
default=None,
help="URL of the LangSmith Debugger instance to connect to. Defaults to https://smith.langchain.com",
)
@click.option(
"--studio-url",
type=str,
default=None,
help="URL of the LangGraph Studio instance to connect to. Defaults to https://smith.langchain.com",
help="(Deprecated: use --debugger-url instead) URL of the LangSmith Debugger instance to connect to.",
)
@click.option(
"--allow-blocking",
@@ -692,12 +698,21 @@ def dev(
no_browser: bool,
debug_port: Optional[int],
wait_for_client: bool,
debugger_url: Optional[str],
studio_url: Optional[str],
allow_blocking: bool,
tunnel: bool,
server_log_level: str,
):
"""CLI entrypoint for running the LangGraph API server."""
if studio_url is not None:
click.secho(
"Warning: --studio-url is deprecated and will be removed in a future version. "
"Please use --debugger-url instead.",
fg="yellow",
)
if debugger_url is None:
debugger_url = studio_url
try:
from langgraph_api.cli import run_server # type: ignore
except ImportError:
@@ -761,7 +776,7 @@ def dev(
http=config_json.get("http"),
ui=config_json.get("ui"),
ui_config=config_json.get("ui_config"),
studio_url=studio_url,
studio_url=debugger_url,
allow_blocking=allow_blocking,
tunnel=tunnel,
server_level=server_log_level,
+6 -9
View File
@@ -89,13 +89,12 @@ def push_ui_message(
The created UI message.
Example:
.. code-block:: python
```python
push_ui_message(
name="component-name",
props={"content": "Hello world"},
)
```
"""
from langgraph._internal._constants import CONFIG_KEY_SEND
@@ -146,10 +145,9 @@ def delete_ui_message(id: str, *, state_key: str = "ui") -> RemoveUIMessage:
The remove UI message.
Example:
.. code-block:: python
```python
delete_ui_message("message-123")
```
"""
from langgraph._internal._constants import CONFIG_KEY_SEND
@@ -183,13 +181,12 @@ def ui_message_reducer(
Combined list of UI messages with removals applied.
Example:
.. code-block:: python
```python
messages = ui_message_reducer(
[{"type": "ui", "id": "1", "name": "Chat", "props": {}}],
{"type": "remove-ui", "id": "1"},
)
```
"""
if not isinstance(left, list):
+97 -29
View File
@@ -242,14 +242,13 @@ class PregelLoop:
self.interrupt_before = interrupt_before
self.manager = manager
self.is_nested = CONFIG_KEY_TASK_ID in self.config.get(CONF, {})
self.skip_done_tasks = CONFIG_KEY_CHECKPOINT_ID not in config[CONF] or (
CONFIG_KEY_RESUMING in self.config[CONF] and self.is_nested
)
self.skip_done_tasks = CONFIG_KEY_CHECKPOINT_ID not in config[CONF]
self._migrate_checkpoint = migrate_checkpoint
self.trigger_to_nodes = trigger_to_nodes
self.retry_policy = retry_policy
self.cache_policy = cache_policy
self.durability = durability
self.skipped_task_ids: set[str] = set()
if self.stream is not None and CONFIG_KEY_STREAM in config[CONF]:
self.stream = DuplexStream(self.stream, config[CONF][CONFIG_KEY_STREAM])
scratchpad: PregelScratchpad | None = config[CONF].get(CONFIG_KEY_SCRATCHPAD)
@@ -317,15 +316,35 @@ class PregelLoop:
]
writes_to_save: WritesT = [
w[1:] for w in self.checkpoint_pending_writes if w[0] == task_id
] + list(writes)
] + [(c, v) for c, v in writes if c != RESUME]
self.checkpoint_pending_writes.extend((task_id, c, v) for c, v in writes)
else:
# remove existing writes for this task
# build map of existing interrupts for this task: interrupt id -> list of interrupts
existing_interrupts_by_id: dict[str, list[Any]] = {
v[0].id: v
for tid, ch, v in self.checkpoint_pending_writes
if tid == task_id and ch == INTERRUPT
}
writes_to_save = []
for ch, v in writes:
if ch == INTERRUPT:
# we merge new interrupt writes with existing interrupts writes if they
# occurred within the same task (which means they have the same interrupt id)
new_interrupts = v if isinstance(v, list) else list(v)
if new_interrupts and (
existing := existing_interrupts_by_id.get(new_interrupts[0].id)
):
v = existing + new_interrupts
writes_to_save.append((ch, v))
else:
# we add non-interrupt writes as-is
writes_to_save.append((ch, v))
# replace all writes for this task_id with the merged writes
self.checkpoint_pending_writes = [
w for w in self.checkpoint_pending_writes if w[0] != task_id
]
writes_to_save = writes
# save writes
self.checkpoint_pending_writes.extend((task_id, c, v) for c, v in writes)
] + [(task_id, c, v) for c, v in writes_to_save]
if self.durability != "exit" and self.checkpointer_put_writes is not None:
config = patch_configurable(
self.checkpoint_config,
@@ -473,6 +492,22 @@ class PregelLoop:
cache_policy=self.cache_policy,
)
resume_map = self.config.get(CONF, {}).get(CONFIG_KEY_RESUME_MAP, {})
if resume_map or self.input is None:
# do not re-execute tasks that have unresumable interrupts
# i.e. when the graph is invoked with None, or the interrupt id is not in the resume map
skipped_interrupt_ids = self._pending_interrupts() - set(resume_map)
self.skipped_task_ids = {
task_id
for task_id, channel, value in self.checkpoint_pending_writes
if channel == INTERRUPT
# interrupts within a task are uncovered sequentially as resumes are provided,
# so we only need to check the last interrupt id
and value[-1].id in skipped_interrupt_ids
}
else:
self.skipped_task_ids = set()
# produce debug output
if self._checkpointer_put_after_previous is not None:
self._emit(
@@ -518,9 +553,45 @@ class PregelLoop:
if task.writes:
self.output_writes(task.id, task.writes, cached=True)
if self.skipped_task_ids:
# remove tasks with writes that have been matched with previous pending writes
self.skipped_task_ids = {
task_id
for task_id in self.skipped_task_ids
if not self.tasks[task_id].writes
}
# output interrupt writes for blocked tasks so they are still visible in the stream
for task_id, channel, value in self.checkpoint_pending_writes:
if task_id in self.skipped_task_ids and channel == INTERRUPT:
# find resume count for this task
resumes = next(
(
v
for tid, ch, v in self.checkpoint_pending_writes
if tid == task_id and ch == RESUME
),
None,
)
resume_count = len(resumes) if resumes is not None else 0
# only output unresumed interrupts
if resume_count < len(value):
self.output_writes(task_id, [(INTERRUPT, value[resume_count:])])
return True
def after_tick(self) -> None:
if self.skipped_task_ids:
# raise early GraphInterrupt for skipped tasks.
# since we know len(resumes) < len(interrupts) for these tasks, we
# can prevent unnecessary node re-execution by raising early
interrupts = []
for task_id, channel, value in self.checkpoint_pending_writes:
if channel == INTERRUPT and task_id in self.skipped_task_ids:
interrupts.extend(value)
if interrupts:
raise GraphInterrupt(interrupts)
self.skipped_task_ids.clear()
# finish superstep
writes = [w for t in self.tasks.values() for w in t.writes]
# all tasks have finished
@@ -572,30 +643,26 @@ class PregelLoop:
def _pending_interrupts(self) -> set[str]:
"""Return the set of interrupt ids that are pending without corresponding resume values."""
# mapping of task ids to interrupt ids
pending_interrupts: dict[str, str] = {}
# mapping of task ids to (interrupt_id, interrupt_count)
pending_interrupts: dict[str, tuple[str, int]] = {}
# mapping of task ids to resume count
pending_resumes: dict[str, int] = {}
# set of resume task ids
pending_resumes: set[str] = set()
for task_id, channel, value in self.checkpoint_pending_writes:
if channel == INTERRUPT:
pending_interrupts[task_id] = (
value[0].id,
len(value),
)
elif channel == RESUME:
resume_list = value if isinstance(value, list) else [value]
pending_resumes[task_id] = len(resume_list)
for task_id, write_type, value in self.checkpoint_pending_writes:
if write_type == INTERRUPT:
# interrupts is always a list, but there should only be one element
pending_interrupts[task_id] = value[0].id
elif write_type == RESUME:
pending_resumes.add(task_id)
resumed_interrupt_ids = {
pending_interrupts[task_id]
for task_id in pending_resumes
if task_id in pending_interrupts
}
# Keep only interrupts whose interrupt_id is not resumed
# keep only interrupt ids where resume_count < interrupt_count
hanging_interrupts: set[str] = {
interrupt_id
for interrupt_id in pending_interrupts.values()
if interrupt_id not in resumed_interrupt_ids
for task_id, (interrupt_id, interrupt_count) in pending_interrupts.items()
if pending_resumes.get(task_id, 0) < interrupt_count
}
return hanging_interrupts
@@ -1028,6 +1095,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
def put_writes(self, task_id: str, writes: WritesT) -> None:
"""Put writes for a task, to be read by the next tick."""
super().put_writes(task_id, writes)
if not writes or self.cache is None or not hasattr(self, "tasks"):
return
+80 -44
View File
@@ -40,7 +40,7 @@ class TaskResultPayload(TypedDict):
name: str
error: str | None
interrupts: list[dict]
result: list[tuple[str, Any]]
result: dict[str, Any]
class CheckpointTask(TypedDict):
@@ -77,6 +77,38 @@ def map_debug_tasks(tasks: Iterable[PregelExecutableTask]) -> Iterator[TaskPaylo
}
def is_multiple_channel_write(value: Any) -> bool:
"""Return True if the payload already wraps multiple writes from the same channel."""
return (
isinstance(value, dict)
and "$writes" in value
and isinstance(value["$writes"], list)
)
def map_task_result_writes(writes: Sequence[tuple[str, Any]]) -> dict[str, Any]:
"""Folds task writes into a result dict and aggregates multiple writes to the same channel.
If the channel contains a single write, we record the write in the result dict as `{channel: write}`
If the channel contains multiple writes, we record the writes in the result dict as `{channel: {'$writes': [write1, write2, ...]}}`"""
result: dict[str, Any] = {}
for channel, value in writes:
existing = result.get(channel)
if existing is not None:
channel_writes = (
existing["$writes"]
if is_multiple_channel_write(existing)
else [existing]
)
channel_writes.append(value)
result[channel] = {"$writes": channel_writes}
else:
result[channel] = value
return result
def map_debug_task_results(
task_tup: tuple[PregelExecutableTask, Sequence[tuple[str, Any]]],
stream_keys: str | Sequence[str],
@@ -90,7 +122,9 @@ def map_debug_task_results(
"id": task.id,
"name": task.name,
"error": next((w[1] for w in writes if w[0] == ERROR), None),
"result": [w for w in writes if w[0] in stream_channels_list or w[0] == RETURN],
"result": map_task_result_writes(
[w for w in writes if w[0] in stream_channels_list or w[0] == RETURN]
),
"interrupts": [
asdict(v)
for w in writes
@@ -196,54 +230,56 @@ def tasks_w_writes(
),
MISSING,
)
task_error = next(
(exc for tid, n, exc in pending_writes if tid == task.id and n == ERROR),
None,
)
task_interrupts = tuple(
v
for tid, n, vv in pending_writes
if tid == task.id and n == INTERRUPT
for v in (vv if isinstance(vv, Sequence) else [vv])
)
task_writes = [
(chan, val)
for tid, chan, val in pending_writes
if tid == task.id and chan not in (ERROR, INTERRUPT, RETURN)
]
if rtn is not MISSING:
task_result = rtn
elif isinstance(output_keys, str):
# unwrap single channel writes to just the write value
filtered_writes = [
(chan, val) for chan, val in task_writes if chan == output_keys
]
mapped_writes = map_task_result_writes(filtered_writes)
task_result = mapped_writes.get(str(output_keys)) if mapped_writes else None
else:
if isinstance(output_keys, str):
output_keys = [output_keys]
# map task result writes to the desired output channels
# repeateed writes to the same channel are aggregated into: {'$writes': [write1, write2, ...]}
filtered_writes = [
(chan, val) for chan, val in task_writes if chan in output_keys
]
mapped_writes = map_task_result_writes(filtered_writes)
task_result = mapped_writes if filtered_writes else {}
has_writes = rtn is not MISSING or any(
w[0] == task.id and w[1] not in (ERROR, INTERRUPT) for w in pending_writes
)
out.append(
PregelTask(
task.id,
task.name,
task.path,
next(
(
exc
for tid, n, exc in pending_writes
if tid == task.id and n == ERROR
),
None,
),
tuple(
v
for tid, n, vv in pending_writes
if tid == task.id and n == INTERRUPT
for v in (vv if isinstance(vv, Sequence) else [vv])
),
task_error,
task_interrupts,
states.get(task.id) if states else None,
(
rtn
if rtn is not MISSING
else next(
(
val
for tid, chan, val in pending_writes
if tid == task.id and chan == output_keys
),
None,
)
if isinstance(output_keys, str)
else {
chan: val
for tid, chan, val in pending_writes
if tid == task.id
and (
chan == output_keys
if isinstance(output_keys, str)
else chan in output_keys
)
}
)
if any(
w[0] == task.id and w[1] not in (ERROR, INTERRUPT)
for w in pending_writes
)
else None,
task_result if has_writes else None,
)
)
return tuple(out)
+47 -22
View File
@@ -1695,12 +1695,10 @@ class Pregel(
return patch_checkpoint_map(next_config, saved.metadata)
# apply pending writes, if not on specific checkpoint
if (
CONFIG_KEY_CHECKPOINT_ID not in config[CONF]
and saved is not None
and saved.pending_writes
):
# task ids can be provided in the StateUpdate, but if not,
# we use the task id generated by prepare_next_tasks
node_to_task_ids: dict[str, deque[str]] = defaultdict(deque)
if saved is not None and saved.pending_writes is not None:
# tasks for this checkpoint
next_tasks = prepare_next_tasks(
checkpoint,
@@ -1716,6 +1714,10 @@ class Pregel(
checkpointer=checkpointer,
manager=None,
)
# collect task ids to reuse so we can properly attach task results
for t in next_tasks.values():
node_to_task_ids[t.name].append(t.id)
# apply null writes
if null_writes := [
w[1:] for w in saved.pending_writes or [] if w[0] == NULL_TASK_ID
@@ -1797,8 +1799,14 @@ class Pregel(
raise InvalidUpdateError(f"Node {as_node} has no writers")
writes: deque[tuple[str, Any]] = deque()
task = PregelTaskWrites((), as_node, writes, [INTERRUPT])
task_id = provided_task_id or str(
uuid5(UUID(checkpoint["id"]), INTERRUPT)
# get the task ids that were prepared for this node
# if a task id was provided in the StateUpdate, we use it
# otherwise, we use the next available task id
prepared_task_ids = node_to_task_ids.get(as_node, deque())
task_id = provided_task_id or (
prepared_task_ids.popleft()
if prepared_task_ids
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
)
run_tasks.append(task)
run_task_ids.append(task_id)
@@ -2151,12 +2159,11 @@ class Pregel(
return patch_checkpoint_map(
next_config, saved.metadata if saved else None
)
# apply pending writes, if not on specific checkpoint
if (
CONFIG_KEY_CHECKPOINT_ID not in config[CONF]
and saved is not None
and saved.pending_writes
):
# task ids can be provided in the StateUpdate, but if not,
# we use the task id generated by prepare_next_tasks
node_to_task_ids: dict[str, deque[str]] = defaultdict(deque)
if saved is not None and saved.pending_writes is not None:
# tasks for this checkpoint
next_tasks = prepare_next_tasks(
checkpoint,
@@ -2172,6 +2179,10 @@ class Pregel(
checkpointer=checkpointer,
manager=None,
)
# collect task ids to reuse so we can properly attach task results
for t in next_tasks.values():
node_to_task_ids[t.name].append(t.id)
# apply null writes
if null_writes := [
w[1:] for w in saved.pending_writes or [] if w[0] == NULL_TASK_ID
@@ -2248,8 +2259,14 @@ class Pregel(
raise InvalidUpdateError(f"Node {as_node} has no writers")
writes: deque[tuple[str, Any]] = deque()
task = PregelTaskWrites((), as_node, writes, [INTERRUPT])
task_id = provided_task_id or str(
uuid5(UUID(checkpoint["id"]), INTERRUPT)
# get the task ids that were prepared for this node
# if a task id was provided in the StateUpdate, we use it
# otherwise, we use the next available task id
prepared_task_ids = node_to_task_ids.get(as_node, deque())
task_id = provided_task_id or (
prepared_task_ids.popleft()
if prepared_task_ids
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
)
run_tasks.append(task)
run_task_ids.append(task_id)
@@ -2445,7 +2462,7 @@ class Pregel(
input: The input to the graph.
config: The configuration to use for the run.
context: The static context to use for the run.
!!! version-added "Added in version 0.6.0."
!!! version-added "Added in version 0.6.0"
stream_mode: The mode to stream output, defaults to `self.stream_mode`.
Options are:
@@ -2655,7 +2672,11 @@ class Pregel(
for task in loop.match_cached_writes():
loop.output_writes(task.id, task.writes, cached=True)
for _ in runner.tick(
[t for t in loop.tasks.values() if not t.writes],
[
t
for t in loop.tasks.values()
if not t.writes and t.id not in loop.skipped_task_ids
],
timeout=self.step_timeout,
get_waiter=get_waiter,
schedule_task=loop.accept_push,
@@ -2711,7 +2732,7 @@ class Pregel(
input: The input to the graph.
config: The configuration to use for the run.
context: The static context to use for the run.
!!! version-added "Added in version 0.6.0."
!!! version-added "Added in version 0.6.0"
stream_mode: The mode to stream output, defaults to `self.stream_mode`.
Options are:
@@ -2974,7 +2995,11 @@ class Pregel(
for task in await loop.amatch_cached_writes():
loop.output_writes(task.id, task.writes, cached=True)
async for _ in runner.atick(
[t for t in loop.tasks.values() if not t.writes],
[
t
for t in loop.tasks.values()
if not t.writes and t.id not in loop.skipped_task_ids
],
timeout=self.step_timeout,
get_waiter=get_waiter,
schedule_task=loop.aaccept_push,
@@ -3043,7 +3068,7 @@ class Pregel(
input: The input data for the graph. It can be a dictionary or any other type.
config: Optional. The configuration for the graph run.
context: The static context to use for the run.
!!! version-added "Added in version 0.6.0."
!!! version-added "Added in version 0.6.0"
stream_mode: Optional[str]. The stream mode for the graph run. Default is "values".
print_mode: Accepts the same values as `stream_mode`, but only prints the output to the console, for debugging purposes. Does not affect the output of the graph in any way.
output_keys: Optional. The output keys to retrieve from the graph run.
@@ -3128,7 +3153,7 @@ class Pregel(
input: The input data for the computation. It can be a dictionary or any other type.
config: Optional. The configuration for the computation.
context: The static context to use for the run.
!!! version-added "Added in version 0.6.0."
!!! version-added "Added in version 0.6.0"
stream_mode: Optional. The stream mode for the computation. Default is "values".
print_mode: Accepts the same values as `stream_mode`, but only prints the output to the console, for debugging purposes. Does not affect the output of the graph in any way.
output_keys: Optional. The output keys to include in the result. Default is None.
+3 -3
View File
@@ -106,7 +106,7 @@ else:
class RetryPolicy(NamedTuple):
"""Configuration for retrying nodes.
!!! version-added "Added in version 0.2.24."
!!! version-added "Added in version 0.2.24"
"""
initial_interval: float = 0.5
@@ -148,7 +148,7 @@ _DEFAULT_INTERRUPT_ID = "placeholder-id"
class Interrupt:
"""Information about an interrupt that occurred in a node.
!!! version-added "Added in version 0.2.24."
!!! version-added "Added in version 0.2.24"
!!! version-changed "Changed in version v0.4.0"
* `interrupt_id` was introduced as a property
@@ -349,7 +349,7 @@ N = TypeVar("N", bound=Hashable)
class Command(Generic[N], ToolOutputMixin):
"""One or more commands to update the graph's state and send messages to nodes.
!!! version-added "Added in version 0.2.24."
!!! version-added "Added in version 0.2.24"
Args:
graph: graph to send the command to. Supported values are:
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph"
version = "0.6.8"
version = "0.6.9"
description = "Building stateful, multi-actor applications with LLMs"
authors = []
requires-python = ">=3.9"
+574 -1
View File
@@ -1,12 +1,21 @@
import operator
import sys
from typing import Annotated
import pytest
from langgraph.checkpoint.base import BaseCheckpointSaver
from typing_extensions import TypedDict
from langgraph.graph import END, START, StateGraph
from langgraph.types import Durability
from langgraph.types import Command, Durability, Send, interrupt
pytestmark = pytest.mark.anyio
NEEDS_CONTEXTVARS = pytest.mark.skipif(
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
def test_interruption_without_state_updates(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
@@ -90,3 +99,567 @@ async def test_interruption_without_state_updates_async(
assert (await graph.aget_state(thread)).next == ()
n_checkpoints = len([c async for c in graph.aget_state_history(thread)])
assert n_checkpoints == (5 if durability != "exit" else 3)
def test_interrupt_with_send_payloads(sync_checkpointer: BaseCheckpointSaver) -> None:
"""Test interruption in map node with Send payloads and human-in-the-loop resume."""
# Global counter to track node executions
node_counter = {"entry": 0, "map_node": 0}
class State(TypedDict):
items: list[str]
processed: Annotated[list[str], operator.add]
def entry_node(state: State):
node_counter["entry"] += 1
return {} # No state updates in entry node
def send_to_map(state: State):
return [Send("map_node", {"item": item}) for item in state["items"]]
def map_node(state: State):
node_counter["map_node"] += 1
if "dangerous" in state["item"]:
value = interrupt({"processing": state["item"]})
return {"processed": [f"processed_{value}"]}
else:
return {"processed": [f"processed_{state['item']}_auto"]}
builder = StateGraph(State)
builder.add_node("entry", entry_node)
builder.add_node("map_node", map_node)
builder.add_edge(START, "entry")
builder.add_conditional_edges("entry", send_to_map, ["map_node"])
builder.add_edge("map_node", END)
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "test_interrupt_send"}}
# Run until interrupts
result = graph.invoke(
{"items": ["item1", "dangerous_item1", "dangerous_item2"]}, config=config
)
# Verify we have interrupts (only one for dangerous_item)
interrupts = result.get("__interrupt__", [])
assert len(interrupts) == 2
assert "dangerous_item" in interrupts[0].value["processing"]
# Resume with mapping of interrupt IDs to values
resume_map = {i.id: f"human_input_{i.value['processing']}" for i in interrupts}
final_result = graph.invoke(Command(resume=resume_map), config=config)
# Verify final result contains processed items
assert "processed" in final_result
processed_items = final_result["processed"]
assert len(processed_items) == 3
assert "processed_item1_auto" in processed_items # item1 processed automatically
assert any(
"processed_human_input_dangerous_item1" in item for item in processed_items
) # dangerous_item1 processed after interrupt
assert any(
"processed_human_input_dangerous_item2" in item for item in processed_items
) # dangerous_item2 processed after interrupt
# Verify node execution counts
assert node_counter["entry"] == 1 # Entry node runs once
# Map node runs 3 times initially (item1 completes, 2 dangerous_items interrupt),
# then 2 times on resume
assert node_counter["map_node"] == 5
@NEEDS_CONTEXTVARS
async def test_interrupt_with_send_payloads_async(
async_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
"""Test interruption in map node with Send payloads and human-in-the-loop resume."""
# Global counter to track node executions
node_counter = {"entry": 0, "map_node": 0}
class State(TypedDict):
items: list[str]
processed: Annotated[list[str], operator.add]
def entry_node(state: State):
node_counter["entry"] += 1
return {} # No state updates in entry node
def send_to_map(state: State):
return [Send("map_node", {"item": item}) for item in state["items"]]
def map_node(state: State):
node_counter["map_node"] += 1
if "dangerous" in state["item"]:
value = interrupt({"processing": state["item"]})
return {"processed": [f"processed_{value}"]}
else:
return {"processed": [f"processed_{state['item']}_auto"]}
builder = StateGraph(State)
builder.add_node("entry", entry_node)
builder.add_node("map_node", map_node)
builder.add_edge(START, "entry")
builder.add_conditional_edges("entry", send_to_map, ["map_node"])
builder.add_edge("map_node", END)
graph = builder.compile(checkpointer=async_checkpointer)
config = {"configurable": {"thread_id": "test_interrupt_send"}}
# Run until interrupts
result = await graph.ainvoke(
{"items": ["item1", "dangerous_item1", "dangerous_item2"]}, config=config
)
# Verify we have interrupts (only one for dangerous_item)
interrupts = result.get("__interrupt__", [])
assert len(interrupts) == 2
assert "dangerous_item" in interrupts[0].value["processing"]
# Resume with mapping of interrupt IDs to values
resume_map = {i.id: f"human_input_{i.value['processing']}" for i in interrupts}
final_result = await graph.ainvoke(Command(resume=resume_map), config=config)
# Verify final result contains processed items
assert "processed" in final_result
processed_items = final_result["processed"]
assert len(processed_items) == 3
assert "processed_item1_auto" in processed_items # item1 processed automatically
assert any(
"processed_human_input_dangerous_item1" in item for item in processed_items
) # dangerous_item1 processed after interrupt
assert any(
"processed_human_input_dangerous_item2" in item for item in processed_items
) # dangerous_item2 processed after interrupt
# Verify node execution counts
assert node_counter["entry"] == 1 # Entry node runs once
# Map node runs 3 times initially (item1 completes, 2 dangerous_items interrupt),
# then 2 times on resume
assert node_counter["map_node"] == 5
def test_interrupt_with_send_payloads_sequential_resume(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test interruption in map node with Send payloads and sequential resume."""
# Global counter to track node executions
node_counter = {"entry": 0, "map_node": 0}
class State(TypedDict):
items: list[str]
processed: Annotated[list[str], operator.add]
def entry_node(state: State):
node_counter["entry"] += 1
return {} # No state updates in entry node
def send_to_map(state: State):
return [Send("map_node", {"item": item}) for item in state["items"]]
def map_node(state: State):
node_counter["map_node"] += 1
if "dangerous" in state["item"]:
value = interrupt({"processing": state["item"]})
return {"processed": [f"processed_{value}"]}
else:
return {"processed": [f"processed_{state['item']}_auto"]}
builder = StateGraph(State)
builder.add_node("entry", entry_node)
builder.add_node("map_node", map_node)
builder.add_edge(START, "entry")
builder.add_conditional_edges("entry", send_to_map, ["map_node"])
builder.add_edge("map_node", END)
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "test_interrupt_send_sequential"}}
# Run until interrupts
result = graph.invoke(
{"items": ["item1", "dangerous_item1", "dangerous_item2"]}, config=config
)
# Verify we have interrupts
interrupts = result.get("__interrupt__", [])
assert len(interrupts) == 2
assert "dangerous_item" in interrupts[0].value["processing"]
# Resume first interrupt only
first_interrupt = interrupts[0]
first_resume_map = {
first_interrupt.id: f"human_input_{first_interrupt.value['processing']}"
}
partial_result = graph.invoke(Command(resume=first_resume_map), config=config)
# Verify we still have one pending interrupt
remaining_interrupts = partial_result.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
# Resume second interrupt
second_interrupt = remaining_interrupts[0]
second_resume_map = {
second_interrupt.id: f"human_input_{second_interrupt.value['processing']}"
}
final_result = graph.invoke(Command(resume=second_resume_map), config=config)
# Verify final result contains processed items
assert "processed" in final_result
processed_items = final_result["processed"]
assert len(processed_items) == 3
assert "processed_item1_auto" in processed_items # item1 processed automatically
assert any(
"processed_human_input_dangerous_item1" in item for item in processed_items
) # dangerous_item1 processed after interrupt
assert any(
"processed_human_input_dangerous_item2" in item for item in processed_items
) # dangerous_item2 processed after interrupt
# Verify node execution counts
assert node_counter["entry"] == 1 # Entry node runs once
# Map node runs 3 times initially (item1 completes, 2 dangerous_items interrupt),
# then 1 time on first resume, then 1 time on second resume
assert node_counter["map_node"] == 5
@NEEDS_CONTEXTVARS
async def test_interrupt_with_send_payloads_sequential_resume_async(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test interruption in map node with Send payloads and sequential resume."""
# Global counter to track node executions
node_counter = {"entry": 0, "map_node": 0}
class State(TypedDict):
items: list[str]
processed: Annotated[list[str], operator.add]
def entry_node(state: State):
node_counter["entry"] += 1
return {} # No state updates in entry node
def send_to_map(state: State):
return [Send("map_node", {"item": item}) for item in state["items"]]
def map_node(state: State):
node_counter["map_node"] += 1
if "dangerous" in state["item"]:
value = interrupt({"processing": state["item"]})
return {"processed": [f"processed_{value}"]}
else:
return {"processed": [f"processed_{state['item']}_auto"]}
builder = StateGraph(State)
builder.add_node("entry", entry_node)
builder.add_node("map_node", map_node)
builder.add_edge(START, "entry")
builder.add_conditional_edges("entry", send_to_map, ["map_node"])
builder.add_edge("map_node", END)
graph = builder.compile(checkpointer=async_checkpointer)
config = {"configurable": {"thread_id": "test_interrupt_send_sequential"}}
# Run until interrupts
result = await graph.ainvoke(
{"items": ["item1", "dangerous_item1", "dangerous_item2"]}, config=config
)
# Verify we have interrupts
interrupts = result.get("__interrupt__", [])
assert len(interrupts) == 2
assert "dangerous_item" in interrupts[0].value["processing"]
# Resume first interrupt only
first_interrupt = interrupts[0]
first_resume_map = {
first_interrupt.id: f"human_input_{first_interrupt.value['processing']}"
}
partial_result = await graph.ainvoke(
Command(resume=first_resume_map), config=config
)
# Verify we still have one pending interrupt
remaining_interrupts = partial_result.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
# Resume second interrupt
second_interrupt = remaining_interrupts[0]
second_resume_map = {
second_interrupt.id: f"human_input_{second_interrupt.value['processing']}"
}
final_result = await graph.ainvoke(Command(resume=second_resume_map), config=config)
# Verify final result contains processed items
assert "processed" in final_result
processed_items = final_result["processed"]
assert len(processed_items) == 3
assert "processed_item1_auto" in processed_items # item1 processed automatically
assert any(
"processed_human_input_dangerous_item1" in item for item in processed_items
) # dangerous_item1 processed after interrupt
assert any(
"processed_human_input_dangerous_item2" in item for item in processed_items
) # dangerous_item2 processed after interrupt
# Verify node execution counts
assert node_counter["entry"] == 1 # Entry node runs once
# Map node runs 3 times initially (item1 completes, 2 dangerous_items interrupt),
# then 1 time on first resume, then 1 time on second resume
assert node_counter["map_node"] == 5
def test_node_with_multiple_interrupts_requires_full_resume(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test a number of different resume patterns for a node with multiple interrupts,
Ensures that a node is not re-executed until valid resume values have been provided to all
discovered interrupts"""
node_counter = 0
class State(TypedDict):
input: str
def double_interrupt_node(state: State):
nonlocal node_counter
node_counter += 1
first = interrupt("first")
second = interrupt("second")
third = interrupt("third")
return {"input": f"{first}-{second}-{third}"}
builder = StateGraph(State)
builder.add_node("double_interrupt", double_interrupt_node)
builder.add_edge(START, "double_interrupt")
builder.add_edge("double_interrupt", END)
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "test_double_interrupt"}}
result = graph.invoke({"input": "start"}, config=config)
interrupts = result.get("__interrupt__", [])
assert len(interrupts) == 1
first_interrupt = interrupts[0]
assert node_counter == 1
# invoke with an interrupt map that matches double_interrupt_node.
# this should execute the node
partial = graph.invoke(
Command(resume={first_interrupt.id: "human_first"}), config=config
)
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert remaining_interrupts[0].value == "second"
assert node_counter == 2
# invoke with an interrupt map that DOES NOT match double_interrupt_node.
# this should not execute the node because the optimization kicks in
partial = graph.invoke(
Command(resume={"00000000000000000000000000000000": "nothing_burger"}),
config=config,
)
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert remaining_interrupts[0].value == "second"
assert node_counter == 2
# invoke with None resume. this should NOT execute the node
partial = graph.invoke(None, config=config)
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert remaining_interrupts[0].value == "second"
assert node_counter == 2
# invoke with nonspecific resume. this should execute the node
partial = graph.invoke(Command(resume="human_second"), config=config)
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
print("REMAINING INTERRUPTS: ", remaining_interrupts)
assert remaining_interrupts[0].value == "third"
assert node_counter == 3
# finally, invoke with an interrupt map that matches double_interrupt_node.
# this should execute the node and all interrupts should be resolved
final_result = graph.invoke(Command(resume="human_third"), config=config)
assert "input" in final_result
assert final_result["input"] == "human_first-human_second-human_third"
assert node_counter == 4
@NEEDS_CONTEXTVARS
async def test_node_with_multiple_interrupts_requires_full_resume_async(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test a number of different resume patterns for a node with multiple interrupts,
Ensures that a node is not re-executed until valid resume values have been provided to all
discovered interrupts"""
node_counter = 0
class State(TypedDict):
input: str
def double_interrupt_node(state: State):
nonlocal node_counter
node_counter += 1
first = interrupt("first")
second = interrupt("second")
third = interrupt("third")
return {"input": f"{first}-{second}-{third}"}
builder = StateGraph(State)
builder.add_node("double_interrupt", double_interrupt_node)
builder.add_edge(START, "double_interrupt")
builder.add_edge("double_interrupt", END)
graph = builder.compile(checkpointer=async_checkpointer)
config = {"configurable": {"thread_id": "test_double_interrupt"}}
result = await graph.ainvoke({"input": "start"}, config=config)
interrupts = result.get("__interrupt__", [])
assert len(interrupts) == 1
first_interrupt = interrupts[0]
assert node_counter == 1
# invoke with an interrupt map that matches double_interrupt_node.
# this should execute the node
partial = await graph.ainvoke(
Command(resume={first_interrupt.id: "human_first"}), config=config
)
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert remaining_interrupts[0].value == "second"
assert node_counter == 2
# invoke with an interrupt map that DOES NOT match double_interrupt_node.
# this should not execute the node because the optimization kicks in
partial = await graph.ainvoke(
Command(resume={"00000000000000000000000000000000": "nothing_burger"}),
config=config,
)
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert remaining_interrupts[0].value == "second"
assert node_counter == 2
# invoke with None resume. this should NOT execute the node
partial = await graph.ainvoke(None, config=config)
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert remaining_interrupts[0].value == "second"
assert node_counter == 2
# invoke with nonspecific resume. this should execute the node
partial = await graph.ainvoke(Command(resume="human_second"), config=config)
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
print("REMAINING INTERRUPTS: ", remaining_interrupts)
assert remaining_interrupts[0].value == "third"
assert node_counter == 3
# finally, invoke with an interrupt map that matches double_interrupt_node.
# this should execute the node and all interrupts should be resolved
final_result = await graph.ainvoke(Command(resume="human_third"), config=config)
assert "input" in final_result
assert final_result["input"] == "human_first-human_second-human_third"
assert node_counter == 4
def test_invoke_interrupted_graph_with_none(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test that invoking an interrupted graph with None does not duplicate interrupt writes"""
node_counter = 0
class State(TypedDict):
input: str
def double_interrupt_node(state: State):
nonlocal node_counter
node_counter += 1
first = interrupt("first")
second = interrupt("second")
return {"input": f"{first}-{second}"}
builder = StateGraph(State)
builder.add_node("double_interrupt", double_interrupt_node)
builder.add_edge(START, "double_interrupt")
builder.add_edge("double_interrupt", END)
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "test_none_resume"}}
result = graph.invoke({"input": "start"}, config=config)
first_history = list(graph.get_state_history(config))
interrupts = result.get("__interrupt__", [])
assert len(interrupts) == 1
assert node_counter == 1
# invoke with None. this should NOT execute the node and the history should
# look the same as the first run
partial = graph.invoke(None, config=config)
second_history = list(graph.get_state_history(config))
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert remaining_interrupts[0].value == "first"
assert node_counter == 1
# history should look the same for tasks and interrupts
print("first_history[0].interrupts: ", first_history[0].interrupts)
print("second_history[0].interrupts: ", second_history[0].interrupts)
print("first_history[0].tasks: ", first_history[0].tasks)
print("second_history[0].tasks: ", second_history[0].tasks)
assert first_history[0].interrupts == second_history[0].interrupts
assert first_history[0].tasks == second_history[0].tasks
# now resume the first interrupt with some value
partial = graph.invoke(Command(resume="weet"), config=config)
print("partial 3", partial)
third_history = list(graph.get_state_history(config))
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert remaining_interrupts[0].value == "second"
assert node_counter == 2
# invoke with None again. the history should look the same as
# the third run
partial = graph.invoke(None, config=config)
print("partial 4", partial)
fourth_history = list(graph.get_state_history(config))
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 1
assert node_counter == 2
print("\nthird_history[0].interrupts: ", third_history[0].interrupts)
print("fourth_history[0].interrupts: ", fourth_history[0].interrupts)
print("third_history[0].tasks: ", third_history[0].tasks)
print("fourth_history[0].tasks: ", fourth_history[0].tasks)
assert third_history[0].interrupts == fourth_history[0].interrupts
assert third_history[0].tasks == fourth_history[0].tasks
# resume the graph once more with a real value
partial = graph.invoke(Command(resume="bix"), config=config)
remaining_interrupts = partial.get("__interrupt__", [])
assert len(remaining_interrupts) == 0
assert node_counter == 3
+13 -5
View File
@@ -4023,7 +4023,9 @@ def test_in_one_fan_out_out_one_graph_state() -> None:
"payload": {
"id": AnyStr(),
"name": "rewrite_query",
"result": [("query", "query: what is weather in sf")],
"result": {
"query": "query: what is weather in sf",
},
"error": None,
"interrupts": [],
},
@@ -4071,7 +4073,9 @@ def test_in_one_fan_out_out_one_graph_state() -> None:
"payload": {
"id": AnyStr(),
"name": "retriever_two",
"result": [("docs", ["doc3", "doc4"])],
"result": {
"docs": ["doc3", "doc4"],
},
"error": None,
"interrupts": [],
},
@@ -4090,7 +4094,9 @@ def test_in_one_fan_out_out_one_graph_state() -> None:
"payload": {
"id": AnyStr(),
"name": "retriever_one",
"result": [("docs", ["doc1", "doc2"])],
"result": {
"docs": ["doc1", "doc2"],
},
"error": None,
"interrupts": [],
},
@@ -4130,7 +4136,9 @@ def test_in_one_fan_out_out_one_graph_state() -> None:
"payload": {
"id": AnyStr(),
"name": "qa",
"result": [("answer", "doc1,doc2,doc3,doc4")],
"result": {
"answer": "doc1,doc2,doc3,doc4",
},
"error": None,
"interrupts": [],
},
@@ -4718,7 +4726,7 @@ def test_send_dedupe_on_resume(
assert len(history) == (4 if durability != "exit" else 1)
# resume execution
assert graph.invoke(None, thread1, durability=durability) == [
assert graph.invoke(Command(resume=""), thread1, durability=durability) == [
"0",
"1",
"3.1",
+12 -4
View File
@@ -2567,7 +2567,9 @@ async def test_in_one_fan_out_out_one_graph_state() -> None:
"payload": {
"id": AnyStr(),
"name": "rewrite_query",
"result": [("query", "query: what is weather in sf")],
"result": {
"query": "query: what is weather in sf",
},
"error": None,
"interrupts": [],
},
@@ -2615,7 +2617,9 @@ async def test_in_one_fan_out_out_one_graph_state() -> None:
"payload": {
"id": AnyStr(),
"name": "retriever_two",
"result": [("docs", ["doc3", "doc4"])],
"result": {
"docs": ["doc3", "doc4"],
},
"error": None,
"interrupts": [],
},
@@ -2634,7 +2638,9 @@ async def test_in_one_fan_out_out_one_graph_state() -> None:
"payload": {
"id": AnyStr(),
"name": "retriever_one",
"result": [("docs", ["doc1", "doc2"])],
"result": {
"docs": ["doc1", "doc2"],
},
"error": None,
"interrupts": [],
},
@@ -2674,7 +2680,9 @@ async def test_in_one_fan_out_out_one_graph_state() -> None:
"payload": {
"id": AnyStr(),
"name": "qa",
"result": [("answer", "doc1,doc2,doc3,doc4")],
"result": {
"answer": "doc1,doc2,doc3,doc4",
},
"error": None,
"interrupts": [],
},
+182 -85
View File
@@ -2119,7 +2119,9 @@ def test_in_one_fan_out_state_graph_defer_node(
"id": AnyStr(),
"name": "rewrite_query",
"error": None,
"result": [("query", "query: what is weather in sf")],
"result": {
"query": "query: what is weather in sf",
},
"interrupts": [],
},
},
@@ -2153,7 +2155,9 @@ def test_in_one_fan_out_state_graph_defer_node(
"id": AnyStr(),
"name": "retriever_one",
"error": None,
"result": [("docs", ["doc1", "doc2"])],
"result": {
"docs": ["doc1", "doc2"],
},
"interrupts": [],
},
},
@@ -2165,7 +2169,9 @@ def test_in_one_fan_out_state_graph_defer_node(
"id": AnyStr(),
"name": "retriever_two",
"error": None,
"result": [("docs", ["doc3", "doc4"])],
"result": {
"docs": ["doc3", "doc4"],
},
"interrupts": [],
},
},
@@ -2191,7 +2197,9 @@ def test_in_one_fan_out_state_graph_defer_node(
"id": AnyStr(),
"name": "analyzer_one",
"error": None,
"result": [("query", "analyzed: query: what is weather in sf")],
"result": {
"query": "analyzed: query: what is weather in sf",
},
"interrupts": [],
},
},
@@ -2219,7 +2227,9 @@ def test_in_one_fan_out_state_graph_defer_node(
"id": AnyStr(),
"name": "qa",
"error": None,
"result": [("answer", "doc1,doc2,doc3,doc4")],
"result": {
"answer": "doc1,doc2,doc3,doc4",
},
"interrupts": [],
},
},
@@ -3445,73 +3455,6 @@ def test_stream_buffering_single_node(sync_checkpointer: BaseCheckpointSaver) ->
]
def test_nested_graph_resume_reuses_cached_task_writes(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
# Reproduces issue where a helper @task inside a nested graph re-executes
# on resume instead of reusing cached writes. Ensures it runs only once.
counter_parent = 0
counter_sub = 0
@task
def get_time_parent() -> float:
nonlocal counter_parent
counter_parent += 1
return time.time()
@task
def get_time_subgraph() -> float:
nonlocal counter_sub
counter_sub += 1
return time.time()
class State(TypedDict):
state_counter: int
# Subgraph that calls a helper task and then interrupts
sub = StateGraph(State)
def human_node(_: State):
_ = get_time_subgraph().result()
interrupt("what is your name?")
sub.add_node("human_node", human_node)
sub.set_entry_point("human_node")
sub.set_finish_point("human_node")
subgraph = sub.compile(checkpointer=sync_checkpointer)
# Parent graph that calls a helper task and interrupts, then enters subgraph
parent = StateGraph(State)
def parent_node(_: State):
_ = get_time_parent().result()
interrupt("what is your parent name?")
parent.add_node("parent_node", parent_node)
parent.add_node("subgraph", subgraph)
parent.add_edge(START, "parent_node")
parent.add_edge("parent_node", "subgraph")
parent.add_edge("subgraph", END)
graph = parent.compile(checkpointer=sync_checkpointer)
cfg_parent = {"configurable": {"thread_id": str(uuid.uuid4())}}
# First run – interrupts in parent node
for _ in graph.stream({"state_counter": 1}, cfg_parent):
pass
# Resume 1 – proceeds into subgraph, interrupts there
for _ in graph.stream(Command(resume="resume-1"), cfg_parent):
pass
# Resume 2 – completes without re-running subgraph helper task
for _ in graph.stream(Command(resume="resume-2"), cfg_parent):
pass
assert counter_parent == 1
assert counter_sub == 1
def test_nested_graph_interrupts_parallel(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
@@ -5606,12 +5549,9 @@ def test_falsy_return_from_task(sync_checkpointer: BaseCheckpointSaver):
"id": AnyStr(),
"interrupts": [],
"name": "falsy_task",
"result": [
(
"__return__",
False,
),
],
"result": {
"__return__": False,
},
},
"step": 0,
"timestamp": AnyStr(),
@@ -5628,7 +5568,7 @@ def test_falsy_return_from_task(sync_checkpointer: BaseCheckpointSaver):
},
],
"name": "graph",
"result": [],
"result": {},
},
"step": 0,
"timestamp": AnyStr(),
@@ -5714,12 +5654,9 @@ def test_falsy_return_from_task(sync_checkpointer: BaseCheckpointSaver):
"id": AnyStr(),
"interrupts": [],
"name": "graph",
"result": [
(
"__end__",
None,
),
],
"result": {
"__end__": None,
},
},
"step": 0,
"timestamp": AnyStr(),
@@ -8516,3 +8453,163 @@ def test_interrupt_stream_mode_values():
result = [*app.stream(State(), stream_mode="values")]
assert "__interrupt__" in result[-1]
def test_supersteps_populate_task_results(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
class State(TypedDict):
num: int
text: str
def double(state: State) -> State:
return {"num": state["num"] * 2, "text": state["text"] * 2}
graph = (
StateGraph(State)
.add_node("double", double)
.add_edge(START, "double")
.add_edge("double", END)
.compile(checkpointer=sync_checkpointer)
)
def first_task_result(history: list[StateSnapshot], node: str) -> Any:
for s in history:
for t in s.tasks:
if t.name == node:
return t.result
return None
# reference run with invoke
ref_cfg = {"configurable": {"thread_id": "ref"}}
graph.invoke({"num": 1, "text": "one"}, ref_cfg)
ref_history = list(graph.get_state_history(ref_cfg))
ref_start_result = first_task_result(ref_history, "__start__")
ref_double_result = first_task_result(ref_history, "double")
assert ref_start_result == {"num": 1, "text": "one"}
assert ref_double_result == {"num": 2, "text": "oneone"}
# using supersteps
bulk_cfg = {"configurable": {"thread_id": "bulk"}}
graph.bulk_update_state(
bulk_cfg,
[
[StateUpdate(values={}, as_node="__input__")],
[StateUpdate(values={"num": 1, "text": "one"}, as_node="__start__")],
[StateUpdate(values={"num": 2, "text": "oneone"}, as_node="double")],
],
)
bulk_history = list(graph.get_state_history(bulk_cfg))
bulk_start_result = first_task_result(bulk_history, "__start__")
bulk_double_result = first_task_result(bulk_history, "double")
assert bulk_start_result == ref_start_result == {"num": 1, "text": "one"}
assert bulk_double_result == ref_double_result == {"num": 2, "text": "oneone"}
def test_multiple_writes_same_channel_from_same_node(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
"""Test that a node can write multiple times to the same channel and that writes are ordered, reduced, and reflected in streamed events and state history."""
class State(TypedDict):
foo: Annotated[str, lambda a, b: ", ".join([x for x in [a, b] if x])]
def one(_: State) -> Command:
return Command(update=[("foo", "one.0"), ("foo", "one.1")])
def two(_: State) -> State:
return {"foo": "two"}
graph = (
StateGraph(State)
.add_node("one", one)
.add_node("two", two)
.add_edge(START, "one")
.add_edge("one", "two")
.add_edge("two", END)
.compile(checkpointer=sync_checkpointer)
)
config = {"configurable": {"thread_id": "1"}}
events = [
(ns, ev)
for ns, ev in graph.stream(
{"foo": "input"}, config, stream_mode=["updates", "tasks"]
)
]
assert events == [
(
"tasks",
{
"id": AnyStr(),
"name": "one",
"input": {"foo": "input"},
"triggers": ("branch:to:one",),
},
),
("updates", {"one": [{"foo": "one.0"}, {"foo": "one.1"}]}),
(
"tasks",
{
"id": AnyStr(),
"name": "one",
"error": None,
"result": {"foo": {"$writes": ["one.0", "one.1"]}},
"interrupts": [],
},
),
(
"tasks",
{
"id": AnyStr(),
"name": "two",
"input": {"foo": "input, one.0, one.1"},
"triggers": ("branch:to:two",),
},
),
("updates", {"two": {"foo": "two"}}),
(
"tasks",
{
"id": AnyStr(),
"name": "two",
"error": None,
"result": {"foo": "two"},
"interrupts": [],
},
),
]
def map_snapshot(s: StateSnapshot) -> dict:
return {
"tasks": [{"name": t.name, "result": t.result} for t in s.tasks],
"values": s.values,
}
history = [map_snapshot(s) for s in graph.get_state_history(config)]
assert history == [
{
"tasks": [],
"values": {"foo": "input, one.0, one.1, two"},
},
{
"tasks": [{"name": "two", "result": {"foo": "two"}}],
"values": {"foo": "input, one.0, one.1"},
},
{
"tasks": [
{"name": "one", "result": {"foo": {"$writes": ["one.0", "one.1"]}}}
],
"values": {"foo": "input"},
},
{
"tasks": [{"name": "__start__", "result": {"foo": "input"}}],
"values": {"foo": ""},
},
]
+56 -1
View File
@@ -2545,7 +2545,7 @@ async def test_send_dedupe_on_resume(
assert builder.nodes["2"].runnable.func.ticks == 3
assert builder.nodes["flaky"].runnable.func.ticks == 1
# resume execution
assert await graph.ainvoke(None, thread1, durability=durability) == [
assert await graph.ainvoke(Command(resume=""), thread1, durability=durability) == [
"0",
"1",
"3.1",
@@ -9211,3 +9211,58 @@ async def test_astream_waiter_cleanup_on_cancel(
assert recorded_tasks, "expected stream.wait() task to be created"
assert set(finished_tasks) == set(recorded_tasks)
assert all(t.done() for t in recorded_tasks)
async def test_supersteps_populate_task_results(
async_checkpointer: BaseCheckpointSaver,
) -> None:
class State(TypedDict):
num: int
text: str
def double(state: State) -> State:
return {"num": state["num"] * 2, "text": state["text"] * 2}
graph = (
StateGraph(State)
.add_node("double", double)
.add_edge(START, "double")
.add_edge("double", END)
.compile(checkpointer=async_checkpointer)
)
# reference run with ainvoke
ref_cfg = {"configurable": {"thread_id": "ref"}}
await graph.ainvoke({"num": 1, "text": "one"}, ref_cfg)
ref_history = [h async for h in graph.aget_state_history(ref_cfg)]
# Helper: pull first task result for a node name from history
def first_task_result(history: list[StateSnapshot], node: str) -> Any:
for s in history:
for t in s.tasks:
if t.name == node:
return t.result
return None
ref_start_result = first_task_result(ref_history, "__start__")
ref_double_result = first_task_result(ref_history, "double")
assert ref_start_result == {"num": 1, "text": "one"}
assert ref_double_result == {"num": 2, "text": "oneone"}
# using supersteps
bulk_cfg = {"configurable": {"thread_id": "bulk"}}
await graph.abulk_update_state(
bulk_cfg,
[
[StateUpdate(values={}, as_node="__input__")],
[StateUpdate(values={"num": 1, "text": "one"}, as_node="__start__")],
[StateUpdate(values={"num": 2, "text": "oneone"}, as_node="double")],
],
)
bulk_history = [h async for h in graph.aget_state_history(bulk_cfg)]
bulk_start_result = first_task_result(bulk_history, "__start__")
bulk_double_result = first_task_result(bulk_history, "double")
assert bulk_start_result == ref_start_result == {"num": 1, "text": "one"}
assert bulk_double_result == ref_double_result == {"num": 2, "text": "oneone"}
+3 -3
View File
@@ -1428,7 +1428,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "0.6.8"
version = "0.6.9"
source = { editable = "." }
dependencies = [
{ name = "langchain-core" },
@@ -1542,7 +1542,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "2.1.1"
version = "2.1.2"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -1573,7 +1573,7 @@ dev = [
[[package]]
name = "langgraph-checkpoint-postgres"
version = "2.0.24"
version = "2.0.25"
source = { editable = "../checkpoint-postgres" }
dependencies = [
{ name = "langgraph-checkpoint" },
+3 -3
View File
@@ -257,7 +257,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "0.6.8"
version = "0.6.9"
source = { editable = "../langgraph" }
dependencies = [
{ name = "langchain-core" },
@@ -309,7 +309,7 @@ dev = [
[[package]]
name = "langgraph-checkpoint"
version = "2.1.1"
version = "2.1.2"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -340,7 +340,7 @@ dev = [
[[package]]
name = "langgraph-checkpoint-postgres"
version = "2.0.24"
version = "2.0.25"
source = { editable = "../checkpoint-postgres" }
dependencies = [
{ name = "langgraph-checkpoint" },
+14 -14
View File
@@ -873,7 +873,7 @@ class AssistantsClient:
config: Configuration to use for the graph.
metadata: Metadata to add to assistant.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
assistant_id: Assistant ID to use, will default to a random UUID if not provided.
if_exists: How to handle duplicate creation. Defaults to 'raise' under the hood.
Must be either 'raise' (raise error if duplicate), or 'do_nothing' (return existing assistant).
@@ -944,7 +944,7 @@ class AssistantsClient:
The graph ID is normally set in your langgraph.json configuration. If None, assistant will keep pointing to same graph.
config: Configuration to use for the graph.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
metadata: Metadata to merge with existing assistant metadata.
name: The new name for the assistant.
headers: Optional custom headers to include with the request.
@@ -1964,7 +1964,7 @@ class RunsClient:
metadata: Metadata to assign to the run.
config: The configuration for the assistant.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
checkpoint: The checkpoint to resume from.
checkpoint_during: (deprecated) Whether to checkpoint during the run (or only at the end/interruption).
interrupt_before: Nodes to interrupt immediately before they get executed.
@@ -2174,7 +2174,7 @@ class RunsClient:
metadata: Metadata to assign to the run.
config: The configuration for the assistant.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
checkpoint: The checkpoint to resume from.
checkpoint_during: (deprecated) Whether to checkpoint during the run (or only at the end/interruption).
interrupt_before: Nodes to interrupt immediately before they get executed.
@@ -2422,7 +2422,7 @@ class RunsClient:
metadata: Metadata to assign to the run.
config: The configuration for the assistant.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
checkpoint: The checkpoint to resume from.
checkpoint_during: (deprecated) Whether to checkpoint during the run (or only at the end/interruption).
interrupt_before: Nodes to interrupt immediately before they get executed.
@@ -2883,7 +2883,7 @@ class CronClient:
metadata: Metadata to assign to the cron job runs.
config: The configuration for the assistant.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
checkpoint_during: Whether to checkpoint during the run (or only at the end/interruption).
interrupt_before: Nodes to interrupt immediately before they get executed.
@@ -2965,7 +2965,7 @@ class CronClient:
metadata: Metadata to assign to the cron job runs.
config: The configuration for the assistant.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
checkpoint_during: Whether to checkpoint during the run (or only at the end/interruption).
interrupt_before: Nodes to interrupt immediately before they get executed.
interrupt_after: Nodes to Nodes to interrupt immediately after they get executed.
@@ -4131,7 +4131,7 @@ class SyncAssistantsClient:
graph_id: The ID of the graph the assistant should use. The graph ID is normally set in your langgraph.json configuration.
config: Configuration to use for the graph.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
metadata: Metadata to add to assistant.
assistant_id: Assistant ID to use, will default to a random UUID if not provided.
if_exists: How to handle duplicate creation. Defaults to 'raise' under the hood.
@@ -4203,7 +4203,7 @@ class SyncAssistantsClient:
The graph ID is normally set in your langgraph.json configuration. If None, assistant will keep pointing to same graph.
config: Configuration to use for the graph.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
metadata: Metadata to merge with existing assistant metadata.
name: The new name for the assistant.
headers: Optional custom headers to include with the request.
@@ -5198,7 +5198,7 @@ class SyncRunsClient:
metadata: Metadata to assign to the run.
config: The configuration for the assistant.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
checkpoint: The checkpoint to resume from.
checkpoint_during: (deprecated) Whether to checkpoint during the run (or only at the end/interruption).
interrupt_before: Nodes to interrupt immediately before they get executed.
@@ -5404,7 +5404,7 @@ class SyncRunsClient:
metadata: Metadata to assign to the run.
config: The configuration for the assistant.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
checkpoint: The checkpoint to resume from.
checkpoint_during: (deprecated) Whether to checkpoint during the run (or only at the end/interruption).
interrupt_before: Nodes to interrupt immediately before they get executed.
@@ -5652,7 +5652,7 @@ class SyncRunsClient:
metadata: Metadata to assign to the run.
config: The configuration for the assistant.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
checkpoint: The checkpoint to resume from.
checkpoint_during: (deprecated) Whether to checkpoint during the run (or only at the end/interruption).
interrupt_before: Nodes to interrupt immediately before they get executed.
@@ -6093,7 +6093,7 @@ class SyncCronClient:
metadata: Metadata to assign to the cron job runs.
config: The configuration for the assistant.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
checkpoint_during: Whether to checkpoint during the run (or only at the end/interruption).
interrupt_before: Nodes to interrupt immediately before they get executed.
interrupt_after: Nodes to Nodes to interrupt immediately after they get executed.
@@ -6171,7 +6171,7 @@ class SyncCronClient:
metadata: Metadata to assign to the cron job runs.
config: The configuration for the assistant.
context: Static context to add to the assistant.
!!! version-added "Supported with langgraph>=0.6.0"
!!! version-added "Added in version 0.6.0"
checkpoint_during: Whether to checkpoint during the run (or only at the end/interruption).
interrupt_before: Nodes to interrupt immediately before they get executed.
interrupt_after: Nodes to Nodes to interrupt immediately after they get executed.