mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-03 23:15:10 +02:00
Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ef15bc7d67 | ||
|
|
dddcb8ad8f | ||
|
|
5ea018e7d6 | ||
|
|
ac39fd2400 | ||
|
|
ce8c512642 | ||
|
|
81efce7fe7 | ||
|
|
b456b82ee4 | ||
|
|
11596dcd10 |
+3
-3
@@ -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.
|
||||
|
||||
```python
|
||||
my_function(1, "hello")
|
||||
\```
|
||||
.. code-block:: python
|
||||
|
||||
my_function(1, "hello")
|
||||
|
||||
Args:
|
||||
arg1: This is a description of arg1. We do not need to specify the type since
|
||||
|
||||
@@ -147,10 +147,10 @@ REDIRECT_MAP = {
|
||||
"cloud/how-tos/invoke_studio.md": "https://docs.langchain.com/langgraph-platform/use-studio#run-application",
|
||||
"cloud/how-tos/studio/manage_assistants.md": "https://docs.langchain.com/langgraph-platform/use-studio#manage-assistants",
|
||||
"cloud/how-tos/threads_studio.md": "https://docs.langchain.com/langgraph-platform/use-studio#manage-threads",
|
||||
"cloud/how-tos/iterate_graph_studio.md": "https://docs.langchain.com/langgraph-platform/observability-studio#iterate-on-prompts",
|
||||
"cloud/how-tos/studio/run_evals.md": "https://docs.langchain.com/langgraph-platform/observability-studio#run-experiments-over-a-dataset",
|
||||
"cloud/how-tos/clone_traces_studio.md": "https://docs.langchain.com/langgraph-platform/observability-studio#debug-langsmith-traces",
|
||||
"cloud/how-tos/datasets_studio.md": "https://docs.langchain.com/langgraph-platform/observability-studio#add-node-to-dataset",
|
||||
"cloud/how-tos/iterate_graph_studio.md": "https://docs.langchain.com/langgraph-platform/iterate-graph-studio",
|
||||
"cloud/how-tos/studio/run_evals.md": "https://docs.langchain.com/langgraph-platform/run-evals-studio",
|
||||
"cloud/how-tos/clone_traces_studio.md": "https://docs.langchain.com/langgraph-platform/clone-traces-studio",
|
||||
"cloud/how-tos/datasets_studio.md": "https://docs.langchain.com/langgraph-platform/datasets-studio",
|
||||
"concepts/sdk.md": "https://docs.langchain.com/langgraph-platform/sdk",
|
||||
"concepts/plans.md": "https://docs.langchain.com/langgraph-platform/plans",
|
||||
"concepts/application_structure.md": "https://docs.langchain.com/langgraph-platform/application-structure",
|
||||
|
||||
@@ -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 "Added in version 0.6.0: `context` replaces `config['configurable']`"
|
||||
!!! version-added "New in LangGraph v0.6: `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']`.
|
||||
|
||||
@@ -211,7 +211,7 @@ output = agent.invoke(
|
||||
print(output["messages"][-1].text())
|
||||
```
|
||||
|
||||
!!! version-added "Added in version 0.6.0"
|
||||
!!! version-added "New in LangGraph v0.6"
|
||||
|
||||
:::
|
||||
|
||||
@@ -351,13 +351,11 @@ 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`.
|
||||
@@ -373,7 +371,6 @@ 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
|
||||
@@ -384,5 +381,4 @@ 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/)
|
||||
|
||||
:::
|
||||
|
||||
@@ -244,7 +244,7 @@ output = agent.invoke(
|
||||
print(output["messages"][-1].text())
|
||||
```
|
||||
|
||||
!!! version-added "Added in version 0.6.0"
|
||||
!!! version-added "New in langgraph>=0.6"
|
||||
|
||||
:::
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@ from langgraph.checkpoint.base import (
|
||||
CheckpointMetadata,
|
||||
CheckpointTuple,
|
||||
get_checkpoint_id,
|
||||
get_serializable_checkpoint_metadata,
|
||||
get_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_serializable_checkpoint_metadata(config, metadata)),
|
||||
Jsonb(get_checkpoint_metadata(config, metadata)),
|
||||
),
|
||||
)
|
||||
return next_config
|
||||
|
||||
@@ -14,7 +14,7 @@ from langgraph.checkpoint.base import (
|
||||
CheckpointMetadata,
|
||||
CheckpointTuple,
|
||||
get_checkpoint_id,
|
||||
get_serializable_checkpoint_metadata,
|
||||
get_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_serializable_checkpoint_metadata(config, metadata)),
|
||||
Jsonb(get_checkpoint_metadata(config, metadata)),
|
||||
),
|
||||
)
|
||||
return next_config
|
||||
|
||||
@@ -1,9 +1,7 @@
|
||||
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
|
||||
@@ -18,18 +16,6 @@ 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_serializable_checkpoint_metadata,
|
||||
get_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_serializable_checkpoint_metadata(config, metadata)),
|
||||
Jsonb(get_checkpoint_metadata(config, metadata)),
|
||||
),
|
||||
)
|
||||
return next_config
|
||||
@@ -774,7 +774,7 @@ class AsyncShallowPostgresSaver(BasePostgresSaver):
|
||||
thread_id,
|
||||
checkpoint_ns,
|
||||
Jsonb(copy),
|
||||
Jsonb(get_serializable_checkpoint_metadata(config, metadata)),
|
||||
Jsonb(get_checkpoint_metadata(config, metadata)),
|
||||
),
|
||||
)
|
||||
return next_config
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "langgraph-checkpoint-postgres"
|
||||
version = "2.0.25"
|
||||
version = "2.0.23"
|
||||
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.1.2,<3.0.0",
|
||||
"langgraph-checkpoint>=2.0.21,<3.0.0",
|
||||
"orjson>=3.10.1",
|
||||
"psycopg>=3.2.0",
|
||||
"psycopg-pool>=3.2.0",
|
||||
|
||||
@@ -187,11 +187,13 @@ 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 = {}
|
||||
@@ -218,6 +220,7 @@ 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, {})
|
||||
@@ -243,6 +246,7 @@ 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,11 +169,13 @@ 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 = {}
|
||||
@@ -200,6 +202,7 @@ 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, {})
|
||||
@@ -225,6 +228,7 @@ 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
|
||||
|
||||
Generated
+2
-2
@@ -245,7 +245,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "2.1.2"
|
||||
version = "2.1.1"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -276,7 +276,7 @@ dev = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint-postgres"
|
||||
version = "2.0.25"
|
||||
version = "2.0.23"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "langgraph-checkpoint" },
|
||||
|
||||
Generated
+1
-1
@@ -257,7 +257,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "2.1.2"
|
||||
version = "2.1.1"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -404,16 +404,6 @@ 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.
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "2.1.2"
|
||||
version = "2.1.1"
|
||||
description = "Library with base interfaces for LangGraph checkpoint savers."
|
||||
authors = []
|
||||
requires-python = ">=3.9"
|
||||
|
||||
Generated
+1
-1
@@ -273,7 +273,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "2.1.2"
|
||||
version = "2.1.1"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -1 +1 @@
|
||||
__version__ = "0.4.3"
|
||||
__version__ = "0.4.2"
|
||||
|
||||
@@ -16,6 +16,7 @@ DEFAULT_PYTHON_VERSION = "3.11"
|
||||
|
||||
DEFAULT_IMAGE_DISTRO = "debian"
|
||||
|
||||
CONSTRAINTS_PATH = "/api/constraints.txt"
|
||||
|
||||
Distros = Literal["debian", "wolfi", "bullseye", "bookworm"]
|
||||
MiddlewareOrders = Literal["auth_first", "middleware_first"]
|
||||
@@ -793,6 +794,15 @@ def validate_config_file(config_path: pathlib.Path) -> Config:
|
||||
return validated
|
||||
|
||||
|
||||
class ReqGenSpec(NamedTuple):
|
||||
host_pkg_path: pathlib.Path
|
||||
container_pkg_path: str
|
||||
container_req_path: str
|
||||
package_type: Literal["pyproject", "setup"]
|
||||
has_uv_lock: bool
|
||||
stage_name: Optional[str]
|
||||
|
||||
|
||||
class LocalDeps(NamedTuple):
|
||||
"""A container for referencing and managing local Python dependencies.
|
||||
|
||||
@@ -840,6 +850,15 @@ class LocalDeps(NamedTuple):
|
||||
additional_contexts: A list of paths to directories that contain local
|
||||
dependencies in parent directories. These directories are added to the
|
||||
Docker build context to ensure that the Dockerfile can access them.
|
||||
|
||||
pkgs_missing_reqs: A list of packages that need requirements.txt generated.
|
||||
Each entry is a ReqGenSpec containing:
|
||||
- host_pkg_path: Absolute host path of the package directory
|
||||
- container_pkg_path: Target path in the container for metadata/requirements
|
||||
- container_req_path: Full path to requirements.txt inside the container
|
||||
- package_type: "pyproject" or "setup"
|
||||
- has_uv_lock: True if the package includes uv.lock
|
||||
- stage_name: BuildKit context name if outside the build context, else None
|
||||
"""
|
||||
|
||||
pip_reqs: list[tuple[pathlib.Path, str]]
|
||||
@@ -848,7 +867,9 @@ class LocalDeps(NamedTuple):
|
||||
# if . is in dependencies, use it as working_dir
|
||||
working_dir: Optional[str] = None
|
||||
# if there are local dependencies in parent directories, use additional_contexts
|
||||
additional_contexts: list[pathlib.Path] = None
|
||||
additional_contexts: Optional[list[pathlib.Path]] = None
|
||||
# Packages that need requirements.txt generation
|
||||
pkgs_missing_reqs: Optional[list[ReqGenSpec]] = None
|
||||
|
||||
|
||||
def _assemble_local_deps(config_path: pathlib.Path, config: Config) -> LocalDeps:
|
||||
@@ -884,6 +905,7 @@ def _assemble_local_deps(config_path: pathlib.Path, config: Config) -> LocalDeps
|
||||
faux_pkgs = {}
|
||||
working_dir: Optional[str] = None
|
||||
additional_contexts: list[pathlib.Path] = []
|
||||
pkgs_missing_reqs: list[ReqGenSpec] = []
|
||||
|
||||
for local_dep in config["dependencies"]:
|
||||
if not local_dep.startswith("."):
|
||||
@@ -923,6 +945,35 @@ def _assemble_local_deps(config_path: pathlib.Path, config: Config) -> LocalDeps
|
||||
# set working_dir
|
||||
if local_dep == ".":
|
||||
working_dir = f"/deps/{container_name}"
|
||||
requirement_path = f"/deps/{container_name}/requirements.txt"
|
||||
|
||||
# Track packages needing requirements.txt generation (real packages only)
|
||||
if "requirements.txt" not in files:
|
||||
has_pyproject = "pyproject.toml" in files
|
||||
has_uv_lock = "uv.lock" in files
|
||||
|
||||
if has_pyproject:
|
||||
pkg_type = "pyproject"
|
||||
has_lock = has_uv_lock
|
||||
else: # setup.py only
|
||||
pkg_type = "setup"
|
||||
has_lock = False
|
||||
|
||||
container_pkg_path = f"/deps/{container_name}"
|
||||
stage_name = None
|
||||
if config_path.parent not in resolved.parents:
|
||||
stage_name = container_name
|
||||
|
||||
pkgs_missing_reqs.append(
|
||||
ReqGenSpec(
|
||||
host_pkg_path=resolved,
|
||||
container_pkg_path=container_pkg_path,
|
||||
container_req_path=requirement_path,
|
||||
package_type=pkg_type,
|
||||
has_uv_lock=has_lock,
|
||||
stage_name=stage_name,
|
||||
)
|
||||
)
|
||||
else:
|
||||
# We could not find a pyproject.toml or setup.py, so treat as a faux package
|
||||
if any(file == "__init__.py" for file in files):
|
||||
@@ -954,19 +1005,27 @@ def _assemble_local_deps(config_path: pathlib.Path, config: Config) -> LocalDeps
|
||||
faux_pkgs[resolved] = (local_dep, container_path)
|
||||
if local_dep == ".":
|
||||
working_dir = container_path
|
||||
requirement_path = f"{container_path}/requirements.txt"
|
||||
|
||||
# If the faux package has a requirements.txt, we'll add
|
||||
# the path to the list of requirements to install.
|
||||
if "requirements.txt" in files:
|
||||
rfile = resolved / "requirements.txt"
|
||||
pip_reqs.append(
|
||||
(
|
||||
rfile,
|
||||
f"{container_path}/requirements.txt",
|
||||
)
|
||||
# If the package has a requirements.txt, we'll add
|
||||
# the path to the list of requirements to install.
|
||||
if "requirements.txt" in files:
|
||||
rfile = resolved / "requirements.txt"
|
||||
pip_reqs.append(
|
||||
(
|
||||
rfile,
|
||||
requirement_path,
|
||||
)
|
||||
)
|
||||
|
||||
return LocalDeps(pip_reqs, real_pkgs, faux_pkgs, working_dir, additional_contexts)
|
||||
return LocalDeps(
|
||||
pip_reqs,
|
||||
real_pkgs,
|
||||
faux_pkgs,
|
||||
working_dir,
|
||||
additional_contexts,
|
||||
pkgs_missing_reqs,
|
||||
)
|
||||
|
||||
|
||||
def _update_graph_paths(
|
||||
@@ -1261,6 +1320,110 @@ def get_build_tools_to_uninstall(config: Config) -> tuple[str]:
|
||||
)
|
||||
|
||||
|
||||
def _metadata_files(spec: ReqGenSpec) -> list[str]:
|
||||
files = ["pyproject.toml"] if spec.package_type == "pyproject" else ["setup.py"]
|
||||
if spec.package_type == "pyproject" and spec.has_uv_lock:
|
||||
files.append("uv.lock")
|
||||
if spec.package_type == "setup" and (spec.host_pkg_path / "setup.cfg").exists():
|
||||
files.append("setup.cfg")
|
||||
return files
|
||||
|
||||
|
||||
def _get_reqs_gen_cmd(spec: ReqGenSpec) -> str:
|
||||
if spec.package_type == "pyproject" and spec.has_uv_lock:
|
||||
return "uv export --no-hashes --no-dev --no-emit-local -o 'requirements.txt'"
|
||||
if spec.package_type == "pyproject":
|
||||
return f"uv pip compile pyproject.toml -o 'requirements.txt' --constraint {CONSTRAINTS_PATH}"
|
||||
return (
|
||||
f"uv pip compile setup.py -o 'requirements.txt' --constraint {CONSTRAINTS_PATH}"
|
||||
)
|
||||
|
||||
|
||||
def _generate_requirements_from_metadata(
|
||||
config_path: pathlib.Path,
|
||||
local_deps: LocalDeps,
|
||||
pip_installer: Literal["uv", "pip"],
|
||||
) -> str:
|
||||
"""Generate requirements.txt from uv.lock, pyproject.toml, or setup.py.
|
||||
|
||||
This function creates Docker layers that:
|
||||
1. Copy packaging metadata files (pyproject.toml, setup.py, uv.lock, etc.)
|
||||
2. Generate requirements.txt using appropriate uv commands:
|
||||
- `uv export` if uv.lock exists (exact locked versions)
|
||||
- `uv pip compile` for pyproject.toml or setup.py (fresh resolution)
|
||||
|
||||
The generated requirements.txt files are then handled by the existing
|
||||
pip_reqs installation logic, maintaining a single code path for all packages.
|
||||
|
||||
Supports:
|
||||
- pyproject.toml with uv.lock → uv export (preserves lock file versions)
|
||||
- pyproject.toml without uv.lock → uv pip compile (resolve dependencies)
|
||||
- setup.py (with optional setup.cfg) → uv pip compile (resolve dependencies)
|
||||
|
||||
Args:
|
||||
config_path: Path to the langgraph.json config file
|
||||
local_deps: LocalDeps object containing package information
|
||||
pip_installer: Either "uv" or "pip"
|
||||
|
||||
Returns:
|
||||
Docker instruction string for requirements.txt generation,
|
||||
or empty string if not applicable (pip installer or no packages to generate)
|
||||
"""
|
||||
if pip_installer != "uv" or not local_deps.pkgs_missing_reqs:
|
||||
# if installer is pip, we need pip-tools to generate requirements.txt
|
||||
# this doesn't come automatically with pip, so we need to install it
|
||||
# in base images, but ci uses uv, which is where we need the caching
|
||||
# so limit to uv.
|
||||
return ""
|
||||
|
||||
docker_lines = ["# -- Generate requirements.txt for packages without one --"]
|
||||
|
||||
# Layer 1: Copy packaging metadata files needed for requirements generation
|
||||
docker_lines.append("# Copy packaging metadata files")
|
||||
for spec in sorted(
|
||||
local_deps.pkgs_missing_reqs, key=lambda s: s.container_pkg_path
|
||||
):
|
||||
for file_name in _metadata_files(spec):
|
||||
if (
|
||||
local_deps.additional_contexts
|
||||
and spec.host_pkg_path in local_deps.additional_contexts
|
||||
):
|
||||
if not spec.stage_name:
|
||||
raise RuntimeError(
|
||||
f"Package {spec.host_pkg_path} in additional_contexts but has no stage_name"
|
||||
)
|
||||
docker_lines.append(
|
||||
f"COPY --from={spec.stage_name} {file_name} {spec.container_pkg_path}/{file_name}"
|
||||
)
|
||||
else:
|
||||
file_relpath = (spec.host_pkg_path / file_name).relative_to(
|
||||
config_path.parent
|
||||
)
|
||||
docker_lines.append(
|
||||
f"ADD {file_relpath} {spec.container_pkg_path}/{file_name}"
|
||||
)
|
||||
|
||||
# Layer 2: Generate requirements.txt files using appropriate uv commands
|
||||
docker_lines.append("")
|
||||
docker_lines.append("# Generate requirements.txt from packaging metadata")
|
||||
for spec in sorted(
|
||||
local_deps.pkgs_missing_reqs, key=lambda s: s.container_pkg_path
|
||||
):
|
||||
pkg_name = spec.host_pkg_path.name
|
||||
if spec.package_type == "pyproject" and spec.has_uv_lock:
|
||||
docker_lines.append(f"# Generate from uv.lock for {pkg_name}")
|
||||
elif spec.package_type == "pyproject":
|
||||
docker_lines.append(f"# Compile from pyproject.toml for {pkg_name}")
|
||||
else:
|
||||
docker_lines.append(f"# Compile from setup.py for {pkg_name}")
|
||||
docker_lines.append(
|
||||
f"RUN cd '{spec.container_pkg_path}' && {_get_reqs_gen_cmd(spec)}"
|
||||
)
|
||||
|
||||
docker_lines.append("# -- End of requirements.txt generation --")
|
||||
return os.linesep.join(docker_lines)
|
||||
|
||||
|
||||
def python_config_to_docker(
|
||||
config_path: pathlib.Path,
|
||||
config: Config,
|
||||
@@ -1283,8 +1446,12 @@ def python_config_to_docker(
|
||||
raise ValueError(f"Invalid pip_installer: {pip_installer}")
|
||||
|
||||
# configure pip
|
||||
local_reqs_pip_install = f"PYTHONDONTWRITEBYTECODE=1 {install_cmd} --no-cache-dir -c /api/constraints.txt"
|
||||
global_reqs_pip_install = f"PYTHONDONTWRITEBYTECODE=1 {install_cmd} --no-cache-dir -c /api/constraints.txt"
|
||||
local_reqs_pip_install = (
|
||||
f"PYTHONDONTWRITEBYTECODE=1 {install_cmd} --no-cache-dir -c {CONSTRAINTS_PATH}"
|
||||
)
|
||||
global_reqs_pip_install = (
|
||||
f"PYTHONDONTWRITEBYTECODE=1 {install_cmd} --no-cache-dir -c {CONSTRAINTS_PATH}"
|
||||
)
|
||||
if config.get("pip_config_file"):
|
||||
local_reqs_pip_install = (
|
||||
f"PIP_CONFIG_FILE=/pipconfig.txt {local_reqs_pip_install}"
|
||||
@@ -1311,22 +1478,59 @@ def python_config_to_docker(
|
||||
pip_pkgs_str = (
|
||||
f"RUN {local_reqs_pip_install} {' '.join(pypi_deps)}" if pypi_deps else ""
|
||||
)
|
||||
if local_deps.pip_reqs:
|
||||
pip_reqs_str = os.linesep.join(
|
||||
(
|
||||
f"COPY --from=outer-{reqpath.name} requirements.txt {destpath}"
|
||||
if reqpath.parent in local_deps.additional_contexts
|
||||
else f"ADD {reqpath.relative_to(config_path.parent)} {destpath}"
|
||||
)
|
||||
for reqpath, destpath in local_deps.pip_reqs
|
||||
)
|
||||
pip_reqs_str += f"{os.linesep}RUN {local_reqs_pip_install} {' '.join('-r ' + r for _, r in local_deps.pip_reqs)}"
|
||||
pip_reqs_str = f"""# -- Installing local requirements --
|
||||
{pip_reqs_str}
|
||||
# -- End of local requirements install --"""
|
||||
|
||||
else:
|
||||
pip_reqs_str = ""
|
||||
# Generate requirements.txt layer for packages that need it
|
||||
# This happens BEFORE copying existing requirements.txt files
|
||||
generated_reqs_str = _generate_requirements_from_metadata(
|
||||
config_path, local_deps, pip_installer
|
||||
)
|
||||
# Combine existing requirements.txt with generated ones in a single deterministic layer
|
||||
all_req_paths: list[str] = []
|
||||
copy_lines: list[str] = []
|
||||
|
||||
# Map additional_contexts path -> stage name
|
||||
additional_ctx_stage: dict[pathlib.Path, str] = {}
|
||||
for p in local_deps.additional_contexts or []:
|
||||
if p in local_deps.real_pkgs:
|
||||
additional_ctx_stage[p] = local_deps.real_pkgs[p][1]
|
||||
elif p in local_deps.faux_pkgs:
|
||||
additional_ctx_stage[p] = f"outer-{p.name}"
|
||||
else:
|
||||
raise RuntimeError(
|
||||
f"Package {p} in additional_contexts but not in real_pkgs or faux_pkgs"
|
||||
)
|
||||
|
||||
# Existing reqs
|
||||
for reqpath, destpath in local_deps.pip_reqs or []:
|
||||
if local_deps.additional_contexts and reqpath.parent in additional_ctx_stage:
|
||||
copy_lines.append(
|
||||
f"COPY --from={additional_ctx_stage[reqpath.parent]} requirements.txt {destpath}"
|
||||
)
|
||||
else:
|
||||
copy_lines.append(
|
||||
f"ADD {reqpath.relative_to(config_path.parent)} {destpath}"
|
||||
)
|
||||
all_req_paths.append(destpath)
|
||||
|
||||
# Generated reqs
|
||||
if (
|
||||
pip_installer == "uv"
|
||||
): # we are only generate a requirements.txt if installer is uv (for now)
|
||||
for spec in local_deps.pkgs_missing_reqs or []:
|
||||
all_req_paths.append(spec.container_req_path)
|
||||
|
||||
pip_reqs_str = ""
|
||||
if all_req_paths:
|
||||
# Stabilize order
|
||||
all_req_paths = sorted(set(all_req_paths))
|
||||
# Install each requirements.txt sequentially. This mimics the previous sequential solver behavior
|
||||
# which allows to adjust/downgrade packages between installs, so no conflict is raised.
|
||||
pip_reqs_str = f"""# -- Installing from requirements.txt files --
|
||||
{os.linesep.join(copy_lines)}
|
||||
{os.linesep.join(f"RUN {local_reqs_pip_install} -r '{p}'" for p in all_req_paths)}
|
||||
# -- End of requirements.txt install --"""
|
||||
|
||||
# generate lock file if real package and lock file missing
|
||||
|
||||
# https://setuptools.pypa.io/en/latest/userguide/datafiles.html#package-data
|
||||
# https://til.simonwillison.net/python/pyproject
|
||||
@@ -1380,6 +1584,7 @@ ADD {relpath} /deps/{name}
|
||||
install_node_str,
|
||||
pip_config_file_str,
|
||||
pip_pkgs_str,
|
||||
generated_reqs_str,
|
||||
pip_reqs_str,
|
||||
local_pkgs_str,
|
||||
faux_pkgs_str,
|
||||
@@ -1417,7 +1622,7 @@ ADD {relpath} /deps/{name}
|
||||
[
|
||||
"# -- Installing JS dependencies --",
|
||||
f"ENV NODE_VERSION={config.get('node_version') or DEFAULT_NODE_VERSION}",
|
||||
f"RUN cd {local_deps.working_dir} && {_get_node_pm_install_cmd(config_path, config)} && tsx /api/langgraph_api/js/build.mts",
|
||||
f"RUN cd '{local_deps.working_dir}' && {_get_node_pm_install_cmd(config_path, config)} && tsx /api/langgraph_api/js/build.mts",
|
||||
"# -- End of JS dependencies install --",
|
||||
]
|
||||
)
|
||||
|
||||
@@ -142,6 +142,21 @@ services:
|
||||
- cli_1: {str(pathlib.Path(__file__).parent.parent.parent.parent.absolute())}
|
||||
dockerfile_inline: |
|
||||
FROM langchain/langgraph-api:3.11
|
||||
# -- Generate requirements.txt for packages without one --
|
||||
# Copy packaging metadata files
|
||||
ADD pyproject.toml /deps/cli/pyproject.toml
|
||||
COPY --from=cli_1 pyproject.toml /deps/cli_1/pyproject.toml
|
||||
COPY --from=cli_1 uv.lock /deps/cli_1/uv.lock
|
||||
# Generate requirements.txt from packaging metadata
|
||||
# Compile from pyproject.toml for cli
|
||||
RUN cd '/deps/cli' && uv pip compile pyproject.toml -o 'requirements.txt' --constraint /api/constraints.txt
|
||||
# Generate from uv.lock for cli
|
||||
RUN cd '/deps/cli_1' && uv export --no-hashes --no-dev -o 'requirements.txt'
|
||||
# -- End of requirements.txt generation --
|
||||
# -- Installing from requirements.txt files --
|
||||
RUN PYTHONDONTWRITEBYTECODE=1 uv pip install --system --no-cache-dir -c /api/constraints.txt -r '/deps/cli/requirements.txt'
|
||||
RUN PYTHONDONTWRITEBYTECODE=1 uv pip install --system --no-cache-dir -c /api/constraints.txt -r '/deps/cli_1/requirements.txt'
|
||||
# -- End of requirements.txt install --
|
||||
# -- Adding local package . --
|
||||
ADD . /deps/cli
|
||||
# -- End of local package . --
|
||||
|
||||
@@ -420,10 +420,18 @@ def test_config_to_docker_simple():
|
||||
)
|
||||
expected_docker_stdin = f"""\
|
||||
FROM langchain/langgraph-api:3.11
|
||||
# -- Installing local requirements --
|
||||
COPY --from=outer-requirements.txt requirements.txt /deps/outer-graphs_reqs_a/graphs_reqs_a/requirements.txt
|
||||
RUN PYTHONDONTWRITEBYTECODE=1 uv pip install --system --no-cache-dir -c /api/constraints.txt -r /deps/outer-graphs_reqs_a/graphs_reqs_a/requirements.txt
|
||||
# -- End of local requirements install --
|
||||
# -- Generate requirements.txt for packages without one --
|
||||
# Copy packaging metadata files
|
||||
COPY --from=examples pyproject.toml /deps/examples/pyproject.toml
|
||||
# Generate requirements.txt from packaging metadata
|
||||
# Compile from pyproject.toml for examples
|
||||
RUN cd '/deps/examples' && uv pip compile pyproject.toml -o 'requirements.txt' --constraint /api/constraints.txt
|
||||
# -- End of requirements.txt generation --
|
||||
# -- Installing from requirements.txt files --
|
||||
COPY --from=outer-graphs_reqs_a requirements.txt /deps/outer-graphs_reqs_a/graphs_reqs_a/requirements.txt
|
||||
RUN PYTHONDONTWRITEBYTECODE=1 uv pip install --system --no-cache-dir -c /api/constraints.txt -r '/deps/examples/requirements.txt'
|
||||
RUN PYTHONDONTWRITEBYTECODE=1 uv pip install --system --no-cache-dir -c /api/constraints.txt -r '/deps/outer-graphs_reqs_a/graphs_reqs_a/requirements.txt'
|
||||
# -- End of requirements.txt install --
|
||||
# -- Adding local package ../../examples --
|
||||
COPY --from=examples . /deps/examples
|
||||
# -- End of local package ../../examples --
|
||||
@@ -653,6 +661,16 @@ dependencies = ["langchain"]"""
|
||||
os.remove(pyproject_path)
|
||||
expected_docker_stdin = (
|
||||
"""FROM langchain/langgraph-api:3.11
|
||||
# -- Generate requirements.txt for packages without one --
|
||||
# Copy packaging metadata files
|
||||
ADD pyproject.toml /deps/unit_tests/pyproject.toml
|
||||
# Generate requirements.txt from packaging metadata
|
||||
# Compile from pyproject.toml for unit_tests
|
||||
RUN cd '/deps/unit_tests' && uv pip compile pyproject.toml -o 'requirements.txt' --constraint /api/constraints.txt
|
||||
# -- End of requirements.txt generation --
|
||||
# -- Installing from requirements.txt files --
|
||||
RUN PYTHONDONTWRITEBYTECODE=1 uv pip install --system --no-cache-dir -c /api/constraints.txt -r '/deps/unit_tests/requirements.txt'
|
||||
# -- End of requirements.txt install --
|
||||
# -- Adding local package . --
|
||||
ADD . /deps/unit_tests
|
||||
# -- End of local package . --
|
||||
@@ -818,7 +836,7 @@ ENV LANGGRAPH_UI_CONFIG='{{"shared": ["nuqs"]}}'
|
||||
ENV LANGSERVE_GRAPHS='{{"agent": "/deps/outer-unit_tests/unit_tests/agent.py:graph"}}'
|
||||
# -- Installing JS dependencies --
|
||||
ENV NODE_VERSION=20
|
||||
RUN cd /deps/outer-unit_tests/unit_tests && npm i && tsx /api/langgraph_api/js/build.mts
|
||||
RUN cd '/deps/outer-unit_tests/unit_tests' && npm i && tsx /api/langgraph_api/js/build.mts
|
||||
# -- End of JS dependencies install --
|
||||
{FORMATTED_CLEANUP_LINES}
|
||||
WORKDIR /deps/outer-unit_tests/unit_tests"""
|
||||
@@ -862,7 +880,7 @@ RUN for dep in /deps/*; do echo "Installing $dep"; if [
|
||||
ENV LANGSERVE_GRAPHS='{{"python": "/deps/outer-unit_tests/unit_tests/multiplatform/python.py:graph", "js": "/deps/outer-unit_tests/unit_tests/multiplatform/js.mts:graph"}}'
|
||||
# -- Installing JS dependencies --
|
||||
ENV NODE_VERSION=22
|
||||
RUN cd /deps/outer-unit_tests/unit_tests && npm i && tsx /api/langgraph_api/js/build.mts
|
||||
RUN cd '/deps/outer-unit_tests/unit_tests' && npm i && tsx /api/langgraph_api/js/build.mts
|
||||
# -- End of JS dependencies install --
|
||||
{FORMATTED_CLEANUP_LINES}
|
||||
WORKDIR /deps/outer-unit_tests/unit_tests"""
|
||||
|
||||
@@ -89,12 +89,13 @@ def push_ui_message(
|
||||
The created UI message.
|
||||
|
||||
Example:
|
||||
```python
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
push_ui_message(
|
||||
name="component-name",
|
||||
props={"content": "Hello world"},
|
||||
)
|
||||
```
|
||||
|
||||
"""
|
||||
from langgraph._internal._constants import CONFIG_KEY_SEND
|
||||
@@ -145,9 +146,10 @@ def delete_ui_message(id: str, *, state_key: str = "ui") -> RemoveUIMessage:
|
||||
The remove UI message.
|
||||
|
||||
Example:
|
||||
```python
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
delete_ui_message("message-123")
|
||||
```
|
||||
|
||||
"""
|
||||
from langgraph._internal._constants import CONFIG_KEY_SEND
|
||||
@@ -181,12 +183,13 @@ def ui_message_reducer(
|
||||
Combined list of UI messages with removals applied.
|
||||
|
||||
Example:
|
||||
```python
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
messages = ui_message_reducer(
|
||||
[{"type": "ui", "id": "1", "name": "Chat", "props": {}}],
|
||||
{"type": "remove-ui", "id": "1"},
|
||||
)
|
||||
```
|
||||
|
||||
"""
|
||||
if not isinstance(left, list):
|
||||
|
||||
@@ -114,7 +114,6 @@ from langgraph.types import (
|
||||
CachePolicy,
|
||||
Command,
|
||||
Durability,
|
||||
Interrupt,
|
||||
PregelExecutableTask,
|
||||
RetryPolicy,
|
||||
StreamMode,
|
||||
@@ -243,13 +242,14 @@ 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]
|
||||
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._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)
|
||||
@@ -318,19 +318,14 @@ class PregelLoop:
|
||||
writes_to_save: WritesT = [
|
||||
w[1:] for w in self.checkpoint_pending_writes if w[0] == task_id
|
||||
] + list(writes)
|
||||
self.checkpoint_pending_writes.extend((task_id, c, v) for c, v in writes)
|
||||
else:
|
||||
writes_to_save = [
|
||||
# aggregate existing interrupts for this task
|
||||
(ch, self._merge_interrupts(task_id, v) if ch == INTERRUPT else v)
|
||||
for ch, v in writes
|
||||
]
|
||||
|
||||
# replace all writes for this task_id in one shot
|
||||
# remove existing writes for this task
|
||||
self.checkpoint_pending_writes = [
|
||||
w for w in self.checkpoint_pending_writes if w[0] != task_id
|
||||
] + [(task_id, c, v) for c, v in writes_to_save]
|
||||
|
||||
]
|
||||
writes_to_save = writes
|
||||
# save writes
|
||||
self.checkpoint_pending_writes.extend((task_id, c, v) for c, v in writes)
|
||||
if self.durability != "exit" and self.checkpointer_put_writes is not None:
|
||||
config = patch_configurable(
|
||||
self.checkpoint_config,
|
||||
@@ -478,20 +473,6 @@ class PregelLoop:
|
||||
cache_policy=self.cache_policy,
|
||||
)
|
||||
|
||||
resume_map = self.config.get(CONF, {}).get(CONFIG_KEY_RESUME_MAP, {})
|
||||
if 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(
|
||||
@@ -537,45 +518,9 @@ class PregelLoop:
|
||||
if task.writes:
|
||||
self.output_writes(task.id, task.writes, cached=True)
|
||||
|
||||
if self.skipped_task_ids:
|
||||
# remove tasks with writes that may have been matched from previous loop
|
||||
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 preemptively
|
||||
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
|
||||
@@ -627,53 +572,34 @@ 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_id, interrupt_count)
|
||||
pending_interrupts: dict[str, tuple[str, int]] = {}
|
||||
# mapping of task ids to resume count
|
||||
pending_resumes: dict[str, int] = {}
|
||||
# mapping of task ids to interrupt ids
|
||||
pending_interrupts: dict[str, str] = {}
|
||||
|
||||
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)
|
||||
# set of resume task ids
|
||||
pending_resumes: set[str] = set()
|
||||
|
||||
# keep only interrupt ids where resume_count < interrupt_count
|
||||
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
|
||||
hanging_interrupts: set[str] = {
|
||||
interrupt_id
|
||||
for task_id, (interrupt_id, interrupt_count) in pending_interrupts.items()
|
||||
if pending_resumes.get(task_id, 0) < interrupt_count
|
||||
for interrupt_id in pending_interrupts.values()
|
||||
if interrupt_id not in resumed_interrupt_ids
|
||||
}
|
||||
|
||||
return hanging_interrupts
|
||||
|
||||
def _merge_interrupts(
|
||||
self, task_id: str, value: Sequence[Interrupt]
|
||||
) -> Sequence[Interrupt]:
|
||||
"""Normalize interrupt value to list and merge with existing interrupts.
|
||||
|
||||
If the interrupt ID matches existing, append; otherwise replace.
|
||||
|
||||
Returns list of Interrupt objects for this task.
|
||||
"""
|
||||
new = value if isinstance(value, list) else list(value)
|
||||
existing = next(
|
||||
(
|
||||
v
|
||||
for tid, ch, v in self.checkpoint_pending_writes
|
||||
if tid == task_id and ch == INTERRUPT
|
||||
),
|
||||
None,
|
||||
)
|
||||
if existing is None:
|
||||
return new
|
||||
old = existing if isinstance(existing, list) else list(existing)
|
||||
return old + new if old and new and old[0].id == new[0].id else new
|
||||
|
||||
def _first(
|
||||
self, *, input_keys: str | Sequence[str], updated_channels: set[str] | None
|
||||
) -> set[str] | None:
|
||||
@@ -1102,7 +1028,6 @@ 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
|
||||
|
||||
@@ -40,7 +40,7 @@ class TaskResultPayload(TypedDict):
|
||||
name: str
|
||||
error: str | None
|
||||
interrupts: list[dict]
|
||||
result: dict[str, Any]
|
||||
result: list[tuple[str, Any]]
|
||||
|
||||
|
||||
class CheckpointTask(TypedDict):
|
||||
@@ -77,38 +77,6 @@ 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],
|
||||
@@ -122,9 +90,7 @@ 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": map_task_result_writes(
|
||||
[w for w in writes if w[0] in stream_channels_list or w[0] == RETURN]
|
||||
),
|
||||
"result": [w for w in writes if w[0] in stream_channels_list or w[0] == RETURN],
|
||||
"interrupts": [
|
||||
asdict(v)
|
||||
for w in writes
|
||||
@@ -230,56 +196,54 @@ 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,
|
||||
task_error,
|
||||
task_interrupts,
|
||||
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])
|
||||
),
|
||||
states.get(task.id) if states else None,
|
||||
task_result if has_writes 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,
|
||||
)
|
||||
)
|
||||
return tuple(out)
|
||||
|
||||
@@ -1695,10 +1695,12 @@ class Pregel(
|
||||
|
||||
return patch_checkpoint_map(next_config, saved.metadata)
|
||||
|
||||
# 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:
|
||||
# apply pending writes, if not on specific checkpoint
|
||||
if (
|
||||
CONFIG_KEY_CHECKPOINT_ID not in config[CONF]
|
||||
and saved is not None
|
||||
and saved.pending_writes
|
||||
):
|
||||
# tasks for this checkpoint
|
||||
next_tasks = prepare_next_tasks(
|
||||
checkpoint,
|
||||
@@ -1714,10 +1716,6 @@ 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
|
||||
@@ -1799,14 +1797,8 @@ class Pregel(
|
||||
raise InvalidUpdateError(f"Node {as_node} has no writers")
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
task = PregelTaskWrites((), as_node, writes, [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))
|
||||
task_id = provided_task_id or str(
|
||||
uuid5(UUID(checkpoint["id"]), INTERRUPT)
|
||||
)
|
||||
run_tasks.append(task)
|
||||
run_task_ids.append(task_id)
|
||||
@@ -2159,11 +2151,12 @@ class Pregel(
|
||||
return patch_checkpoint_map(
|
||||
next_config, saved.metadata if saved else None
|
||||
)
|
||||
|
||||
# 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:
|
||||
# apply pending writes, if not on specific checkpoint
|
||||
if (
|
||||
CONFIG_KEY_CHECKPOINT_ID not in config[CONF]
|
||||
and saved is not None
|
||||
and saved.pending_writes
|
||||
):
|
||||
# tasks for this checkpoint
|
||||
next_tasks = prepare_next_tasks(
|
||||
checkpoint,
|
||||
@@ -2179,10 +2172,6 @@ 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
|
||||
@@ -2259,14 +2248,8 @@ class Pregel(
|
||||
raise InvalidUpdateError(f"Node {as_node} has no writers")
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
task = PregelTaskWrites((), as_node, writes, [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))
|
||||
task_id = provided_task_id or str(
|
||||
uuid5(UUID(checkpoint["id"]), INTERRUPT)
|
||||
)
|
||||
run_tasks.append(task)
|
||||
run_task_ids.append(task_id)
|
||||
@@ -2462,7 +2445,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:
|
||||
|
||||
@@ -2672,11 +2655,7 @@ 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 and t.id not in loop.skipped_task_ids
|
||||
],
|
||||
[t for t in loop.tasks.values() if not t.writes],
|
||||
timeout=self.step_timeout,
|
||||
get_waiter=get_waiter,
|
||||
schedule_task=loop.accept_push,
|
||||
@@ -2732,7 +2711,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:
|
||||
|
||||
@@ -2995,11 +2974,7 @@ 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 and t.id not in loop.skipped_task_ids
|
||||
],
|
||||
[t for t in loop.tasks.values() if not t.writes],
|
||||
timeout=self.step_timeout,
|
||||
get_waiter=get_waiter,
|
||||
schedule_task=loop.aaccept_push,
|
||||
@@ -3068,7 +3043,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.
|
||||
@@ -3153,7 +3128,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.
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "langgraph"
|
||||
version = "0.6.9"
|
||||
version = "0.6.8"
|
||||
description = "Building stateful, multi-actor applications with LLMs"
|
||||
authors = []
|
||||
requires-python = ">=3.9"
|
||||
|
||||
@@ -1,21 +1,12 @@
|
||||
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 Command, Durability, Send, interrupt
|
||||
from langgraph.types import Durability
|
||||
|
||||
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
|
||||
@@ -99,486 +90,3 @@ 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 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 == 3
|
||||
|
||||
# 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 == 4
|
||||
|
||||
# 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 == 5
|
||||
|
||||
|
||||
@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 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 == 3
|
||||
|
||||
# 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 == 4
|
||||
|
||||
# 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 == 5
|
||||
|
||||
@@ -4023,9 +4023,7 @@ 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": [],
|
||||
},
|
||||
@@ -4073,9 +4071,7 @@ 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": [],
|
||||
},
|
||||
@@ -4094,9 +4090,7 @@ 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": [],
|
||||
},
|
||||
@@ -4136,9 +4130,7 @@ 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": [],
|
||||
},
|
||||
|
||||
@@ -2567,9 +2567,7 @@ 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": [],
|
||||
},
|
||||
@@ -2617,9 +2615,7 @@ 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": [],
|
||||
},
|
||||
@@ -2638,9 +2634,7 @@ 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": [],
|
||||
},
|
||||
@@ -2680,9 +2674,7 @@ 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": [],
|
||||
},
|
||||
|
||||
@@ -2119,9 +2119,7 @@ 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": [],
|
||||
},
|
||||
},
|
||||
@@ -2155,9 +2153,7 @@ 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": [],
|
||||
},
|
||||
},
|
||||
@@ -2169,9 +2165,7 @@ 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": [],
|
||||
},
|
||||
},
|
||||
@@ -2197,9 +2191,7 @@ 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": [],
|
||||
},
|
||||
},
|
||||
@@ -2227,9 +2219,7 @@ 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": [],
|
||||
},
|
||||
},
|
||||
@@ -3455,6 +3445,73 @@ 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:
|
||||
@@ -5549,9 +5606,12 @@ 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(),
|
||||
@@ -5568,7 +5628,7 @@ def test_falsy_return_from_task(sync_checkpointer: BaseCheckpointSaver):
|
||||
},
|
||||
],
|
||||
"name": "graph",
|
||||
"result": {},
|
||||
"result": [],
|
||||
},
|
||||
"step": 0,
|
||||
"timestamp": AnyStr(),
|
||||
@@ -5654,9 +5714,12 @@ def test_falsy_return_from_task(sync_checkpointer: BaseCheckpointSaver):
|
||||
"id": AnyStr(),
|
||||
"interrupts": [],
|
||||
"name": "graph",
|
||||
"result": {
|
||||
"__end__": None,
|
||||
},
|
||||
"result": [
|
||||
(
|
||||
"__end__",
|
||||
None,
|
||||
),
|
||||
],
|
||||
},
|
||||
"step": 0,
|
||||
"timestamp": AnyStr(),
|
||||
@@ -8453,163 +8516,3 @@ 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": ""},
|
||||
},
|
||||
]
|
||||
|
||||
@@ -9211,58 +9211,3 @@ 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"}
|
||||
|
||||
Generated
+4
-4
@@ -1,5 +1,5 @@
|
||||
version = 1
|
||||
revision = 3
|
||||
revision = 2
|
||||
requires-python = ">=3.9"
|
||||
resolution-markers = [
|
||||
"python_full_version >= '3.14'",
|
||||
@@ -1428,7 +1428,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph"
|
||||
version = "0.6.9"
|
||||
version = "0.6.8"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -1542,7 +1542,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "2.1.2"
|
||||
version = "2.1.1"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -1573,7 +1573,7 @@ dev = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint-postgres"
|
||||
version = "2.0.25"
|
||||
version = "2.0.23"
|
||||
source = { editable = "../checkpoint-postgres" }
|
||||
dependencies = [
|
||||
{ name = "langgraph-checkpoint" },
|
||||
|
||||
Generated
+4
-4
@@ -1,5 +1,5 @@
|
||||
version = 1
|
||||
revision = 3
|
||||
revision = 2
|
||||
requires-python = ">=3.9"
|
||||
|
||||
[[package]]
|
||||
@@ -257,7 +257,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph"
|
||||
version = "0.6.9"
|
||||
version = "0.6.8"
|
||||
source = { editable = "../langgraph" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -309,7 +309,7 @@ dev = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "2.1.2"
|
||||
version = "2.1.1"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -340,7 +340,7 @@ dev = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint-postgres"
|
||||
version = "2.0.25"
|
||||
version = "2.0.23"
|
||||
source = { editable = "../checkpoint-postgres" }
|
||||
dependencies = [
|
||||
{ name = "langgraph-checkpoint" },
|
||||
|
||||
@@ -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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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 "Added in version 0.6.0"
|
||||
!!! version-added "Supported with langgraph>=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.
|
||||
|
||||
Reference in New Issue
Block a user