mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-01 22:15:11 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
22083959f4 | ||
|
|
858e55f232 | ||
|
|
e59ecbc23d | ||
|
|
719a4d71bc | ||
|
|
ddaf708cd0 | ||
|
|
9d5b0f1991 | ||
|
|
9d16b52955 | ||
|
|
2de0c47c1f | ||
|
|
377083220e | ||
|
|
9c5914861b | ||
|
|
24cf33f348 | ||
|
|
07b33185ea |
@@ -32,6 +32,8 @@ def test(config: pathlib.Path, port: int, tag: str, verbose: bool):
|
|||||||
docker_compose=None,
|
docker_compose=None,
|
||||||
port=port,
|
port=port,
|
||||||
watch=False,
|
watch=False,
|
||||||
|
debugger_port=None,
|
||||||
|
debugger_base_url=f"http://127.0.0.1:{port}",
|
||||||
postgres_uri=None,
|
postgres_uri=None,
|
||||||
api_version=None,
|
api_version=None,
|
||||||
image=tag,
|
image=tag,
|
||||||
@@ -171,5 +173,5 @@ if __name__ == "__main__":
|
|||||||
except BaseException:
|
except BaseException:
|
||||||
logger.exception("Test failed")
|
logger.exception("Test failed")
|
||||||
raise
|
raise
|
||||||
|
|
||||||
logger.info("Test execution finished")
|
logger.info("Test execution finished")
|
||||||
|
|||||||
@@ -103,6 +103,8 @@ The CLI uses a `langgraph.json` configuration file with these key settings:
|
|||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Git dependencies should use credential-free URLs. The CLI conservatively scans direct `langgraph.json` dependencies, common Python package files, uv project and lock files, and common Node.js package and lock files for HTTP Git URLs with userinfo. This check is not exhaustive: generated Docker builds can copy other files, including nested requirement or constraint files, into image layers without scanning them. For private dependencies, provide short-lived credentials through your build environment's secret-backed Git credential helper. Do not store credentials in copied files such as `langgraph.json` or `pip_config_file`.
|
||||||
|
|
||||||
See the [full documentation](https://reference.langchain.com/python/langgraph-cli) for detailed configuration options.
|
See the [full documentation](https://reference.langchain.com/python/langgraph-cli) for detailed configuration options.
|
||||||
|
|
||||||
## Development
|
## Development
|
||||||
|
|||||||
@@ -48,6 +48,9 @@ def get_anonymized_params(
|
|||||||
if kwargs.get("docker_compose"):
|
if kwargs.get("docker_compose"):
|
||||||
params["docker_compose"] = True
|
params["docker_compose"] = True
|
||||||
|
|
||||||
|
if kwargs.get("debugger_port"):
|
||||||
|
params["debugger_port"] = True
|
||||||
|
|
||||||
if kwargs.get("postgres_uri"):
|
if kwargs.get("postgres_uri"):
|
||||||
params["postgres_uri"] = True
|
params["postgres_uri"] = True
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ import pathlib
|
|||||||
import shutil
|
import shutil
|
||||||
import sys
|
import sys
|
||||||
from collections.abc import Sequence
|
from collections.abc import Sequence
|
||||||
from urllib.parse import SplitResult, urlencode, urlsplit, urlunsplit
|
|
||||||
|
|
||||||
import click
|
import click
|
||||||
import click.exceptions
|
import click.exceptions
|
||||||
@@ -141,6 +140,17 @@ OPT_VERBOSE = click.option(
|
|||||||
help="Show more output from the server logs",
|
help="Show more output from the server logs",
|
||||||
)
|
)
|
||||||
OPT_WATCH = click.option("--watch", is_flag=True, help="Restart on file changes")
|
OPT_WATCH = click.option("--watch", is_flag=True, help="Restart on file changes")
|
||||||
|
OPT_DEBUGGER_PORT = click.option(
|
||||||
|
"--debugger-port",
|
||||||
|
type=int,
|
||||||
|
help="Pull the debugger image locally and serve the UI on specified port",
|
||||||
|
)
|
||||||
|
OPT_DEBUGGER_BASE_URL = click.option(
|
||||||
|
"--debugger-base-url",
|
||||||
|
type=str,
|
||||||
|
help="URL used by the debugger to access LangGraph API. Defaults to http://127.0.0.1:[PORT]",
|
||||||
|
)
|
||||||
|
|
||||||
OPT_POSTGRES_URI = click.option(
|
OPT_POSTGRES_URI = click.option(
|
||||||
"--postgres-uri",
|
"--postgres-uri",
|
||||||
help="Postgres URI to use for the database. Defaults to launching a local database",
|
help="Postgres URI to use for the database. Defaults to launching a local database",
|
||||||
@@ -232,94 +242,18 @@ cli.add_command(deploy)
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def _validated_http_url(value: str, option_name: str) -> SplitResult:
|
|
||||||
try:
|
|
||||||
parsed = urlsplit(value)
|
|
||||||
hostname = parsed.hostname
|
|
||||||
_ = parsed.port
|
|
||||||
except ValueError as exc:
|
|
||||||
raise click.UsageError(
|
|
||||||
f"{option_name} must be a valid HTTP(S) URL without credentials."
|
|
||||||
) from exc
|
|
||||||
|
|
||||||
if (
|
|
||||||
value != value.strip()
|
|
||||||
or parsed.scheme not in {"http", "https"}
|
|
||||||
or not parsed.netloc
|
|
||||||
or not hostname
|
|
||||||
or parsed.username is not None
|
|
||||||
or parsed.password is not None
|
|
||||||
):
|
|
||||||
raise click.UsageError(
|
|
||||||
f"{option_name} must be a valid HTTP(S) URL without credentials."
|
|
||||||
)
|
|
||||||
return parsed
|
|
||||||
|
|
||||||
|
|
||||||
def _studio_link(
|
|
||||||
*,
|
|
||||||
port: int,
|
|
||||||
studio_url: str | None,
|
|
||||||
api_url: str | None,
|
|
||||||
debugger_base_url: str | None,
|
|
||||||
) -> str:
|
|
||||||
if debugger_base_url is not None:
|
|
||||||
if api_url is not None and api_url != debugger_base_url:
|
|
||||||
raise click.UsageError(
|
|
||||||
"--api-url and --debugger-base-url cannot specify different URLs."
|
|
||||||
)
|
|
||||||
click.echo(
|
|
||||||
"Warning: --debugger-base-url is deprecated; use --api-url instead.",
|
|
||||||
err=True,
|
|
||||||
)
|
|
||||||
api_url = debugger_base_url
|
|
||||||
|
|
||||||
studio_url = "https://smith.langchain.com" if studio_url is None else studio_url
|
|
||||||
api_url = f"http://127.0.0.1:{port}" if api_url is None else api_url
|
|
||||||
studio_parts = _validated_http_url(studio_url, "--studio-url")
|
|
||||||
_validated_http_url(api_url, "--api-url")
|
|
||||||
if studio_parts.query or studio_parts.fragment:
|
|
||||||
raise click.UsageError(
|
|
||||||
"--studio-url must not include a query string or fragment."
|
|
||||||
)
|
|
||||||
|
|
||||||
studio_path = f"{studio_parts.path.rstrip('/')}/studio/"
|
|
||||||
return urlunsplit(
|
|
||||||
studio_parts._replace(
|
|
||||||
path=studio_path,
|
|
||||||
query=urlencode({"baseUrl": api_url}),
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@OPT_RECREATE
|
@OPT_RECREATE
|
||||||
@OPT_PULL
|
@OPT_PULL
|
||||||
@OPT_PORT
|
@OPT_PORT
|
||||||
@OPT_DOCKER_COMPOSE
|
@OPT_DOCKER_COMPOSE
|
||||||
@OPT_CONFIG
|
@OPT_CONFIG
|
||||||
@OPT_VERBOSE
|
@OPT_VERBOSE
|
||||||
|
@OPT_DEBUGGER_PORT
|
||||||
|
@OPT_DEBUGGER_BASE_URL
|
||||||
@OPT_WATCH
|
@OPT_WATCH
|
||||||
@OPT_POSTGRES_URI
|
@OPT_POSTGRES_URI
|
||||||
@OPT_API_VERSION
|
@OPT_API_VERSION
|
||||||
@OPT_ENGINE_RUNTIME_MODE
|
@OPT_ENGINE_RUNTIME_MODE
|
||||||
@click.option(
|
|
||||||
"--studio-url",
|
|
||||||
type=str,
|
|
||||||
default=None,
|
|
||||||
help="URL of the LangGraph Studio instance. Defaults to https://smith.langchain.com",
|
|
||||||
)
|
|
||||||
@click.option(
|
|
||||||
"--api-url",
|
|
||||||
type=str,
|
|
||||||
default=None,
|
|
||||||
help="URL that LangGraph Studio uses to access the API. Defaults to http://127.0.0.1:[PORT]",
|
|
||||||
)
|
|
||||||
@click.option(
|
|
||||||
"--debugger-base-url",
|
|
||||||
type=str,
|
|
||||||
default=None,
|
|
||||||
hidden=True,
|
|
||||||
)
|
|
||||||
@click.option(
|
@click.option(
|
||||||
"--image",
|
"--image",
|
||||||
type=str,
|
type=str,
|
||||||
@@ -350,21 +284,14 @@ def up(
|
|||||||
watch: bool,
|
watch: bool,
|
||||||
wait: bool,
|
wait: bool,
|
||||||
verbose: bool,
|
verbose: bool,
|
||||||
|
debugger_port: int | None,
|
||||||
|
debugger_base_url: str | None,
|
||||||
postgres_uri: str | None,
|
postgres_uri: str | None,
|
||||||
api_version: str | None,
|
api_version: str | None,
|
||||||
engine_runtime_mode: str,
|
engine_runtime_mode: str,
|
||||||
studio_url: str | None,
|
|
||||||
api_url: str | None,
|
|
||||||
debugger_base_url: str | None,
|
|
||||||
image: str | None,
|
image: str | None,
|
||||||
base_image: str | None,
|
base_image: str | None,
|
||||||
):
|
):
|
||||||
studio_link = _studio_link(
|
|
||||||
port=port,
|
|
||||||
studio_url=studio_url,
|
|
||||||
api_url=api_url,
|
|
||||||
debugger_base_url=debugger_base_url,
|
|
||||||
)
|
|
||||||
click.secho("Starting LangGraph API server...", fg="green")
|
click.secho("Starting LangGraph API server...", fg="green")
|
||||||
click.secho(
|
click.secho(
|
||||||
"""For local dev, requires env var LANGSMITH_API_KEY with access to LangSmith Deployment.
|
"""For local dev, requires env var LANGSMITH_API_KEY with access to LangSmith Deployment.
|
||||||
@@ -381,6 +308,8 @@ For production use, requires a license key in env var LANGGRAPH_CLOUD_LICENSE_KE
|
|||||||
pull=pull,
|
pull=pull,
|
||||||
watch=watch,
|
watch=watch,
|
||||||
verbose=verbose,
|
verbose=verbose,
|
||||||
|
debugger_port=debugger_port,
|
||||||
|
debugger_base_url=debugger_base_url,
|
||||||
postgres_uri=postgres_uri,
|
postgres_uri=postgres_uri,
|
||||||
api_version=api_version,
|
api_version=api_version,
|
||||||
engine_runtime_mode=engine_runtime_mode,
|
engine_runtime_mode=engine_runtime_mode,
|
||||||
@@ -408,12 +337,20 @@ For production use, requires a license key in env var LANGGRAPH_CLOUD_LICENSE_KE
|
|||||||
if "unpacking to docker.io" in line:
|
if "unpacking to docker.io" in line:
|
||||||
set("Starting...")
|
set("Starting...")
|
||||||
elif "Application startup complete" in line:
|
elif "Application startup complete" in line:
|
||||||
|
debugger_origin = (
|
||||||
|
f"http://localhost:{debugger_port}"
|
||||||
|
if debugger_port
|
||||||
|
else "https://smith.langchain.com"
|
||||||
|
)
|
||||||
|
debugger_base_url_query = (
|
||||||
|
debugger_base_url or f"http://127.0.0.1:{port}"
|
||||||
|
)
|
||||||
set("")
|
set("")
|
||||||
sys.stdout.write(
|
sys.stdout.write(
|
||||||
f"""Ready!
|
f"""Ready!
|
||||||
- API: http://localhost:{port}
|
- API: http://localhost:{port}
|
||||||
- Docs: http://localhost:{port}/docs
|
- Docs: http://localhost:{port}/docs
|
||||||
- LangGraph Studio: {studio_link}
|
- LangGraph Studio: {debugger_origin}/studio/?baseUrl={debugger_base_url_query}
|
||||||
"""
|
"""
|
||||||
)
|
)
|
||||||
sys.stdout.flush()
|
sys.stdout.flush()
|
||||||
@@ -998,6 +935,8 @@ def prepare_args_and_stdin(
|
|||||||
docker_compose: pathlib.Path | None,
|
docker_compose: pathlib.Path | None,
|
||||||
port: int,
|
port: int,
|
||||||
watch: bool,
|
watch: bool,
|
||||||
|
debugger_port: int | None = None,
|
||||||
|
debugger_base_url: str | None = None,
|
||||||
postgres_uri: str | None = None,
|
postgres_uri: str | None = None,
|
||||||
api_version: str | None = None,
|
api_version: str | None = None,
|
||||||
engine_runtime_mode: str = "combined_queue_worker",
|
engine_runtime_mode: str = "combined_queue_worker",
|
||||||
@@ -1011,6 +950,8 @@ def prepare_args_and_stdin(
|
|||||||
stdin = langgraph_cli.docker.compose(
|
stdin = langgraph_cli.docker.compose(
|
||||||
capabilities,
|
capabilities,
|
||||||
port=port,
|
port=port,
|
||||||
|
debugger_port=debugger_port,
|
||||||
|
debugger_base_url=debugger_base_url,
|
||||||
postgres_uri=postgres_uri,
|
postgres_uri=postgres_uri,
|
||||||
image=image,
|
image=image,
|
||||||
base_image=base_image,
|
base_image=base_image,
|
||||||
@@ -1048,6 +989,8 @@ def prepare(
|
|||||||
pull: bool,
|
pull: bool,
|
||||||
watch: bool,
|
watch: bool,
|
||||||
verbose: bool,
|
verbose: bool,
|
||||||
|
debugger_port: int | None = None,
|
||||||
|
debugger_base_url: str | None = None,
|
||||||
postgres_uri: str | None = None,
|
postgres_uri: str | None = None,
|
||||||
api_version: str | None = None,
|
api_version: str | None = None,
|
||||||
engine_runtime_mode: str = "combined_queue_worker",
|
engine_runtime_mode: str = "combined_queue_worker",
|
||||||
@@ -1089,6 +1032,8 @@ def prepare(
|
|||||||
docker_compose=docker_compose,
|
docker_compose=docker_compose,
|
||||||
port=port,
|
port=port,
|
||||||
watch=watch,
|
watch=watch,
|
||||||
|
debugger_port=debugger_port,
|
||||||
|
debugger_base_url=debugger_base_url or f"http://127.0.0.1:{port}",
|
||||||
postgres_uri=postgres_uri,
|
postgres_uri=postgres_uri,
|
||||||
api_version=api_version,
|
api_version=api_version,
|
||||||
engine_runtime_mode=engine_runtime_mode,
|
engine_runtime_mode=engine_runtime_mode,
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import re
|
|||||||
import shlex
|
import shlex
|
||||||
import textwrap
|
import textwrap
|
||||||
from collections import Counter
|
from collections import Counter
|
||||||
|
from collections.abc import Iterable
|
||||||
from typing import Literal, NamedTuple
|
from typing import Literal, NamedTuple
|
||||||
|
|
||||||
import click
|
import click
|
||||||
@@ -36,6 +37,10 @@ DISALLOWED_BUILD_COMMAND_CHARS = [
|
|||||||
# This blocks background execution (cmd &) while allowing command
|
# This blocks background execution (cmd &) while allowing command
|
||||||
# chaining (cmd1 && cmd2) which is common in build commands.
|
# chaining (cmd1 && cmd2) which is common in build commands.
|
||||||
_SINGLE_AMPERSAND_RE = re.compile(r"(?<!&)&(?:&&)*(?!&)")
|
_SINGLE_AMPERSAND_RE = re.compile(r"(?<!&)&(?:&&)*(?!&)")
|
||||||
|
_GIT_HTTP_AUTHORITY_RES = (
|
||||||
|
re.compile(r"git\+https?://(?P<authority>[^/\s\"']+)", re.I),
|
||||||
|
re.compile(r"\bgit\s*=\s*[\"']https?://(?P<authority>[^/\s\"']+)", re.I),
|
||||||
|
)
|
||||||
_API_VERSION_PATTERN = re.compile(
|
_API_VERSION_PATTERN = re.compile(
|
||||||
r"^(?P<major>\d+)"
|
r"^(?P<major>\d+)"
|
||||||
r"(?:\.(?P<minor>\d+))?"
|
r"(?:\.(?P<minor>\d+))?"
|
||||||
@@ -78,6 +83,62 @@ def has_disallowed_build_command_content(command: str) -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _has_git_http_url_userinfo(dependency: str) -> bool:
|
||||||
|
"""Check whether a Git HTTP URL contains userinfo."""
|
||||||
|
return any(
|
||||||
|
"@" in match.group("authority")
|
||||||
|
for pattern in _GIT_HTTP_AUTHORITY_RES
|
||||||
|
for match in pattern.finditer(dependency)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_git_http_url_userinfo(
|
||||||
|
values: Iterable[str], *, source: pathlib.Path | None = None
|
||||||
|
) -> None:
|
||||||
|
"""Reject credential-bearing Git HTTP URLs without echoing their values."""
|
||||||
|
if not any(_has_git_http_url_userinfo(value) for value in values):
|
||||||
|
return
|
||||||
|
message = (
|
||||||
|
"Git dependency URLs must not contain credentials or other URL "
|
||||||
|
"userinfo because generated Dockerfiles and image layers can retain "
|
||||||
|
"them. Use a credential-free Git URL and provide short-lived "
|
||||||
|
"credentials through your build environment's secret-backed Git "
|
||||||
|
"credential helper."
|
||||||
|
)
|
||||||
|
if source is not None:
|
||||||
|
message += f" Found in: {source}"
|
||||||
|
raise click.UsageError(message)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_git_http_url_userinfo_files(paths: Iterable[pathlib.Path]) -> None:
|
||||||
|
"""Reject credential-bearing Git HTTP URLs in dependency files."""
|
||||||
|
for path in paths:
|
||||||
|
path = path.resolve()
|
||||||
|
if not path.is_file():
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
contents = path.read_text(encoding="utf-8", errors="replace")
|
||||||
|
except OSError:
|
||||||
|
raise click.UsageError(
|
||||||
|
f"Could not inspect dependency file for embedded credentials: {path}"
|
||||||
|
) from None
|
||||||
|
_validate_git_http_url_userinfo([contents], source=path)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_local_dependency_files(config_path: pathlib.Path, config: Config) -> None:
|
||||||
|
"""Validate dependency files copied into a non-uv Python image."""
|
||||||
|
paths: list[pathlib.Path] = []
|
||||||
|
for dependency in config["dependencies"]:
|
||||||
|
if not isinstance(dependency, str) or not dependency.startswith("."):
|
||||||
|
continue
|
||||||
|
root = (config_path.parent / dependency).resolve()
|
||||||
|
paths.extend(
|
||||||
|
root / name
|
||||||
|
for name in ("requirements.txt", "pyproject.toml", "setup.py", "setup.cfg")
|
||||||
|
)
|
||||||
|
_validate_git_http_url_userinfo_files(paths)
|
||||||
|
|
||||||
|
|
||||||
MIN_PYTHON_VERSION = "3.11"
|
MIN_PYTHON_VERSION = "3.11"
|
||||||
DEFAULT_PYTHON_VERSION = "3.11"
|
DEFAULT_PYTHON_VERSION = "3.11"
|
||||||
|
|
||||||
@@ -320,7 +381,9 @@ def _get_source_kind(config: Config) -> str | None:
|
|||||||
return kind if isinstance(kind, str) else None
|
return kind if isinstance(kind, str) else None
|
||||||
|
|
||||||
|
|
||||||
def validate_config(config: Config) -> Config:
|
def validate_config(
|
||||||
|
config: Config, *, source_path: pathlib.Path | None = None
|
||||||
|
) -> Config:
|
||||||
"""Validate a configuration dictionary."""
|
"""Validate a configuration dictionary."""
|
||||||
|
|
||||||
graphs = config.get("graphs", {})
|
graphs = config.get("graphs", {})
|
||||||
@@ -415,6 +478,15 @@ def validate_config(config: Config) -> Config:
|
|||||||
' "source": {"kind": "uv", "root": ".."}'
|
' "source": {"kind": "uv", "root": ".."}'
|
||||||
)
|
)
|
||||||
|
|
||||||
|
_validate_git_http_url_userinfo(
|
||||||
|
(
|
||||||
|
dependency
|
||||||
|
for dependency in config["dependencies"]
|
||||||
|
if isinstance(dependency, str)
|
||||||
|
),
|
||||||
|
source=source_path,
|
||||||
|
)
|
||||||
|
|
||||||
source = config.get("source")
|
source = config.get("source")
|
||||||
source_kind = _get_source_kind(config)
|
source_kind = _get_source_kind(config)
|
||||||
if source is not None and not isinstance(source, dict):
|
if source is not None and not isinstance(source, dict):
|
||||||
@@ -609,7 +681,7 @@ def validate_config_file(config_path: pathlib.Path) -> Config:
|
|||||||
"""Load and validate a configuration file."""
|
"""Load and validate a configuration file."""
|
||||||
with open(config_path) as f:
|
with open(config_path) as f:
|
||||||
config = json.load(f)
|
config = json.load(f)
|
||||||
validated = validate_config(config)
|
validated = validate_config(config, source_path=config_path.resolve())
|
||||||
# Enforce the package.json doesn't enforce an
|
# Enforce the package.json doesn't enforce an
|
||||||
# incompatible Node.js version
|
# incompatible Node.js version
|
||||||
if validated.get("node_version"):
|
if validated.get("node_version"):
|
||||||
@@ -1280,6 +1352,7 @@ def python_config_to_docker(
|
|||||||
api_version=api_version,
|
api_version=api_version,
|
||||||
build_tools_to_uninstall=build_tools_to_uninstall,
|
build_tools_to_uninstall=build_tools_to_uninstall,
|
||||||
)
|
)
|
||||||
|
_validate_local_dependency_files(config_path, config)
|
||||||
if pip_installer == "auto":
|
if pip_installer == "auto":
|
||||||
if _image_supports_uv(base_image):
|
if _image_supports_uv(base_image):
|
||||||
pip_installer = "uv"
|
pip_installer = "uv"
|
||||||
@@ -1490,7 +1563,18 @@ def node_config_to_docker(
|
|||||||
) -> tuple[str, dict[str, str]]:
|
) -> tuple[str, dict[str, str]]:
|
||||||
# Calculate paths for monorepo support
|
# Calculate paths for monorepo support
|
||||||
install_root = (
|
install_root = (
|
||||||
pathlib.Path(build_context).resolve() if build_context else config_path.parent
|
pathlib.Path(build_context).resolve()
|
||||||
|
if build_context
|
||||||
|
else config_path.parent.resolve()
|
||||||
|
)
|
||||||
|
config_root = config_path.parent.resolve()
|
||||||
|
dependency_roots = (
|
||||||
|
(install_root, config_root) if install_root != config_root else (install_root,)
|
||||||
|
)
|
||||||
|
_validate_git_http_url_userinfo_files(
|
||||||
|
root / name
|
||||||
|
for root in dependency_roots
|
||||||
|
for name in ("package.json", "package-lock.json", "yarn.lock", "pnpm-lock.yaml")
|
||||||
)
|
)
|
||||||
install_cmd = install_command or _get_node_pm_install_cmd(install_root)
|
install_cmd = install_command or _get_node_pm_install_cmd(install_root)
|
||||||
if build_context:
|
if build_context:
|
||||||
|
|||||||
@@ -142,6 +142,29 @@ def check_capabilities(runner) -> DockerCapabilities:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def debugger_compose(*, port: int | None = None, base_url: str | None = None) -> dict:
|
||||||
|
if port is None:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
config = {
|
||||||
|
"langgraph-debugger": {
|
||||||
|
"image": "langchain/langgraph-debugger",
|
||||||
|
"restart": "on-failure",
|
||||||
|
"depends_on": {
|
||||||
|
"langgraph-postgres": {"condition": "service_healthy"},
|
||||||
|
},
|
||||||
|
"ports": [f'"{port}:3968"'],
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if base_url:
|
||||||
|
config["langgraph-debugger"]["environment"] = {
|
||||||
|
"VITE_STUDIO_LOCAL_GRAPH_URL": base_url
|
||||||
|
}
|
||||||
|
|
||||||
|
return config
|
||||||
|
|
||||||
|
|
||||||
# Function to convert dictionary to YAML
|
# Function to convert dictionary to YAML
|
||||||
def dict_to_yaml(d: dict, *, indent: int = 0) -> str:
|
def dict_to_yaml(d: dict, *, indent: int = 0) -> str:
|
||||||
"""Convert a dictionary to a YAML string."""
|
"""Convert a dictionary to a YAML string."""
|
||||||
@@ -168,6 +191,8 @@ def compose_as_dict(
|
|||||||
capabilities: DockerCapabilities,
|
capabilities: DockerCapabilities,
|
||||||
*,
|
*,
|
||||||
port: int,
|
port: int,
|
||||||
|
debugger_port: int | None = None,
|
||||||
|
debugger_base_url: str | None = None,
|
||||||
# postgres://user:password@host:port/database?option=value
|
# postgres://user:password@host:port/database?option=value
|
||||||
postgres_uri: str | None = None,
|
postgres_uri: str | None = None,
|
||||||
# If you are running against an already-built image, you can pass it here
|
# If you are running against an already-built image, you can pass it here
|
||||||
@@ -228,6 +253,12 @@ def compose_as_dict(
|
|||||||
else:
|
else:
|
||||||
services["langgraph-postgres"]["healthcheck"]["interval"] = "5s"
|
services["langgraph-postgres"]["healthcheck"]["interval"] = "5s"
|
||||||
|
|
||||||
|
# Add optional debugger service if debugger_port is specified
|
||||||
|
if debugger_port:
|
||||||
|
services["langgraph-debugger"] = debugger_compose(
|
||||||
|
port=debugger_port, base_url=debugger_base_url
|
||||||
|
)["langgraph-debugger"]
|
||||||
|
|
||||||
# Add langgraph-api service
|
# Add langgraph-api service
|
||||||
api_environment = {
|
api_environment = {
|
||||||
"REDIS_URI": "redis://langgraph-redis:6379",
|
"REDIS_URI": "redis://langgraph-redis:6379",
|
||||||
@@ -258,7 +289,7 @@ def compose_as_dict(
|
|||||||
"test": "python /api/healthcheck.py",
|
"test": "python /api/healthcheck.py",
|
||||||
"interval": "60s",
|
"interval": "60s",
|
||||||
"start_interval": "1s",
|
"start_interval": "1s",
|
||||||
"start_period": "60s",
|
"start_period": "10s",
|
||||||
}
|
}
|
||||||
|
|
||||||
# Final compose dictionary with volumes included if needed
|
# Final compose dictionary with volumes included if needed
|
||||||
@@ -274,6 +305,8 @@ def compose(
|
|||||||
capabilities: DockerCapabilities,
|
capabilities: DockerCapabilities,
|
||||||
*,
|
*,
|
||||||
port: int,
|
port: int,
|
||||||
|
debugger_port: int | None = None,
|
||||||
|
debugger_base_url: str | None = None,
|
||||||
# postgres://user:password@host:port/database?option=value
|
# postgres://user:password@host:port/database?option=value
|
||||||
postgres_uri: str | None = None,
|
postgres_uri: str | None = None,
|
||||||
image: str | None = None,
|
image: str | None = None,
|
||||||
@@ -285,6 +318,8 @@ def compose(
|
|||||||
compose_content = compose_as_dict(
|
compose_content = compose_as_dict(
|
||||||
capabilities,
|
capabilities,
|
||||||
port=port,
|
port=port,
|
||||||
|
debugger_port=debugger_port,
|
||||||
|
debugger_base_url=debugger_base_url,
|
||||||
postgres_uri=postgres_uri,
|
postgres_uri=postgres_uri,
|
||||||
image=image,
|
image=image,
|
||||||
base_image=base_image,
|
base_image=base_image,
|
||||||
|
|||||||
@@ -650,7 +650,8 @@ class Config(TypedDict, total=False):
|
|||||||
|
|
||||||
pip_config_file: str | None
|
pip_config_file: str | None
|
||||||
"""Optional. Path to a pip config file (e.g., "/etc/pip.conf" or "pip.ini") for controlling
|
"""Optional. Path to a pip config file (e.g., "/etc/pip.conf" or "pip.ini") for controlling
|
||||||
package installation (custom indices, credentials, etc.).
|
package installation (custom indices, timeouts, etc.). The file is copied into the
|
||||||
|
generated image, so it must not contain credentials or other secrets.
|
||||||
|
|
||||||
Only relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.
|
Only relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.
|
||||||
"""
|
"""
|
||||||
@@ -689,6 +690,9 @@ class Config(TypedDict, total=False):
|
|||||||
- "." or "./src" if you have a local Python package
|
- "." or "./src" if you have a local Python package
|
||||||
- str (aka "anthropic") for a PyPI package
|
- str (aka "anthropic") for a PyPI package
|
||||||
- "git+https://github.com/org/repo.git@main" for a Git-based package
|
- "git+https://github.com/org/repo.git@main" for a Git-based package
|
||||||
|
Git HTTP URLs must not contain userinfo such as a username or token. For private
|
||||||
|
dependencies, provide short-lived credentials through the build environment's
|
||||||
|
secret-backed Git credential helper.
|
||||||
Defaults to an empty list, meaning no additional packages installed beyond your base environment.
|
Defaults to an empty list, meaning no additional packages installed beyond your base environment.
|
||||||
|
|
||||||
This field is not supported when `source.kind` is `uv`.
|
This field is not supported when `source.kind` is `uv`.
|
||||||
|
|||||||
@@ -880,6 +880,7 @@ def python_config_to_docker_uv_lock(
|
|||||||
_get_node_pm_install_cmd,
|
_get_node_pm_install_cmd,
|
||||||
_get_pip_cleanup_lines,
|
_get_pip_cleanup_lines,
|
||||||
_image_supports_uv,
|
_image_supports_uv,
|
||||||
|
_validate_git_http_url_userinfo_files,
|
||||||
docker_tag,
|
docker_tag,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -890,11 +891,20 @@ def python_config_to_docker_uv_lock(
|
|||||||
)
|
)
|
||||||
|
|
||||||
config_root = config_path.parent.resolve()
|
config_root = config_path.parent.resolve()
|
||||||
|
source_root = config["source"].get("root", ".")
|
||||||
|
project_root = (config_root / source_root).resolve()
|
||||||
|
_validate_git_http_url_userinfo_files(
|
||||||
|
[project_root / "pyproject.toml", project_root / "uv.lock"]
|
||||||
|
)
|
||||||
|
|
||||||
install_cmd = "uv pip install --system"
|
install_cmd = "uv pip install --system"
|
||||||
_, global_reqs_pip_install, pip_config_file_str = _build_python_install_commands(
|
_, global_reqs_pip_install, pip_config_file_str = _build_python_install_commands(
|
||||||
config, install_cmd
|
config, install_cmd
|
||||||
)
|
)
|
||||||
plan = _plan_uv_lock_workspace(config_path, config)
|
plan = _plan_uv_lock_workspace(config_path, config)
|
||||||
|
_validate_git_http_url_userinfo_files(
|
||||||
|
package.pyproject_path for package in plan.install_order
|
||||||
|
)
|
||||||
|
|
||||||
_update_uv_lock_graph_paths(config_path, config, plan)
|
_update_uv_lock_graph_paths(config_path, config, plan)
|
||||||
for section, key in [
|
for section, key in [
|
||||||
|
|||||||
@@ -28,7 +28,7 @@
|
|||||||
"type": "null"
|
"type": "null"
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, credentials, etc.).\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, timeouts, etc.). The file is copied into the\ngenerated image, so it must not contain credentials or other secrets.\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
||||||
},
|
},
|
||||||
"_INTERNAL_docker_tag": {
|
"_INTERNAL_docker_tag": {
|
||||||
"anyOf": [
|
"anyOf": [
|
||||||
@@ -270,7 +270,7 @@
|
|||||||
"type": "null"
|
"type": "null"
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, credentials, etc.).\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, timeouts, etc.). The file is copied into the\ngenerated image, so it must not contain credentials or other secrets.\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
||||||
},
|
},
|
||||||
"_INTERNAL_docker_tag": {
|
"_INTERNAL_docker_tag": {
|
||||||
"anyOf": [
|
"anyOf": [
|
||||||
|
|||||||
@@ -28,7 +28,7 @@
|
|||||||
"type": "null"
|
"type": "null"
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, credentials, etc.).\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, timeouts, etc.). The file is copied into the\ngenerated image, so it must not contain credentials or other secrets.\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
||||||
},
|
},
|
||||||
"_INTERNAL_docker_tag": {
|
"_INTERNAL_docker_tag": {
|
||||||
"anyOf": [
|
"anyOf": [
|
||||||
@@ -270,7 +270,7 @@
|
|||||||
"type": "null"
|
"type": "null"
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, credentials, etc.).\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
"description": "Optional. Path to a pip config file (e.g., \"/etc/pip.conf\" or \"pip.ini\") for controlling\npackage installation (custom indices, timeouts, etc.). The file is copied into the\ngenerated image, so it must not contain credentials or other secrets.\n\nOnly relevant if Python dependencies are installed via pip. If omitted, default pip settings are used.\n"
|
||||||
},
|
},
|
||||||
"_INTERNAL_docker_tag": {
|
"_INTERNAL_docker_tag": {
|
||||||
"anyOf": [
|
"anyOf": [
|
||||||
|
|||||||
@@ -8,11 +8,10 @@ from contextlib import contextmanager
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import click
|
import click
|
||||||
import pytest
|
|
||||||
from click.testing import CliRunner
|
from click.testing import CliRunner
|
||||||
|
|
||||||
import langgraph_cli.deploy as deploy_module
|
import langgraph_cli.deploy as deploy_module
|
||||||
from langgraph_cli.cli import _studio_link, cli, prepare_args_and_stdin
|
from langgraph_cli.cli import cli, prepare_args_and_stdin
|
||||||
from langgraph_cli.config import Config, _get_pip_cleanup_lines, validate_config
|
from langgraph_cli.config import Config, _get_pip_cleanup_lines, validate_config
|
||||||
from langgraph_cli.docker import DEFAULT_POSTGRES_URI, DockerCapabilities, Version
|
from langgraph_cli.docker import DEFAULT_POSTGRES_URI, DockerCapabilities, Version
|
||||||
from langgraph_cli.util import clean_empty_lines
|
from langgraph_cli.util import clean_empty_lines
|
||||||
@@ -57,6 +56,8 @@ def test_prepare_args_and_stdin() -> None:
|
|||||||
Config(dependencies=[".", "../../.."], graphs={"agent": "agent.py:graph"})
|
Config(dependencies=[".", "../../.."], graphs={"agent": "agent.py:graph"})
|
||||||
)
|
)
|
||||||
port = 8000
|
port = 8000
|
||||||
|
debugger_port = 8001
|
||||||
|
debugger_graph_url = f"http://127.0.0.1:{port}"
|
||||||
|
|
||||||
actual_args, actual_stdin = prepare_args_and_stdin(
|
actual_args, actual_stdin = prepare_args_and_stdin(
|
||||||
capabilities=DEFAULT_DOCKER_CAPABILITIES,
|
capabilities=DEFAULT_DOCKER_CAPABILITIES,
|
||||||
@@ -64,6 +65,8 @@ def test_prepare_args_and_stdin() -> None:
|
|||||||
config=config,
|
config=config,
|
||||||
docker_compose=pathlib.Path("custom-docker-compose.yml"),
|
docker_compose=pathlib.Path("custom-docker-compose.yml"),
|
||||||
port=port,
|
port=port,
|
||||||
|
debugger_port=debugger_port,
|
||||||
|
debugger_base_url=debugger_graph_url,
|
||||||
watch=True,
|
watch=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -107,6 +110,16 @@ services:
|
|||||||
retries: 5
|
retries: 5
|
||||||
interval: 60s
|
interval: 60s
|
||||||
start_interval: 1s
|
start_interval: 1s
|
||||||
|
langgraph-debugger:
|
||||||
|
image: langchain/langgraph-debugger
|
||||||
|
restart: on-failure
|
||||||
|
depends_on:
|
||||||
|
langgraph-postgres:
|
||||||
|
condition: service_healthy
|
||||||
|
ports:
|
||||||
|
- "{debugger_port}:3968"
|
||||||
|
environment:
|
||||||
|
VITE_STUDIO_LOCAL_GRAPH_URL: {debugger_graph_url}
|
||||||
langgraph-api:
|
langgraph-api:
|
||||||
ports:
|
ports:
|
||||||
- "8000:8000"
|
- "8000:8000"
|
||||||
@@ -122,7 +135,7 @@ services:
|
|||||||
test: python /api/healthcheck.py
|
test: python /api/healthcheck.py
|
||||||
interval: 60s
|
interval: 60s
|
||||||
start_interval: 1s
|
start_interval: 1s
|
||||||
start_period: 60s
|
start_period: 10s
|
||||||
|
|
||||||
pull_policy: build
|
pull_policy: build
|
||||||
build:
|
build:
|
||||||
@@ -165,6 +178,8 @@ def test_prepare_args_and_stdin_with_image() -> None:
|
|||||||
Config(dependencies=[".", "../../.."], graphs={"agent": "agent.py:graph"})
|
Config(dependencies=[".", "../../.."], graphs={"agent": "agent.py:graph"})
|
||||||
)
|
)
|
||||||
port = 8000
|
port = 8000
|
||||||
|
debugger_port = 8001
|
||||||
|
debugger_graph_url = f"http://127.0.0.1:{port}"
|
||||||
|
|
||||||
actual_args, actual_stdin = prepare_args_and_stdin(
|
actual_args, actual_stdin = prepare_args_and_stdin(
|
||||||
capabilities=DEFAULT_DOCKER_CAPABILITIES,
|
capabilities=DEFAULT_DOCKER_CAPABILITIES,
|
||||||
@@ -172,6 +187,8 @@ def test_prepare_args_and_stdin_with_image() -> None:
|
|||||||
config=config,
|
config=config,
|
||||||
docker_compose=pathlib.Path("custom-docker-compose.yml"),
|
docker_compose=pathlib.Path("custom-docker-compose.yml"),
|
||||||
port=port,
|
port=port,
|
||||||
|
debugger_port=debugger_port,
|
||||||
|
debugger_base_url=debugger_graph_url,
|
||||||
watch=True,
|
watch=True,
|
||||||
image="my-cool-image",
|
image="my-cool-image",
|
||||||
)
|
)
|
||||||
@@ -216,6 +233,16 @@ services:
|
|||||||
retries: 5
|
retries: 5
|
||||||
interval: 60s
|
interval: 60s
|
||||||
start_interval: 1s
|
start_interval: 1s
|
||||||
|
langgraph-debugger:
|
||||||
|
image: langchain/langgraph-debugger
|
||||||
|
restart: on-failure
|
||||||
|
depends_on:
|
||||||
|
langgraph-postgres:
|
||||||
|
condition: service_healthy
|
||||||
|
ports:
|
||||||
|
- "{debugger_port}:3968"
|
||||||
|
environment:
|
||||||
|
VITE_STUDIO_LOCAL_GRAPH_URL: {debugger_graph_url}
|
||||||
langgraph-api:
|
langgraph-api:
|
||||||
ports:
|
ports:
|
||||||
- "8000:8000"
|
- "8000:8000"
|
||||||
@@ -232,7 +259,7 @@ services:
|
|||||||
test: python /api/healthcheck.py
|
test: python /api/healthcheck.py
|
||||||
interval: 60s
|
interval: 60s
|
||||||
start_interval: 1s
|
start_interval: 1s
|
||||||
start_period: 60s
|
start_period: 10s
|
||||||
|
|
||||||
|
|
||||||
develop:
|
develop:
|
||||||
@@ -262,82 +289,6 @@ def test_version_option() -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_up_help_shows_hosted_studio_options() -> None:
|
|
||||||
result = CliRunner().invoke(cli, ["up", "--help"])
|
|
||||||
|
|
||||||
assert result.exit_code == 0, result.output
|
|
||||||
assert "--studio-url" in result.output
|
|
||||||
assert "--api-url" in result.output
|
|
||||||
assert "--debugger-port" not in result.output
|
|
||||||
assert "--debugger-base-url" not in result.output
|
|
||||||
|
|
||||||
|
|
||||||
def test_studio_link_defaults_to_hosted_studio() -> None:
|
|
||||||
assert _studio_link(
|
|
||||||
port=8123,
|
|
||||||
studio_url=None,
|
|
||||||
api_url=None,
|
|
||||||
debugger_base_url=None,
|
|
||||||
) == ("https://smith.langchain.com/studio/?baseUrl=http%3A%2F%2F127.0.0.1%3A8123")
|
|
||||||
|
|
||||||
|
|
||||||
def test_studio_link_supports_self_hosted_and_remote_urls() -> None:
|
|
||||||
assert _studio_link(
|
|
||||||
port=8123,
|
|
||||||
studio_url="https://langsmith.example.com/prefix/",
|
|
||||||
api_url="https://api.example.com/graph?tenant=a®ion=eu",
|
|
||||||
debugger_base_url=None,
|
|
||||||
) == (
|
|
||||||
"https://langsmith.example.com/prefix/studio/"
|
|
||||||
"?baseUrl=https%3A%2F%2Fapi.example.com%2Fgraph%3Ftenant%3Da%26region%3Deu"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_studio_link_supports_deprecated_debugger_base_url(capsys) -> None:
|
|
||||||
assert _studio_link(
|
|
||||||
port=8123,
|
|
||||||
studio_url=None,
|
|
||||||
api_url=None,
|
|
||||||
debugger_base_url="https://api.example.com",
|
|
||||||
).endswith("?baseUrl=https%3A%2F%2Fapi.example.com")
|
|
||||||
assert "--debugger-base-url is deprecated; use --api-url" in capsys.readouterr().err
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
("studio_url", "api_url"),
|
|
||||||
[
|
|
||||||
("javascript:alert(1)", None),
|
|
||||||
("https://user:password@example.com", None),
|
|
||||||
("https://smith.langchain.com?workspace=test", None),
|
|
||||||
(None, "file:///tmp/langgraph.sock"),
|
|
||||||
(None, "https://user:password@example.com"),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_studio_link_rejects_unsafe_urls(
|
|
||||||
studio_url: str | None, api_url: str | None
|
|
||||||
) -> None:
|
|
||||||
with pytest.raises(click.UsageError):
|
|
||||||
_studio_link(
|
|
||||||
port=8123,
|
|
||||||
studio_url=studio_url,
|
|
||||||
api_url=api_url,
|
|
||||||
debugger_base_url=None,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_studio_link_rejects_conflicting_api_url_aliases() -> None:
|
|
||||||
with pytest.raises(
|
|
||||||
click.UsageError,
|
|
||||||
match="cannot specify different URLs",
|
|
||||||
):
|
|
||||||
_studio_link(
|
|
||||||
port=8123,
|
|
||||||
studio_url=None,
|
|
||||||
api_url="https://api.example.com",
|
|
||||||
debugger_base_url="https://other.example.com",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_top_level_help_shows_deploy_subcommands() -> None:
|
def test_top_level_help_shows_deploy_subcommands() -> None:
|
||||||
runner = CliRunner()
|
runner = CliRunner()
|
||||||
|
|
||||||
|
|||||||
@@ -255,6 +255,243 @@ def test_validate_config():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"dependency",
|
||||||
|
[
|
||||||
|
"git+https://user:secret-token@github.com/org/private.git@main",
|
||||||
|
"private-package @ git+http://token@github.com/org/private.git",
|
||||||
|
"git+HTTPS://user%40example.com:secret%2Ftoken@github.com/org/private.git",
|
||||||
|
"git+https://${GIT_TOKEN}@github.com/org/private.git",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_validate_config_rejects_git_http_url_userinfo(dependency: str):
|
||||||
|
with pytest.raises(click.UsageError) as exc_info:
|
||||||
|
validate_config(
|
||||||
|
{
|
||||||
|
"python_version": "3.11",
|
||||||
|
"dependencies": [dependency],
|
||||||
|
"graphs": {"agent": "./agent.py:graph"},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
message = str(exc_info.value)
|
||||||
|
assert "must not contain credentials or other URL userinfo" in message
|
||||||
|
assert "secret-token" not in message
|
||||||
|
assert "secret%2Ftoken" not in message
|
||||||
|
|
||||||
|
|
||||||
|
def test_validate_config_file_reports_source_for_git_http_url_userinfo(
|
||||||
|
tmp_path: pathlib.Path,
|
||||||
|
):
|
||||||
|
config_path = tmp_path / "langgraph.json"
|
||||||
|
config_path.write_text(
|
||||||
|
json.dumps(
|
||||||
|
{
|
||||||
|
"python_version": "3.11",
|
||||||
|
"dependencies": ["git+https://secret-token@github.com/org/private.git"],
|
||||||
|
"graphs": {"agent": "./agent.py:graph"},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(click.UsageError) as exc_info:
|
||||||
|
validate_config_file(config_path)
|
||||||
|
|
||||||
|
message = str(exc_info.value)
|
||||||
|
assert "secret-token" not in message
|
||||||
|
assert f"Found in: {config_path.resolve()}" in message
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"manifest", ["package.json", "package-lock.json", "yarn.lock", "pnpm-lock.yaml"]
|
||||||
|
)
|
||||||
|
def test_config_to_docker_rejects_git_http_url_userinfo_in_node_files(
|
||||||
|
tmp_path: pathlib.Path, manifest: str
|
||||||
|
):
|
||||||
|
config_path = tmp_path / "langgraph.json"
|
||||||
|
config_path.write_text("{}\n")
|
||||||
|
(tmp_path / "agent.js").write_text("export const graph = {};\n")
|
||||||
|
(tmp_path / "package.json").write_text('{"name":"agent"}\n')
|
||||||
|
(tmp_path / manifest).write_text(
|
||||||
|
'"priv": "git+https://user:secret-token@github.com/org/private.git"\n'
|
||||||
|
)
|
||||||
|
config = validate_config(
|
||||||
|
{
|
||||||
|
"node_version": "20",
|
||||||
|
"graphs": {"agent": "./agent.js:graph"},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(click.UsageError) as exc_info:
|
||||||
|
config_to_docker(
|
||||||
|
config_path,
|
||||||
|
config,
|
||||||
|
base_image="langchain/langgraphjs-api",
|
||||||
|
)
|
||||||
|
|
||||||
|
message = str(exc_info.value)
|
||||||
|
assert "must not contain credentials or other URL userinfo" in message
|
||||||
|
assert "secret-token" not in message
|
||||||
|
assert f"Found in: {(tmp_path / manifest).resolve()}" in message
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_to_docker_allows_node_git_urls_without_http_userinfo(
|
||||||
|
tmp_path: pathlib.Path,
|
||||||
|
):
|
||||||
|
config_path = tmp_path / "langgraph.json"
|
||||||
|
config_path.write_text("{}\n")
|
||||||
|
(tmp_path / "agent.js").write_text("export const graph = {};\n")
|
||||||
|
(tmp_path / "package.json").write_text(
|
||||||
|
'{"dependencies":{"public":"git+https://github.com/org/public.git"}}\n'
|
||||||
|
)
|
||||||
|
config = validate_config(
|
||||||
|
{
|
||||||
|
"node_version": "20",
|
||||||
|
"graphs": {"agent": "./agent.js:graph"},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
docker, _ = config_to_docker(
|
||||||
|
config_path,
|
||||||
|
config,
|
||||||
|
base_image="langchain/langgraphjs-api",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert f"ADD . /deps/{tmp_path.name}" in docker
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_to_docker_rejects_git_http_url_userinfo_in_node_workspace(
|
||||||
|
tmp_path: pathlib.Path,
|
||||||
|
):
|
||||||
|
config_root = tmp_path / "apps" / "agent"
|
||||||
|
config_root.mkdir(parents=True)
|
||||||
|
config_path = config_root / "langgraph.json"
|
||||||
|
config_path.write_text("{}\n")
|
||||||
|
(config_root / "agent.js").write_text("export const graph = {};\n")
|
||||||
|
(config_root / "package.json").write_text(
|
||||||
|
'{"dependencies":{"priv":"git+https://secret-token@github.com/org/private.git"}}\n'
|
||||||
|
)
|
||||||
|
(tmp_path / "package.json").write_text('{"name":"workspace"}\n')
|
||||||
|
config = validate_config(
|
||||||
|
{
|
||||||
|
"node_version": "20",
|
||||||
|
"graphs": {"agent": "./agent.js:graph"},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(click.UsageError) as exc_info:
|
||||||
|
config_to_docker(
|
||||||
|
config_path,
|
||||||
|
config,
|
||||||
|
base_image="langchain/langgraphjs-api",
|
||||||
|
build_context=str(tmp_path),
|
||||||
|
)
|
||||||
|
|
||||||
|
message = str(exc_info.value)
|
||||||
|
assert "secret-token" not in message
|
||||||
|
assert f"Found in: {(config_root / 'package.json').resolve()}" in message
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"dependency",
|
||||||
|
[
|
||||||
|
"git+https://github.com/org/public.git@main",
|
||||||
|
"private-package @ git+https://github.com/org/private.git@main",
|
||||||
|
"git+ssh://git@github.com/org/private.git@main",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_validate_config_allows_git_urls_without_http_userinfo(dependency: str):
|
||||||
|
config = validate_config(
|
||||||
|
{
|
||||||
|
"python_version": "3.11",
|
||||||
|
"dependencies": [dependency],
|
||||||
|
"graphs": {"agent": "./agent.py:graph"},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert config["dependencies"] == [dependency]
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_to_docker_rejects_git_http_url_userinfo_in_requirements(
|
||||||
|
tmp_path: pathlib.Path,
|
||||||
|
):
|
||||||
|
config_path = tmp_path / "langgraph.json"
|
||||||
|
config_path.write_text("{}\n")
|
||||||
|
(tmp_path / "agent.py").write_text("graph = object()\n")
|
||||||
|
(tmp_path / "requirements.txt").write_text(
|
||||||
|
"private @ git+https://secret-token@github.com/org/private.git\n"
|
||||||
|
)
|
||||||
|
config = validate_config(
|
||||||
|
{
|
||||||
|
"python_version": "3.11",
|
||||||
|
"dependencies": ["."],
|
||||||
|
"graphs": {"agent": "./agent.py:graph"},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(click.UsageError) as exc_info:
|
||||||
|
config_to_docker(
|
||||||
|
config_path,
|
||||||
|
config,
|
||||||
|
base_image="langchain/langgraph-api:0.2.47",
|
||||||
|
)
|
||||||
|
|
||||||
|
message = str(exc_info.value)
|
||||||
|
assert "must not contain credentials or other URL userinfo" in message
|
||||||
|
assert "secret-token" not in message
|
||||||
|
assert f"Found in: {(tmp_path / 'requirements.txt').resolve()}" in message
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("manifest", ["pyproject.toml", "uv.lock"])
|
||||||
|
def test_config_to_docker_rejects_git_http_url_userinfo_in_uv_files(
|
||||||
|
tmp_path: pathlib.Path, manifest: str
|
||||||
|
):
|
||||||
|
config_path = tmp_path / "langgraph.json"
|
||||||
|
config_path.write_text("{}\n")
|
||||||
|
(tmp_path / "src").mkdir()
|
||||||
|
(tmp_path / "src" / "agent.py").write_text("graph = object()\n")
|
||||||
|
pyproject = textwrap.dedent(
|
||||||
|
"""
|
||||||
|
[project]
|
||||||
|
name = "agent"
|
||||||
|
version = "0.1.0"
|
||||||
|
dependencies = ["private"]
|
||||||
|
|
||||||
|
[tool.uv.sources]
|
||||||
|
private = { git = "https://github.com/org/private.git" }
|
||||||
|
"""
|
||||||
|
).strip()
|
||||||
|
uv_lock = "# uv lock file\n"
|
||||||
|
if manifest == "pyproject.toml":
|
||||||
|
pyproject = pyproject.replace(
|
||||||
|
"https://github.com", "https://secret-token@github.com"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
uv_lock += (
|
||||||
|
'source = { git = "https://secret-token@github.com/org/private.git" }\n'
|
||||||
|
)
|
||||||
|
(tmp_path / "pyproject.toml").write_text(pyproject + "\n")
|
||||||
|
(tmp_path / "uv.lock").write_text(uv_lock)
|
||||||
|
config = validate_config(
|
||||||
|
{
|
||||||
|
"python_version": "3.11",
|
||||||
|
"graphs": {"agent": "./src/agent.py:graph"},
|
||||||
|
"source": {"kind": "uv"},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(click.UsageError) as exc_info:
|
||||||
|
config_to_docker(
|
||||||
|
config_path,
|
||||||
|
config,
|
||||||
|
base_image="langchain/langgraph-api:0.2.47",
|
||||||
|
)
|
||||||
|
|
||||||
|
message = str(exc_info.value)
|
||||||
|
assert "must not contain credentials or other URL userinfo" in message
|
||||||
|
assert "secret-token" not in message
|
||||||
|
|
||||||
|
|
||||||
def test_validate_config_image_distro():
|
def test_validate_config_image_distro():
|
||||||
"""Test validation of image_distro field."""
|
"""Test validation of image_distro field."""
|
||||||
# Valid image_distro values should work
|
# Valid image_distro values should work
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ DEFAULT_DOCKER_CAPABILITIES = DockerCapabilities(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_compose_with_custom_db():
|
def test_compose_with_no_debugger_and_custom_db():
|
||||||
port = 8123
|
port = 8123
|
||||||
custom_postgres_uri = "custom_postgres_uri"
|
custom_postgres_uri = "custom_postgres_uri"
|
||||||
actual_compose_str = compose(
|
actual_compose_str = compose(
|
||||||
@@ -42,7 +42,7 @@ def test_compose_with_custom_db():
|
|||||||
assert clean_empty_lines(actual_compose_str) == expected_compose_str
|
assert clean_empty_lines(actual_compose_str) == expected_compose_str
|
||||||
|
|
||||||
|
|
||||||
def test_compose_with_custom_db_and_healthcheck():
|
def test_compose_with_no_debugger_and_custom_db_with_healthcheck():
|
||||||
port = 8123
|
port = 8123
|
||||||
custom_postgres_uri = "custom_postgres_uri"
|
custom_postgres_uri = "custom_postgres_uri"
|
||||||
actual_compose_str = compose(
|
actual_compose_str = compose(
|
||||||
@@ -71,11 +71,39 @@ def test_compose_with_custom_db_and_healthcheck():
|
|||||||
test: python /api/healthcheck.py
|
test: python /api/healthcheck.py
|
||||||
interval: 60s
|
interval: 60s
|
||||||
start_interval: 1s
|
start_interval: 1s
|
||||||
start_period: 60s"""
|
start_period: 10s"""
|
||||||
assert clean_empty_lines(actual_compose_str) == expected_compose_str
|
assert clean_empty_lines(actual_compose_str) == expected_compose_str
|
||||||
|
|
||||||
|
|
||||||
def test_compose_with_default_db():
|
def test_compose_with_debugger_and_custom_db():
|
||||||
|
port = 8123
|
||||||
|
custom_postgres_uri = "custom_postgres_uri"
|
||||||
|
actual_compose_str = compose(
|
||||||
|
DEFAULT_DOCKER_CAPABILITIES,
|
||||||
|
port=port,
|
||||||
|
postgres_uri=custom_postgres_uri,
|
||||||
|
)
|
||||||
|
expected_compose_str = f"""services:
|
||||||
|
langgraph-redis:
|
||||||
|
image: redis:6
|
||||||
|
healthcheck:
|
||||||
|
test: redis-cli ping
|
||||||
|
interval: 5s
|
||||||
|
timeout: 1s
|
||||||
|
retries: 5
|
||||||
|
langgraph-api:
|
||||||
|
ports:
|
||||||
|
- "{port}:8000"
|
||||||
|
depends_on:
|
||||||
|
langgraph-redis:
|
||||||
|
condition: service_healthy
|
||||||
|
environment:
|
||||||
|
REDIS_URI: redis://langgraph-redis:6379
|
||||||
|
POSTGRES_URI: {custom_postgres_uri}"""
|
||||||
|
assert clean_empty_lines(actual_compose_str) == expected_compose_str
|
||||||
|
|
||||||
|
|
||||||
|
def test_compose_with_debugger_and_default_db():
|
||||||
port = 8123
|
port = 8123
|
||||||
actual_compose_str = compose(DEFAULT_DOCKER_CAPABILITIES, port=port)
|
actual_compose_str = compose(DEFAULT_DOCKER_CAPABILITIES, port=port)
|
||||||
expected_compose_str = f"""volumes:
|
expected_compose_str = f"""volumes:
|
||||||
@@ -274,6 +302,72 @@ def test_compose_with_api_version_and_custom_postgres():
|
|||||||
assert clean_empty_lines(actual_compose_str) == expected_compose_str
|
assert clean_empty_lines(actual_compose_str) == expected_compose_str
|
||||||
|
|
||||||
|
|
||||||
|
def test_compose_with_api_version_and_debugger():
|
||||||
|
"""Test compose function with api_version and debugger port."""
|
||||||
|
port = 8123
|
||||||
|
debugger_port = 8001
|
||||||
|
api_version = "0.2.74"
|
||||||
|
|
||||||
|
actual_compose_str = compose(
|
||||||
|
DEFAULT_DOCKER_CAPABILITIES,
|
||||||
|
port=port,
|
||||||
|
api_version=api_version,
|
||||||
|
debugger_port=debugger_port,
|
||||||
|
)
|
||||||
|
|
||||||
|
expected_compose_str = f"""volumes:
|
||||||
|
langgraph-data:
|
||||||
|
driver: local
|
||||||
|
services:
|
||||||
|
langgraph-redis:
|
||||||
|
image: redis:6
|
||||||
|
healthcheck:
|
||||||
|
test: redis-cli ping
|
||||||
|
interval: 5s
|
||||||
|
timeout: 1s
|
||||||
|
retries: 5
|
||||||
|
langgraph-postgres:
|
||||||
|
image: pgvector/pgvector:pg16
|
||||||
|
ports:
|
||||||
|
- "5433:5432"
|
||||||
|
environment:
|
||||||
|
POSTGRES_DB: postgres
|
||||||
|
POSTGRES_USER: postgres
|
||||||
|
POSTGRES_PASSWORD: postgres
|
||||||
|
command:
|
||||||
|
- postgres
|
||||||
|
- -c
|
||||||
|
- shared_preload_libraries=vector
|
||||||
|
volumes:
|
||||||
|
- langgraph-data:/var/lib/postgresql/data
|
||||||
|
healthcheck:
|
||||||
|
test: pg_isready -U postgres
|
||||||
|
start_period: 10s
|
||||||
|
timeout: 1s
|
||||||
|
retries: 5
|
||||||
|
interval: 5s
|
||||||
|
langgraph-debugger:
|
||||||
|
image: langchain/langgraph-debugger
|
||||||
|
restart: on-failure
|
||||||
|
depends_on:
|
||||||
|
langgraph-postgres:
|
||||||
|
condition: service_healthy
|
||||||
|
ports:
|
||||||
|
- "{debugger_port}:3968"
|
||||||
|
langgraph-api:
|
||||||
|
ports:
|
||||||
|
- "{port}:8000"
|
||||||
|
depends_on:
|
||||||
|
langgraph-redis:
|
||||||
|
condition: service_healthy
|
||||||
|
langgraph-postgres:
|
||||||
|
condition: service_healthy
|
||||||
|
environment:
|
||||||
|
REDIS_URI: redis://langgraph-redis:6379
|
||||||
|
POSTGRES_URI: {DEFAULT_POSTGRES_URI}"""
|
||||||
|
assert clean_empty_lines(actual_compose_str) == expected_compose_str
|
||||||
|
|
||||||
|
|
||||||
def test_compose_distributed_mode_with_custom_db():
|
def test_compose_distributed_mode_with_custom_db():
|
||||||
"""Test compose with engine_runtime_mode='distributed' adds N_JOBS_PER_WORKER=0."""
|
"""Test compose with engine_runtime_mode='distributed' adds N_JOBS_PER_WORKER=0."""
|
||||||
port = 8123
|
port = 8123
|
||||||
|
|||||||
@@ -8,13 +8,15 @@ from typing import Any, cast
|
|||||||
from langchain_core.runnables import RunnableConfig
|
from langchain_core.runnables import RunnableConfig
|
||||||
from langgraph.checkpoint.base import (
|
from langgraph.checkpoint.base import (
|
||||||
BaseCheckpointSaver,
|
BaseCheckpointSaver,
|
||||||
|
ChannelVersions,
|
||||||
Checkpoint,
|
Checkpoint,
|
||||||
|
PendingWrite,
|
||||||
)
|
)
|
||||||
from langgraph.checkpoint.base.id import uuid6
|
from langgraph.checkpoint.base.id import uuid6
|
||||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||||
|
|
||||||
from langgraph._internal._config import DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT
|
from langgraph._internal._config import DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT
|
||||||
from langgraph._internal._constants import PUSH
|
from langgraph._internal._constants import INTERRUPT, PUSH
|
||||||
from langgraph._internal._typing import MISSING
|
from langgraph._internal._typing import MISSING
|
||||||
from langgraph.channels.base import BaseChannel
|
from langgraph.channels.base import BaseChannel
|
||||||
from langgraph.channels.delta import DeltaChannel
|
from langgraph.channels.delta import DeltaChannel
|
||||||
@@ -89,6 +91,23 @@ def get_delta_channels_from_all_channels(
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def delta_channels_with_pending_writes(
|
||||||
|
specs: Mapping[str, Any],
|
||||||
|
pending_writes: Iterable[PendingWrite] | None,
|
||||||
|
) -> set[str]:
|
||||||
|
"""DeltaChannels a branch starting from this checkpoint must snapshot.
|
||||||
|
|
||||||
|
A checkpoint's pending writes belong to the child that consumed them, and
|
||||||
|
nothing records which child that was. A new branch snapshots every delta
|
||||||
|
channel they touch, so its ancestor walk never replays them.
|
||||||
|
"""
|
||||||
|
return {
|
||||||
|
ch
|
||||||
|
for _, ch, _ in pending_writes or ()
|
||||||
|
if isinstance(specs.get(ch), DeltaChannel)
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def create_metadata_for_update_state_api(
|
def create_metadata_for_update_state_api(
|
||||||
channels: Mapping[str, BaseChannel],
|
channels: Mapping[str, BaseChannel],
|
||||||
updated_channels: set[str],
|
updated_channels: set[str],
|
||||||
@@ -122,6 +141,7 @@ def create_checkpoint_plan_for_update_state_api(
|
|||||||
parents: dict[str, Any],
|
parents: dict[str, Any],
|
||||||
saved_metadata: Mapping[str, Any] | None,
|
saved_metadata: Mapping[str, Any] | None,
|
||||||
is_fresh_thread: bool,
|
is_fresh_thread: bool,
|
||||||
|
fork_channels: set[str],
|
||||||
) -> tuple[set[str], dict[str, Any]]:
|
) -> tuple[set[str], dict[str, Any]]:
|
||||||
"""Return ``(channels_to_snapshot, metadata)`` for an update_state head."""
|
"""Return ``(channels_to_snapshot, metadata)`` for an update_state head."""
|
||||||
metadata: dict[str, Any] = {
|
metadata: dict[str, Any] = {
|
||||||
@@ -137,7 +157,9 @@ def create_checkpoint_plan_for_update_state_api(
|
|||||||
updated_channels,
|
updated_channels,
|
||||||
prev_metadata=saved_metadata,
|
prev_metadata=saved_metadata,
|
||||||
)
|
)
|
||||||
channels_to_snapshot = delta_channels_to_snapshot(channels, new_counters)
|
channels_to_snapshot = (
|
||||||
|
delta_channels_to_snapshot(channels, new_counters) | fork_channels
|
||||||
|
)
|
||||||
for k in channels_to_snapshot:
|
for k in channels_to_snapshot:
|
||||||
new_counters[k] = (0, 0)
|
new_counters[k] = (0, 0)
|
||||||
non_zero = {k: v for k, v in new_counters.items() if v != (0, 0)}
|
non_zero = {k: v for k, v in new_counters.items() if v != (0, 0)}
|
||||||
@@ -167,6 +189,7 @@ def create_checkpoint(
|
|||||||
"""
|
"""
|
||||||
ts = datetime.now(timezone.utc).isoformat()
|
ts = datetime.now(timezone.utc).isoformat()
|
||||||
channels_to_snapshot = channels_to_snapshot or set()
|
channels_to_snapshot = channels_to_snapshot or set()
|
||||||
|
bumped: dict[str, tuple[Any, Any]] = {}
|
||||||
if channels is None:
|
if channels is None:
|
||||||
values = checkpoint["channel_values"]
|
values = checkpoint["channel_values"]
|
||||||
channel_versions = checkpoint["channel_versions"]
|
channel_versions = checkpoint["channel_versions"]
|
||||||
@@ -174,30 +197,29 @@ def create_checkpoint(
|
|||||||
values = {}
|
values = {}
|
||||||
channel_versions = dict(checkpoint["channel_versions"])
|
channel_versions = dict(checkpoint["channel_versions"])
|
||||||
for k in channels:
|
for k in channels:
|
||||||
if k not in channel_versions:
|
|
||||||
continue
|
|
||||||
ch = channels[k]
|
ch = channels[k]
|
||||||
|
if k not in channel_versions:
|
||||||
|
# A forced snapshot of a never-written channel still has to
|
||||||
|
# land to stop the ancestor walk, and `put` only stores blobs
|
||||||
|
# for versioned channels.
|
||||||
|
if k in channels_to_snapshot and get_next_version is not None:
|
||||||
|
channel_versions[k] = get_next_version(None, None)
|
||||||
|
bumped[k] = (None, channel_versions[k])
|
||||||
|
values[k] = _DeltaSnapshot(
|
||||||
|
ch.get() if ch.is_available() else ch.typ()
|
||||||
|
)
|
||||||
|
continue
|
||||||
if k in channels_to_snapshot:
|
if k in channels_to_snapshot:
|
||||||
# Callers force a full snapshot blob here: exit mode when a
|
# `put` only stores a blob for a channel whose version moved,
|
||||||
# delta channel reaches its snapshot cadence, and update_state
|
# so snapshotting a channel this step did not write needs a
|
||||||
# on a fresh thread (no ancestor to replay writes from). The
|
# bump: exit mode reaching the cadence on a superstep that
|
||||||
# manual version-bump below only applies to the exit-mode case.
|
# skipped the channel, and a fork's first checkpoint.
|
||||||
#
|
|
||||||
# In exit mode, the snapshot decision is deferred to exit
|
|
||||||
# time (intermediate steps have do_checkpoint=False). The
|
|
||||||
# channel's count may have reached snapshot_frequency over
|
|
||||||
# several supersteps, but the LAST superstep may not have
|
|
||||||
# written to this channel. In that case apply_writes()
|
|
||||||
# (in _algo.py) didn't bump this channel's version, so
|
|
||||||
# saver.put() wouldn't include it in new_versions and
|
|
||||||
# the snapshot blob would be silently dropped. The manual
|
|
||||||
# bump below closes the gap. In sync/async durability this
|
|
||||||
# branch is effectively dead code (the step that pushes
|
|
||||||
# the count to freq always writes the channel).
|
|
||||||
if get_next_version is not None and (
|
if get_next_version is not None and (
|
||||||
updated_channels is None or k not in updated_channels
|
updated_channels is None or k not in updated_channels
|
||||||
):
|
):
|
||||||
channel_versions[k] = get_next_version(channel_versions[k], None)
|
old = channel_versions[k]
|
||||||
|
channel_versions[k] = get_next_version(old, None)
|
||||||
|
bumped[k] = (old, channel_versions[k])
|
||||||
values[k] = _DeltaSnapshot(ch.get())
|
values[k] = _DeltaSnapshot(ch.get())
|
||||||
else:
|
else:
|
||||||
v = ch.checkpoint()
|
v = ch.checkpoint()
|
||||||
@@ -209,11 +231,30 @@ def create_checkpoint(
|
|||||||
id=id or str(uuid6(clock_seq=step)),
|
id=id or str(uuid6(clock_seq=step)),
|
||||||
channel_values=values,
|
channel_values=values,
|
||||||
channel_versions=channel_versions,
|
channel_versions=channel_versions,
|
||||||
versions_seen=checkpoint["versions_seen"],
|
versions_seen=_mark_bumps_seen(checkpoint["versions_seen"], bumped),
|
||||||
updated_channels=None if updated_channels is None else sorted(updated_channels),
|
updated_channels=None if updated_channels is None else sorted(updated_channels),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _mark_bumps_seen(
|
||||||
|
versions_seen: dict[str, ChannelVersions],
|
||||||
|
bumped: Mapping[str, tuple[Any, Any]],
|
||||||
|
) -> dict[str, ChannelVersions]:
|
||||||
|
"""Advance whoever had seen a bumped channel's old version to the new one.
|
||||||
|
|
||||||
|
A bump that only stores a snapshot is not a write. Left unseen, it would
|
||||||
|
re-fire `interrupt_before` and rerun the channel's subscribers.
|
||||||
|
"""
|
||||||
|
if not bumped:
|
||||||
|
return versions_seen
|
||||||
|
out: dict[str, ChannelVersions] = {}
|
||||||
|
for node, seen in {INTERRUPT: {}, **versions_seen}.items():
|
||||||
|
advanced = {k: new for k, (old, new) in bumped.items() if seen.get(k) == old}
|
||||||
|
if advanced or node in versions_seen:
|
||||||
|
out[node] = {**seen, **advanced}
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
def _needs_replay(spec: BaseChannel, stored: object) -> bool:
|
def _needs_replay(spec: BaseChannel, stored: object) -> bool:
|
||||||
"""True if `spec` is a `DeltaChannel` and no value is stored at this
|
"""True if `spec` is a `DeltaChannel` and no value is stored at this
|
||||||
checkpoint, requiring an ancestor walk to reconstruct.
|
checkpoint, requiring an ancestor walk to reconstruct.
|
||||||
|
|||||||
@@ -102,6 +102,7 @@ from langgraph.pregel._checkpoint import (
|
|||||||
copy_checkpoint,
|
copy_checkpoint,
|
||||||
create_checkpoint,
|
create_checkpoint,
|
||||||
delta_channels_to_snapshot,
|
delta_channels_to_snapshot,
|
||||||
|
delta_channels_with_pending_writes,
|
||||||
empty_checkpoint,
|
empty_checkpoint,
|
||||||
exit_delta_task_id,
|
exit_delta_task_id,
|
||||||
)
|
)
|
||||||
@@ -222,10 +223,13 @@ class PregelLoop:
|
|||||||
# under the saver's `ORDER BY task_id, idx` sorting.
|
# under the saver's `ORDER BY task_id, idx` sorting.
|
||||||
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
|
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
|
||||||
|
|
||||||
# Delta channels that saw an Overwrite since the last checkpoint. These
|
# Delta channels that must snapshot at the next checkpoint, whatever their
|
||||||
# channels must snapshot after live update applies overwrite semantics so
|
# cadence counters say:
|
||||||
# sparse replay starts from the same post-overwrite value.
|
# * an Overwrite arrived since the last checkpoint, so sparse replay has to
|
||||||
_delta_channels_with_overwrite: set[str]
|
# start from the post-overwrite value;
|
||||||
|
# * the checkpoint this run starts from has pending writes to them; see
|
||||||
|
# `delta_channels_with_pending_writes`.
|
||||||
|
_delta_channels_forced_snapshot: set[str]
|
||||||
|
|
||||||
# The checkpoint_config that points at the parent loaded at `__enter__`
|
# The checkpoint_config that points at the parent loaded at `__enter__`
|
||||||
# (or the synthetic-empty checkpoint, on first run). We capture it
|
# (or the synthetic-empty checkpoint, on first run). We capture it
|
||||||
@@ -683,7 +687,7 @@ class PregelLoop:
|
|||||||
def after_tick(self) -> None:
|
def after_tick(self) -> None:
|
||||||
# finish superstep
|
# finish superstep
|
||||||
writes = [w for t in self.tasks.values() for w in t.writes]
|
writes = [w for t in self.tasks.values() for w in t.writes]
|
||||||
self._delta_channels_with_overwrite.update(
|
self._delta_channels_forced_snapshot.update(
|
||||||
ch
|
ch
|
||||||
for ch, v in writes
|
for ch, v in writes
|
||||||
if isinstance(self.specs.get(ch), DeltaChannel) and _get_overwrite(v)[0]
|
if isinstance(self.specs.get(ch), DeltaChannel) and _get_overwrite(v)[0]
|
||||||
@@ -898,6 +902,15 @@ class PregelLoop:
|
|||||||
self.checkpoint_pending_writes = [
|
self.checkpoint_pending_writes = [
|
||||||
w for w in self.checkpoint_pending_writes if w[1] != RESUME
|
w for w in self.checkpoint_pending_writes if w[1] != RESUME
|
||||||
]
|
]
|
||||||
|
# A resume that is not replaying reuses the head's pending writes
|
||||||
|
# instead of rerunning their tasks, so none of them can leak.
|
||||||
|
self._delta_channels_forced_snapshot = (
|
||||||
|
set()
|
||||||
|
if is_resuming and not self.is_replaying
|
||||||
|
else delta_channels_with_pending_writes(
|
||||||
|
self.specs, self.checkpoint_pending_writes
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
# map command to writes
|
# map command to writes
|
||||||
if input_is_command:
|
if input_is_command:
|
||||||
@@ -991,7 +1004,7 @@ class PregelLoop:
|
|||||||
manager=None,
|
manager=None,
|
||||||
updated_channels=updated_channels,
|
updated_channels=updated_channels,
|
||||||
)
|
)
|
||||||
self._delta_channels_with_overwrite.update(
|
self._delta_channels_forced_snapshot.update(
|
||||||
c
|
c
|
||||||
for c, v in input_writes
|
for c, v in input_writes
|
||||||
if isinstance(self.specs.get(c), DeltaChannel) and _get_overwrite(v)[0]
|
if isinstance(self.specs.get(c), DeltaChannel) and _get_overwrite(v)[0]
|
||||||
@@ -1136,7 +1149,7 @@ class PregelLoop:
|
|||||||
# create new checkpoint
|
# create new checkpoint
|
||||||
channels_to_snapshot = (
|
channels_to_snapshot = (
|
||||||
delta_channels_to_snapshot(self.channels, new_counters)
|
delta_channels_to_snapshot(self.channels, new_counters)
|
||||||
| self._delta_channels_with_overwrite
|
| self._delta_channels_forced_snapshot
|
||||||
if do_checkpoint
|
if do_checkpoint
|
||||||
else set()
|
else set()
|
||||||
)
|
)
|
||||||
@@ -1154,7 +1167,7 @@ class PregelLoop:
|
|||||||
for k in channels_to_snapshot:
|
for k in channels_to_snapshot:
|
||||||
new_counters[k] = (0, 0)
|
new_counters[k] = (0, 0)
|
||||||
if do_checkpoint:
|
if do_checkpoint:
|
||||||
self._delta_channels_with_overwrite.difference_update(channels_to_snapshot)
|
self._delta_channels_forced_snapshot.difference_update(channels_to_snapshot)
|
||||||
non_zero = {k: v for k, v in new_counters.items() if v != (0, 0)}
|
non_zero = {k: v for k, v in new_counters.items() if v != (0, 0)}
|
||||||
if non_zero:
|
if non_zero:
|
||||||
self.checkpoint_metadata["counters_since_delta_snapshot"] = non_zero
|
self.checkpoint_metadata["counters_since_delta_snapshot"] = non_zero
|
||||||
@@ -1239,7 +1252,7 @@ class PregelLoop:
|
|||||||
)
|
)
|
||||||
channels_to_snapshot = (
|
channels_to_snapshot = (
|
||||||
delta_channels_to_snapshot(self.channels, counters)
|
delta_channels_to_snapshot(self.channels, counters)
|
||||||
| self._delta_channels_with_overwrite
|
| self._delta_channels_forced_snapshot
|
||||||
)
|
)
|
||||||
|
|
||||||
pending = [
|
pending = [
|
||||||
@@ -1684,7 +1697,6 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
|||||||
)
|
)
|
||||||
self._delta_write_futs = []
|
self._delta_write_futs = []
|
||||||
self._error_handler_write_futs = []
|
self._error_handler_write_futs = []
|
||||||
self._delta_channels_with_overwrite = set()
|
|
||||||
self._exit_delta_writes = (
|
self._exit_delta_writes = (
|
||||||
[] if self.durability == "exit" and self.checkpointer is not None else None
|
[] if self.durability == "exit" and self.checkpointer is not None else None
|
||||||
)
|
)
|
||||||
@@ -1942,7 +1954,6 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
|||||||
)
|
)
|
||||||
self._delta_write_futs = []
|
self._delta_write_futs = []
|
||||||
self._error_handler_write_futs = []
|
self._error_handler_write_futs = []
|
||||||
self._delta_channels_with_overwrite = set()
|
|
||||||
self._exit_delta_writes = (
|
self._exit_delta_writes = (
|
||||||
[] if self.durability == "exit" and self.checkpointer is not None else None
|
[] if self.durability == "exit" and self.checkpointer is not None else None
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -108,6 +108,7 @@ from langgraph.callbacks import (
|
|||||||
get_sync_graph_callback_manager_for_config,
|
get_sync_graph_callback_manager_for_config,
|
||||||
)
|
)
|
||||||
from langgraph.channels.base import BaseChannel
|
from langgraph.channels.base import BaseChannel
|
||||||
|
from langgraph.channels.delta import DeltaChannel
|
||||||
from langgraph.channels.topic import Topic
|
from langgraph.channels.topic import Topic
|
||||||
from langgraph.config import get_config
|
from langgraph.config import get_config
|
||||||
from langgraph.constants import END
|
from langgraph.constants import END
|
||||||
@@ -133,6 +134,7 @@ from langgraph.pregel._checkpoint import (
|
|||||||
copy_checkpoint,
|
copy_checkpoint,
|
||||||
create_checkpoint,
|
create_checkpoint,
|
||||||
create_checkpoint_plan_for_update_state_api,
|
create_checkpoint_plan_for_update_state_api,
|
||||||
|
delta_channels_with_pending_writes,
|
||||||
empty_checkpoint,
|
empty_checkpoint,
|
||||||
get_updated_channels_from_tasks,
|
get_updated_channels_from_tasks,
|
||||||
)
|
)
|
||||||
@@ -1637,12 +1639,22 @@ class Pregel(
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Subgraph {recast} not found")
|
raise ValueError(f"Subgraph {recast} not found")
|
||||||
|
|
||||||
|
# Taken from the first superstep's base, and cleared by the first
|
||||||
|
# checkpoint that carries the snapshots, which `__copy__` does not write.
|
||||||
|
fork_pending: set[str] | None = None
|
||||||
|
|
||||||
def perform_superstep(
|
def perform_superstep(
|
||||||
input_config: RunnableConfig, updates: Sequence[StateUpdate]
|
input_config: RunnableConfig, updates: Sequence[StateUpdate]
|
||||||
) -> RunnableConfig:
|
) -> RunnableConfig:
|
||||||
|
nonlocal fork_pending
|
||||||
# get last checkpoint
|
# get last checkpoint
|
||||||
config = ensure_config(self.config, input_config)
|
config = ensure_config(self.config, input_config)
|
||||||
saved = checkpointer.get_tuple(config)
|
saved = checkpointer.get_tuple(config)
|
||||||
|
first_superstep = fork_pending is None
|
||||||
|
if fork_pending is None:
|
||||||
|
fork_pending = delta_channels_with_pending_writes(
|
||||||
|
self.channels, saved.pending_writes if saved else None
|
||||||
|
)
|
||||||
if saved is not None:
|
if saved is not None:
|
||||||
self._migrate_checkpoint(saved.checkpoint)
|
self._migrate_checkpoint(saved.checkpoint)
|
||||||
checkpoint = (
|
checkpoint = (
|
||||||
@@ -1726,9 +1738,17 @@ class Pregel(
|
|||||||
self.trigger_to_nodes,
|
self.trigger_to_nodes,
|
||||||
)
|
)
|
||||||
# save checkpoint
|
# save checkpoint
|
||||||
|
next_checkpoint = create_checkpoint(
|
||||||
|
checkpoint,
|
||||||
|
channels,
|
||||||
|
step,
|
||||||
|
get_next_version=checkpointer.get_next_version,
|
||||||
|
channels_to_snapshot=fork_pending,
|
||||||
|
)
|
||||||
|
fork_pending.difference_update(next_checkpoint["channel_values"])
|
||||||
next_config = checkpointer.put(
|
next_config = checkpointer.put(
|
||||||
checkpoint_config,
|
checkpoint_config,
|
||||||
create_checkpoint(checkpoint, channels, step),
|
next_checkpoint,
|
||||||
{
|
{
|
||||||
"source": "update",
|
"source": "update",
|
||||||
"step": step + 1,
|
"step": step + 1,
|
||||||
@@ -1736,7 +1756,7 @@ class Pregel(
|
|||||||
},
|
},
|
||||||
get_new_channel_versions(
|
get_new_channel_versions(
|
||||||
checkpoint_previous_versions,
|
checkpoint_previous_versions,
|
||||||
checkpoint["channel_versions"],
|
next_checkpoint["channel_versions"],
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
return patch_checkpoint_map(
|
return patch_checkpoint_map(
|
||||||
@@ -1765,9 +1785,17 @@ class Pregel(
|
|||||||
if saved and saved.metadata.get("step") is not None
|
if saved and saved.metadata.get("step") is not None
|
||||||
else -1
|
else -1
|
||||||
)
|
)
|
||||||
|
next_checkpoint = create_checkpoint(
|
||||||
|
checkpoint,
|
||||||
|
channels,
|
||||||
|
next_step,
|
||||||
|
get_next_version=checkpointer.get_next_version,
|
||||||
|
channels_to_snapshot=fork_pending,
|
||||||
|
)
|
||||||
|
fork_pending.difference_update(next_checkpoint["channel_values"])
|
||||||
next_config = checkpointer.put(
|
next_config = checkpointer.put(
|
||||||
checkpoint_config,
|
checkpoint_config,
|
||||||
create_checkpoint(checkpoint, channels, next_step),
|
next_checkpoint,
|
||||||
{
|
{
|
||||||
"source": "input",
|
"source": "input",
|
||||||
"step": next_step,
|
"step": next_step,
|
||||||
@@ -1777,7 +1805,7 @@ class Pregel(
|
|||||||
},
|
},
|
||||||
get_new_channel_versions(
|
get_new_channel_versions(
|
||||||
checkpoint_previous_versions,
|
checkpoint_previous_versions,
|
||||||
checkpoint["channel_versions"],
|
next_checkpoint["channel_versions"],
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1998,13 +2026,21 @@ class Pregel(
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
updated_channels = get_updated_channels_from_tasks(run_tasks)
|
updated_channels = get_updated_channels_from_tasks(run_tasks)
|
||||||
if saved is not None:
|
edited_delta_channels = {
|
||||||
for task_id, task in zip(run_task_ids, run_tasks):
|
ch
|
||||||
channel_writes = [w for w in task.writes if w[0] != PUSH]
|
for ch in updated_channels
|
||||||
if channel_writes:
|
if isinstance(self.channels.get(ch), DeltaChannel)
|
||||||
checkpointer.put_writes(
|
}
|
||||||
checkpoint_config, channel_writes, task_id
|
# The base's other children replay whatever is stored on it, so an
|
||||||
)
|
# edit of an older checkpoint snapshots its delta channels here
|
||||||
|
# instead. Later supersteps address the checkpoint just written.
|
||||||
|
if (
|
||||||
|
first_superstep
|
||||||
|
and saved is not None
|
||||||
|
and edited_delta_channels
|
||||||
|
and _is_older_checkpoint(checkpointer, config, saved)
|
||||||
|
):
|
||||||
|
fork_pending.update(edited_delta_channels)
|
||||||
apply_writes(
|
apply_writes(
|
||||||
checkpoint,
|
checkpoint,
|
||||||
channels,
|
channels,
|
||||||
@@ -2020,18 +2056,29 @@ class Pregel(
|
|||||||
parents=saved.metadata.get("parents", {}) if saved else {},
|
parents=saved.metadata.get("parents", {}) if saved else {},
|
||||||
saved_metadata=saved.metadata if saved else None,
|
saved_metadata=saved.metadata if saved else None,
|
||||||
is_fresh_thread=saved is None,
|
is_fresh_thread=saved is None,
|
||||||
|
fork_channels=fork_pending,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
checkpoint = create_checkpoint(
|
checkpoint = create_checkpoint(
|
||||||
checkpoint,
|
checkpoint,
|
||||||
channels,
|
channels,
|
||||||
step + 1,
|
step + 1,
|
||||||
updated_channels=updated_channels if channels_to_snapshot else None,
|
|
||||||
get_next_version=checkpointer.get_next_version
|
get_next_version=checkpointer.get_next_version
|
||||||
if channels_to_snapshot
|
if channels_to_snapshot
|
||||||
else None,
|
else None,
|
||||||
channels_to_snapshot=channels_to_snapshot,
|
channels_to_snapshot=channels_to_snapshot,
|
||||||
)
|
)
|
||||||
|
sealed = fork_pending.intersection(checkpoint["channel_values"])
|
||||||
|
fork_pending.difference_update(checkpoint["channel_values"])
|
||||||
|
if saved is not None:
|
||||||
|
for task_id, task in zip(run_task_ids, run_tasks):
|
||||||
|
channel_writes = [
|
||||||
|
w for w in task.writes if w[0] != PUSH and w[0] not in sealed
|
||||||
|
]
|
||||||
|
if channel_writes:
|
||||||
|
checkpointer.put_writes(
|
||||||
|
checkpoint_config, channel_writes, task_id
|
||||||
|
)
|
||||||
next_config = checkpointer.put(
|
next_config = checkpointer.put(
|
||||||
checkpoint_config,
|
checkpoint_config,
|
||||||
checkpoint,
|
checkpoint,
|
||||||
@@ -2103,12 +2150,22 @@ class Pregel(
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Subgraph {recast} not found")
|
raise ValueError(f"Subgraph {recast} not found")
|
||||||
|
|
||||||
|
# Taken from the first superstep's base, and cleared by the first
|
||||||
|
# checkpoint that carries the snapshots, which `__copy__` does not write.
|
||||||
|
fork_pending: set[str] | None = None
|
||||||
|
|
||||||
async def aperform_superstep(
|
async def aperform_superstep(
|
||||||
input_config: RunnableConfig, updates: Sequence[StateUpdate]
|
input_config: RunnableConfig, updates: Sequence[StateUpdate]
|
||||||
) -> RunnableConfig:
|
) -> RunnableConfig:
|
||||||
|
nonlocal fork_pending
|
||||||
# get last checkpoint
|
# get last checkpoint
|
||||||
config = ensure_config(self.config, input_config)
|
config = ensure_config(self.config, input_config)
|
||||||
saved = await checkpointer.aget_tuple(config)
|
saved = await checkpointer.aget_tuple(config)
|
||||||
|
first_superstep = fork_pending is None
|
||||||
|
if fork_pending is None:
|
||||||
|
fork_pending = delta_channels_with_pending_writes(
|
||||||
|
self.channels, saved.pending_writes if saved else None
|
||||||
|
)
|
||||||
if saved is not None:
|
if saved is not None:
|
||||||
self._migrate_checkpoint(saved.checkpoint)
|
self._migrate_checkpoint(saved.checkpoint)
|
||||||
checkpoint = (
|
checkpoint = (
|
||||||
@@ -2190,16 +2247,25 @@ class Pregel(
|
|||||||
self.trigger_to_nodes,
|
self.trigger_to_nodes,
|
||||||
)
|
)
|
||||||
# save checkpoint
|
# save checkpoint
|
||||||
|
next_checkpoint = create_checkpoint(
|
||||||
|
checkpoint,
|
||||||
|
channels,
|
||||||
|
step,
|
||||||
|
get_next_version=checkpointer.get_next_version,
|
||||||
|
channels_to_snapshot=fork_pending,
|
||||||
|
)
|
||||||
|
fork_pending.difference_update(next_checkpoint["channel_values"])
|
||||||
next_config = await checkpointer.aput(
|
next_config = await checkpointer.aput(
|
||||||
checkpoint_config,
|
checkpoint_config,
|
||||||
create_checkpoint(checkpoint, channels, step),
|
next_checkpoint,
|
||||||
{
|
{
|
||||||
"source": "update",
|
"source": "update",
|
||||||
"step": step + 1,
|
"step": step + 1,
|
||||||
"parents": saved.metadata.get("parents", {}) if saved else {},
|
"parents": saved.metadata.get("parents", {}) if saved else {},
|
||||||
},
|
},
|
||||||
get_new_channel_versions(
|
get_new_channel_versions(
|
||||||
checkpoint_previous_versions, checkpoint["channel_versions"]
|
checkpoint_previous_versions,
|
||||||
|
next_checkpoint["channel_versions"],
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
return patch_checkpoint_map(
|
return patch_checkpoint_map(
|
||||||
@@ -2228,9 +2294,17 @@ class Pregel(
|
|||||||
if saved and saved.metadata.get("step") is not None
|
if saved and saved.metadata.get("step") is not None
|
||||||
else -1
|
else -1
|
||||||
)
|
)
|
||||||
|
next_checkpoint = create_checkpoint(
|
||||||
|
checkpoint,
|
||||||
|
channels,
|
||||||
|
next_step,
|
||||||
|
get_next_version=checkpointer.get_next_version,
|
||||||
|
channels_to_snapshot=fork_pending,
|
||||||
|
)
|
||||||
|
fork_pending.difference_update(next_checkpoint["channel_values"])
|
||||||
next_config = await checkpointer.aput(
|
next_config = await checkpointer.aput(
|
||||||
checkpoint_config,
|
checkpoint_config,
|
||||||
create_checkpoint(checkpoint, channels, next_step),
|
next_checkpoint,
|
||||||
{
|
{
|
||||||
"source": "input",
|
"source": "input",
|
||||||
"step": next_step,
|
"step": next_step,
|
||||||
@@ -2240,7 +2314,7 @@ class Pregel(
|
|||||||
},
|
},
|
||||||
get_new_channel_versions(
|
get_new_channel_versions(
|
||||||
checkpoint_previous_versions,
|
checkpoint_previous_versions,
|
||||||
checkpoint["channel_versions"],
|
next_checkpoint["channel_versions"],
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -2458,13 +2532,21 @@ class Pregel(
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
updated_channels = get_updated_channels_from_tasks(run_tasks)
|
updated_channels = get_updated_channels_from_tasks(run_tasks)
|
||||||
if saved is not None:
|
edited_delta_channels = {
|
||||||
for task_id, task in zip(run_task_ids, run_tasks):
|
ch
|
||||||
channel_writes = [w for w in task.writes if w[0] != PUSH]
|
for ch in updated_channels
|
||||||
if channel_writes:
|
if isinstance(self.channels.get(ch), DeltaChannel)
|
||||||
await checkpointer.aput_writes(
|
}
|
||||||
checkpoint_config, channel_writes, task_id
|
# The base's other children replay whatever is stored on it, so an
|
||||||
)
|
# edit of an older checkpoint snapshots its delta channels here
|
||||||
|
# instead. Later supersteps address the checkpoint just written.
|
||||||
|
if (
|
||||||
|
first_superstep
|
||||||
|
and saved is not None
|
||||||
|
and edited_delta_channels
|
||||||
|
and await _ais_older_checkpoint(checkpointer, config, saved)
|
||||||
|
):
|
||||||
|
fork_pending.update(edited_delta_channels)
|
||||||
apply_writes(
|
apply_writes(
|
||||||
checkpoint,
|
checkpoint,
|
||||||
channels,
|
channels,
|
||||||
@@ -2480,18 +2562,29 @@ class Pregel(
|
|||||||
parents=saved.metadata.get("parents", {}) if saved else {},
|
parents=saved.metadata.get("parents", {}) if saved else {},
|
||||||
saved_metadata=saved.metadata if saved else None,
|
saved_metadata=saved.metadata if saved else None,
|
||||||
is_fresh_thread=saved is None,
|
is_fresh_thread=saved is None,
|
||||||
|
fork_channels=fork_pending,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
checkpoint = create_checkpoint(
|
checkpoint = create_checkpoint(
|
||||||
checkpoint,
|
checkpoint,
|
||||||
channels,
|
channels,
|
||||||
step + 1,
|
step + 1,
|
||||||
updated_channels=updated_channels if channels_to_snapshot else None,
|
|
||||||
get_next_version=checkpointer.get_next_version
|
get_next_version=checkpointer.get_next_version
|
||||||
if channels_to_snapshot
|
if channels_to_snapshot
|
||||||
else None,
|
else None,
|
||||||
channels_to_snapshot=channels_to_snapshot,
|
channels_to_snapshot=channels_to_snapshot,
|
||||||
)
|
)
|
||||||
|
sealed = fork_pending.intersection(checkpoint["channel_values"])
|
||||||
|
fork_pending.difference_update(checkpoint["channel_values"])
|
||||||
|
if saved is not None:
|
||||||
|
for task_id, task in zip(run_task_ids, run_tasks):
|
||||||
|
channel_writes = [
|
||||||
|
w for w in task.writes if w[0] != PUSH and w[0] not in sealed
|
||||||
|
]
|
||||||
|
if channel_writes:
|
||||||
|
await checkpointer.aput_writes(
|
||||||
|
checkpoint_config, channel_writes, task_id
|
||||||
|
)
|
||||||
next_config = await checkpointer.aput(
|
next_config = await checkpointer.aput(
|
||||||
checkpoint_config,
|
checkpoint_config,
|
||||||
checkpoint,
|
checkpoint,
|
||||||
@@ -4181,6 +4274,30 @@ def _trigger_to_nodes(nodes: dict[str, PregelNode]) -> Mapping[str, Sequence[str
|
|||||||
return dict(trigger_to_nodes)
|
return dict(trigger_to_nodes)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_older_checkpoint(
|
||||||
|
checkpointer: BaseCheckpointSaver, config: RunnableConfig, saved: CheckpointTuple
|
||||||
|
) -> bool:
|
||||||
|
"""Whether `config` addressed a checkpoint the thread has moved past."""
|
||||||
|
if not config[CONF].get(CONFIG_KEY_CHECKPOINT_ID):
|
||||||
|
return False
|
||||||
|
latest = checkpointer.get_tuple(
|
||||||
|
patch_configurable(config, {CONFIG_KEY_CHECKPOINT_ID: None})
|
||||||
|
)
|
||||||
|
return latest is not None and latest.checkpoint["id"] != saved.checkpoint["id"]
|
||||||
|
|
||||||
|
|
||||||
|
async def _ais_older_checkpoint(
|
||||||
|
checkpointer: BaseCheckpointSaver, config: RunnableConfig, saved: CheckpointTuple
|
||||||
|
) -> bool:
|
||||||
|
"""Whether `config` addressed a checkpoint the thread has moved past."""
|
||||||
|
if not config[CONF].get(CONFIG_KEY_CHECKPOINT_ID):
|
||||||
|
return False
|
||||||
|
latest = await checkpointer.aget_tuple(
|
||||||
|
patch_configurable(config, {CONFIG_KEY_CHECKPOINT_ID: None})
|
||||||
|
)
|
||||||
|
return latest is not None and latest.checkpoint["id"] != saved.checkpoint["id"]
|
||||||
|
|
||||||
|
|
||||||
def _output(
|
def _output(
|
||||||
stream_mode: StreamMode | Sequence[StreamMode],
|
stream_mode: StreamMode | Sequence[StreamMode],
|
||||||
print_mode: StreamMode | Sequence[StreamMode],
|
print_mode: StreamMode | Sequence[StreamMode],
|
||||||
|
|||||||
@@ -85,11 +85,13 @@ class MemorySaverAssertImmutable(InMemorySaver):
|
|||||||
)
|
)
|
||||||
== saved
|
== saved
|
||||||
), config["configurable"]["checkpoint_ns"]
|
), config["configurable"]["checkpoint_ns"]
|
||||||
|
next_config = super().put(config, checkpoint, metadata, new_versions)
|
||||||
|
# Read back, not the object handed in: a DeltaChannel a step did not
|
||||||
|
# write is refilled on read from the blob its inherited version points at.
|
||||||
self.storage_for_copies[thread_id][checkpoint_ns][checkpoint["id"]] = (
|
self.storage_for_copies[thread_id][checkpoint_ns][checkpoint["id"]] = (
|
||||||
self.serde.dumps_typed(checkpoint)
|
self.serde.dumps_typed(super().get(next_config))
|
||||||
)
|
)
|
||||||
# call super to write checkpoint
|
return next_config
|
||||||
return super().put(config, checkpoint, metadata, new_versions)
|
|
||||||
|
|
||||||
|
|
||||||
class MemorySaverNoPending(InMemorySaver):
|
class MemorySaverNoPending(InMemorySaver):
|
||||||
|
|||||||
@@ -0,0 +1,645 @@
|
|||||||
|
"""Forking a thread must not replay the abandoned branch into the fork.
|
||||||
|
|
||||||
|
Every graph carries a `DeltaChannel` and a plain reducer channel fed the same
|
||||||
|
values; the plain channel needs no replay, so it is the oracle.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
from operator import add
|
||||||
|
from typing import Annotated, Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from langchain_core.runnables import RunnableConfig
|
||||||
|
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||||
|
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||||
|
from typing_extensions import TypedDict
|
||||||
|
|
||||||
|
from langgraph._internal._constants import INPUT
|
||||||
|
from langgraph.channels.delta import DeltaChannel
|
||||||
|
from langgraph.graph import END, START, StateGraph
|
||||||
|
from langgraph.types import Command, Durability, StateSnapshot, StateUpdate, interrupt
|
||||||
|
|
||||||
|
pytestmark = pytest.mark.anyio
|
||||||
|
|
||||||
|
|
||||||
|
def _append(current: list | None, writes: Sequence[Any]) -> list:
|
||||||
|
out = list(current or [])
|
||||||
|
for write in writes:
|
||||||
|
out.extend(write if isinstance(write, list) else [write])
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
class _State(TypedDict):
|
||||||
|
log: Annotated[list, DeltaChannel(_append, snapshot_frequency=1000)]
|
||||||
|
plain: Annotated[list, add]
|
||||||
|
other: Annotated[list, add]
|
||||||
|
|
||||||
|
|
||||||
|
def _build(checkpointer: BaseCheckpointSaver, tag: str) -> Any:
|
||||||
|
def node(state: _State) -> dict:
|
||||||
|
return {"log": [f"{tag}-out"], "plain": [f"{tag}-out"]}
|
||||||
|
|
||||||
|
builder = StateGraph(_State)
|
||||||
|
builder.add_node("n", node)
|
||||||
|
builder.set_entry_point("n")
|
||||||
|
builder.set_finish_point("n")
|
||||||
|
return builder.compile(checkpointer=checkpointer)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_without_delta_writes(checkpointer: BaseCheckpointSaver, tag: str) -> Any:
|
||||||
|
def node(state: _State) -> dict:
|
||||||
|
return {"other": [f"{tag}-other"]}
|
||||||
|
|
||||||
|
builder = StateGraph(_State)
|
||||||
|
builder.add_node("n", node)
|
||||||
|
builder.set_entry_point("n")
|
||||||
|
builder.set_finish_point("n")
|
||||||
|
return builder.compile(checkpointer=checkpointer)
|
||||||
|
|
||||||
|
|
||||||
|
def _thread(thread_id: str) -> RunnableConfig:
|
||||||
|
return {"configurable": {"thread_id": thread_id}}
|
||||||
|
|
||||||
|
|
||||||
|
def _at(config: RunnableConfig, snapshot: StateSnapshot) -> RunnableConfig:
|
||||||
|
return {
|
||||||
|
"configurable": {
|
||||||
|
**config["configurable"],
|
||||||
|
"checkpoint_ns": "",
|
||||||
|
"checkpoint_id": snapshot.config["configurable"]["checkpoint_id"],
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _both(marker: str) -> dict:
|
||||||
|
return {"log": [marker], "plain": [marker]}
|
||||||
|
|
||||||
|
|
||||||
|
def _snapshotted_checkpoints(
|
||||||
|
checkpointer: BaseCheckpointSaver, config: RunnableConfig
|
||||||
|
) -> list[str]:
|
||||||
|
return [
|
||||||
|
tuple_.config["configurable"]["checkpoint_id"]
|
||||||
|
for tuple_ in checkpointer.list(config)
|
||||||
|
if isinstance(tuple_.checkpoint["channel_values"].get("log"), _DeltaSnapshot)
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _assert_fork_is_clean(state: StateSnapshot, abandoned: str) -> None:
|
||||||
|
assert state.values["log"] == state.values["plain"], (
|
||||||
|
f"delta channel diverged from the plain channel: "
|
||||||
|
f"{state.values['log']} != {state.values['plain']}"
|
||||||
|
)
|
||||||
|
assert abandoned not in state.values["log"], (
|
||||||
|
f"{abandoned!r} belongs to the branch the fork replaced, "
|
||||||
|
f"but was replayed into {state.values['log']}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_fork_by_invoke(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
_build(sync_checkpointer, "first").invoke(
|
||||||
|
_both("in-1"), config, durability=durability
|
||||||
|
)
|
||||||
|
graph = _build(sync_checkpointer, "second")
|
||||||
|
graph.invoke(_both("in-2"), config, durability=durability)
|
||||||
|
abandoned_head = graph.get_state(config)
|
||||||
|
|
||||||
|
base = next(
|
||||||
|
snapshot
|
||||||
|
for snapshot in graph.get_state_history(config)
|
||||||
|
if "in-2" not in snapshot.values["log"]
|
||||||
|
)
|
||||||
|
_build(sync_checkpointer, "third").invoke(
|
||||||
|
_both("in-3"), _at(config, base), durability=durability
|
||||||
|
)
|
||||||
|
|
||||||
|
state = graph.get_state(config)
|
||||||
|
_assert_fork_is_clean(state, "in-2")
|
||||||
|
assert state.values["log"] == [*base.values["log"], "in-3", "third-out"]
|
||||||
|
|
||||||
|
abandoned = graph.get_state(abandoned_head.config).values
|
||||||
|
assert abandoned["log"] == abandoned["plain"] == abandoned_head.values["log"]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_afork_by_invoke(
|
||||||
|
async_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
await _build(async_checkpointer, "first").ainvoke(
|
||||||
|
_both("in-1"), config, durability=durability
|
||||||
|
)
|
||||||
|
graph = _build(async_checkpointer, "second")
|
||||||
|
await graph.ainvoke(_both("in-2"), config, durability=durability)
|
||||||
|
abandoned_head = await graph.aget_state(config)
|
||||||
|
|
||||||
|
base = await anext(
|
||||||
|
snapshot
|
||||||
|
async for snapshot in graph.aget_state_history(config)
|
||||||
|
if "in-2" not in snapshot.values["log"]
|
||||||
|
)
|
||||||
|
await _build(async_checkpointer, "third").ainvoke(
|
||||||
|
_both("in-3"), _at(config, base), durability=durability
|
||||||
|
)
|
||||||
|
|
||||||
|
state = await graph.aget_state(config)
|
||||||
|
_assert_fork_is_clean(state, "in-2")
|
||||||
|
assert state.values["log"] == [*base.values["log"], "in-3", "third-out"]
|
||||||
|
|
||||||
|
abandoned = (await graph.aget_state(abandoned_head.config)).values
|
||||||
|
assert abandoned["log"] == abandoned["plain"] == abandoned_head.values["log"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_fork_off_checkpoint_before_first_input(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build(sync_checkpointer, "first")
|
||||||
|
graph.invoke(_both("in-1"), config, durability=durability)
|
||||||
|
|
||||||
|
root = list(graph.get_state_history(config))[-1]
|
||||||
|
assert root.values["log"] == []
|
||||||
|
|
||||||
|
_build(sync_checkpointer, "third").invoke(
|
||||||
|
_both("in-9"), _at(config, root), durability=durability
|
||||||
|
)
|
||||||
|
|
||||||
|
state = graph.get_state(config)
|
||||||
|
_assert_fork_is_clean(state, "in-1")
|
||||||
|
assert state.values["log"] == ["in-9", "third-out"]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_afork_off_checkpoint_before_first_input(
|
||||||
|
async_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build(async_checkpointer, "first")
|
||||||
|
await graph.ainvoke(_both("in-1"), config, durability=durability)
|
||||||
|
|
||||||
|
root = [snapshot async for snapshot in graph.aget_state_history(config)][-1]
|
||||||
|
assert root.values["log"] == []
|
||||||
|
|
||||||
|
await _build(async_checkpointer, "third").ainvoke(
|
||||||
|
_both("in-9"), _at(config, root), durability=durability
|
||||||
|
)
|
||||||
|
|
||||||
|
state = await graph.aget_state(config)
|
||||||
|
_assert_fork_is_clean(state, "in-1")
|
||||||
|
assert state.values["log"] == ["in-9", "third-out"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_fork_by_update_state(sync_checkpointer: BaseCheckpointSaver) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
_build(sync_checkpointer, "first").invoke(_both("in-1"), config)
|
||||||
|
graph = _build(sync_checkpointer, "second")
|
||||||
|
graph.invoke(_both("in-2"), config)
|
||||||
|
|
||||||
|
base = next(
|
||||||
|
snapshot
|
||||||
|
for snapshot in graph.get_state_history(config)
|
||||||
|
if "in-2" not in snapshot.values["log"]
|
||||||
|
)
|
||||||
|
forked = graph.update_state(_at(config, base), _both("patched"))
|
||||||
|
|
||||||
|
state = graph.get_state(forked)
|
||||||
|
_assert_fork_is_clean(state, "in-2")
|
||||||
|
assert state.values["log"] == [*base.values["log"], "patched"]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_afork_by_update_state(
|
||||||
|
async_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
await _build(async_checkpointer, "first").ainvoke(_both("in-1"), config)
|
||||||
|
graph = _build(async_checkpointer, "second")
|
||||||
|
await graph.ainvoke(_both("in-2"), config)
|
||||||
|
|
||||||
|
base = await anext(
|
||||||
|
snapshot
|
||||||
|
async for snapshot in graph.aget_state_history(config)
|
||||||
|
if "in-2" not in snapshot.values["log"]
|
||||||
|
)
|
||||||
|
forked = await graph.aupdate_state(_at(config, base), _both("patched"))
|
||||||
|
|
||||||
|
state = await graph.aget_state(forked)
|
||||||
|
_assert_fork_is_clean(state, "in-2")
|
||||||
|
assert state.values["log"] == [*base.values["log"], "patched"]
|
||||||
|
|
||||||
|
|
||||||
|
def _assert_branch_unchanged(state: StateSnapshot, expected: list, edit: str) -> None:
|
||||||
|
assert state.values["log"] == state.values["plain"] == expected, (
|
||||||
|
f"{edit!r} was written by an update_state on this branch's base, "
|
||||||
|
f"but this branch now reads {state.values['log']}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# The old checkpoint is either a finished turn, which saved no writes, or one
|
||||||
|
# whose next node already ran there, so the edit reuses that task's id.
|
||||||
|
@pytest.mark.parametrize("next_node_ran", [False, True])
|
||||||
|
def test_update_state_on_an_old_checkpoint_leaves_its_other_branch_alone(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver, next_node_ran: bool
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build(sync_checkpointer, "first")
|
||||||
|
graph.invoke(_both("in-1"), config)
|
||||||
|
_build(sync_checkpointer, "second").invoke(_both("in-2"), config)
|
||||||
|
branch = graph.get_state(config)
|
||||||
|
base = next(
|
||||||
|
snapshot
|
||||||
|
for snapshot in graph.get_state_history(config)
|
||||||
|
if "in-2" not in snapshot.values["log"]
|
||||||
|
and snapshot.next == (("n",) if next_node_ran else ())
|
||||||
|
)
|
||||||
|
|
||||||
|
edited = graph.update_state(_at(config, base), _both("edit"), as_node="n")
|
||||||
|
|
||||||
|
_assert_branch_unchanged(
|
||||||
|
graph.get_state(branch.config), branch.values["log"], "edit"
|
||||||
|
)
|
||||||
|
assert graph.get_state(edited).values["log"] == [*base.values["log"], "edit"]
|
||||||
|
|
||||||
|
_build(sync_checkpointer, "third").invoke(_both("in-3"), branch.config)
|
||||||
|
_assert_branch_unchanged(
|
||||||
|
graph.get_state(config),
|
||||||
|
[*branch.values["log"], "in-3", "third-out"],
|
||||||
|
"edit",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("next_node_ran", [False, True])
|
||||||
|
async def test_aupdate_state_on_an_old_checkpoint_leaves_its_other_branch_alone(
|
||||||
|
async_checkpointer: BaseCheckpointSaver, next_node_ran: bool
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build(async_checkpointer, "first")
|
||||||
|
await graph.ainvoke(_both("in-1"), config)
|
||||||
|
await _build(async_checkpointer, "second").ainvoke(_both("in-2"), config)
|
||||||
|
branch = await graph.aget_state(config)
|
||||||
|
base = await anext(
|
||||||
|
snapshot
|
||||||
|
async for snapshot in graph.aget_state_history(config)
|
||||||
|
if "in-2" not in snapshot.values["log"]
|
||||||
|
and snapshot.next == (("n",) if next_node_ran else ())
|
||||||
|
)
|
||||||
|
|
||||||
|
edited = await graph.aupdate_state(_at(config, base), _both("edit"), as_node="n")
|
||||||
|
|
||||||
|
_assert_branch_unchanged(
|
||||||
|
await graph.aget_state(branch.config), branch.values["log"], "edit"
|
||||||
|
)
|
||||||
|
assert (await graph.aget_state(edited)).values["log"] == [
|
||||||
|
*base.values["log"],
|
||||||
|
"edit",
|
||||||
|
]
|
||||||
|
|
||||||
|
await _build(async_checkpointer, "third").ainvoke(_both("in-3"), branch.config)
|
||||||
|
_assert_branch_unchanged(
|
||||||
|
await graph.aget_state(config),
|
||||||
|
[*branch.values["log"], "in-3", "third-out"],
|
||||||
|
"edit",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_bulk_update_on_an_old_checkpoint_leaves_its_other_branch_alone(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build(sync_checkpointer, "first")
|
||||||
|
graph.invoke(_both("in-1"), config)
|
||||||
|
base = graph.get_state(config)
|
||||||
|
_build(sync_checkpointer, "second").invoke(_both("in-2"), config)
|
||||||
|
branch = graph.get_state(config)
|
||||||
|
|
||||||
|
edited = graph.bulk_update_state(
|
||||||
|
_at(config, base),
|
||||||
|
[[StateUpdate(_both("s1"), "n")], [StateUpdate(_both("s2"), "n")]],
|
||||||
|
)
|
||||||
|
|
||||||
|
_assert_branch_unchanged(graph.get_state(branch.config), branch.values["log"], "s1")
|
||||||
|
assert graph.get_state(edited).values["log"] == [*base.values["log"], "s1", "s2"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_update_state_with_the_head_checkpoint_id_stores_no_snapshot(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build(sync_checkpointer, "first")
|
||||||
|
graph.invoke(_both("in-1"), config)
|
||||||
|
for i in range(3):
|
||||||
|
graph.update_state(graph.get_state(config).config, _both(f"u{i}"))
|
||||||
|
|
||||||
|
assert not _snapshotted_checkpoints(sync_checkpointer, config)
|
||||||
|
assert graph.get_state(config).values["log"] == [
|
||||||
|
"in-1",
|
||||||
|
"first-out",
|
||||||
|
"u0",
|
||||||
|
"u1",
|
||||||
|
"u2",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_unaddressed_run_keeps_snapshot_cadence(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build(sync_checkpointer, "first")
|
||||||
|
graph.invoke(_both("in-1"), config, durability=durability)
|
||||||
|
graph.invoke(_both("in-2"), config, durability=durability)
|
||||||
|
|
||||||
|
assert not _snapshotted_checkpoints(sync_checkpointer, config)
|
||||||
|
|
||||||
|
|
||||||
|
def test_fork_before_first_value_when_fork_never_writes_the_channel(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build(sync_checkpointer, "first")
|
||||||
|
graph.invoke(_both("in-1"), config, durability=durability)
|
||||||
|
|
||||||
|
root = list(graph.get_state_history(config))[-1]
|
||||||
|
assert root.values["log"] == []
|
||||||
|
|
||||||
|
_build_without_delta_writes(sync_checkpointer, "third").invoke(
|
||||||
|
{"other": ["in-9"]}, _at(config, root), durability=durability
|
||||||
|
)
|
||||||
|
|
||||||
|
state = graph.get_state(config)
|
||||||
|
_assert_fork_is_clean(state, "in-1")
|
||||||
|
assert state.values["log"] == []
|
||||||
|
|
||||||
|
|
||||||
|
async def test_afork_before_first_value_when_fork_never_writes_the_channel(
|
||||||
|
async_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build(async_checkpointer, "first")
|
||||||
|
await graph.ainvoke(_both("in-1"), config, durability=durability)
|
||||||
|
|
||||||
|
root = [snapshot async for snapshot in graph.aget_state_history(config)][-1]
|
||||||
|
assert root.values["log"] == []
|
||||||
|
|
||||||
|
await _build_without_delta_writes(async_checkpointer, "third").ainvoke(
|
||||||
|
{"other": ["in-9"]}, _at(config, root), durability=durability
|
||||||
|
)
|
||||||
|
|
||||||
|
state = await graph.aget_state(config)
|
||||||
|
_assert_fork_is_clean(state, "in-1")
|
||||||
|
assert state.values["log"] == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_fork_before_first_value_by_bulk_update(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build(sync_checkpointer, "first")
|
||||||
|
graph.invoke(_both("in-1"), config)
|
||||||
|
|
||||||
|
root = list(graph.get_state_history(config))[-1]
|
||||||
|
assert root.values["log"] == []
|
||||||
|
|
||||||
|
forked = graph.bulk_update_state(
|
||||||
|
_at(config, root),
|
||||||
|
[
|
||||||
|
[StateUpdate({"other": ["s1"]}, "n")],
|
||||||
|
[StateUpdate(_both("s2"), "n")],
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
state = graph.get_state(forked)
|
||||||
|
_assert_fork_is_clean(state, "in-1")
|
||||||
|
assert state.values["log"] == ["s2"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("first_as_node", [INPUT, END, "__copy__"])
|
||||||
|
def test_fork_by_bulk_update_whose_first_superstep_skips_the_plan(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver, first_as_node: str
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
_build(sync_checkpointer, "first").invoke(_both("in-1"), config)
|
||||||
|
graph = _build(sync_checkpointer, "second")
|
||||||
|
graph.invoke(_both("in-2"), config)
|
||||||
|
|
||||||
|
base = next(
|
||||||
|
snapshot
|
||||||
|
for snapshot in graph.get_state_history(config)
|
||||||
|
if "in-2" not in snapshot.values["log"]
|
||||||
|
)
|
||||||
|
first = (
|
||||||
|
StateUpdate(_both("first-step"), first_as_node)
|
||||||
|
if first_as_node == INPUT
|
||||||
|
else StateUpdate(None, first_as_node)
|
||||||
|
)
|
||||||
|
forked = graph.bulk_update_state(
|
||||||
|
_at(config, base),
|
||||||
|
[[first], [StateUpdate(_both("second-step"), "n")]],
|
||||||
|
)
|
||||||
|
|
||||||
|
state = graph.get_state(forked)
|
||||||
|
assert state.values["log"] == state.values["plain"], (
|
||||||
|
f"delta channel diverged from the plain channel: "
|
||||||
|
f"{state.values['log']} != {state.values['plain']}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_unaddressed_bulk_update_keeps_snapshot_cadence(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build(sync_checkpointer, "first")
|
||||||
|
graph.invoke(_both("in-1"), config)
|
||||||
|
|
||||||
|
graph.bulk_update_state(
|
||||||
|
config,
|
||||||
|
[[StateUpdate(_both(f"u{i}"), "n")] for i in range(4)],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert not _snapshotted_checkpoints(sync_checkpointer, config)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_paused_before_b(checkpointer: BaseCheckpointSaver) -> Any:
|
||||||
|
builder = StateGraph(_State)
|
||||||
|
builder.add_node("a", lambda state: _both("a"))
|
||||||
|
builder.add_node("b", lambda state: _both("b"))
|
||||||
|
builder.add_edge(START, "a")
|
||||||
|
builder.add_edge("a", "b")
|
||||||
|
builder.add_edge("b", END)
|
||||||
|
return builder.compile(checkpointer=checkpointer, interrupt_before=["b"])
|
||||||
|
|
||||||
|
|
||||||
|
def _build_parallel_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
|
||||||
|
def ask(state: _State) -> dict:
|
||||||
|
interrupt("approve?")
|
||||||
|
return {"other": ["q"]}
|
||||||
|
|
||||||
|
builder = StateGraph(_State)
|
||||||
|
builder.add_node("p", lambda state: _both("p"))
|
||||||
|
builder.add_node("q", ask)
|
||||||
|
builder.add_edge(START, "p")
|
||||||
|
builder.add_edge(START, "q")
|
||||||
|
return builder.compile(checkpointer=checkpointer)
|
||||||
|
|
||||||
|
|
||||||
|
def test_resume_at_interrupt_before_with_the_head_checkpoint_id_runs_the_node(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build_paused_before_b(sync_checkpointer)
|
||||||
|
graph.invoke(_both("in"), config, durability=durability)
|
||||||
|
|
||||||
|
graph.invoke(None, graph.get_state(config).config, durability=durability)
|
||||||
|
|
||||||
|
state = graph.get_state(config)
|
||||||
|
assert state.next == (), f"resume paused again before {state.next}"
|
||||||
|
assert state.values["log"] == state.values["plain"] == ["in", "a", "b"]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_aresume_at_interrupt_before_with_the_head_checkpoint_id_runs_the_node(
|
||||||
|
async_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build_paused_before_b(async_checkpointer)
|
||||||
|
await graph.ainvoke(_both("in"), config, durability=durability)
|
||||||
|
|
||||||
|
await graph.ainvoke(
|
||||||
|
None, (await graph.aget_state(config)).config, durability=durability
|
||||||
|
)
|
||||||
|
|
||||||
|
state = await graph.aget_state(config)
|
||||||
|
assert state.next == (), f"resume paused again before {state.next}"
|
||||||
|
assert state.values["log"] == state.values["plain"] == ["in", "a", "b"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_replay_from_a_paused_checkpoint_runs_the_node_once(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build_paused_before_b(sync_checkpointer)
|
||||||
|
graph.invoke(_both("in"), config)
|
||||||
|
paused = graph.get_state(config).config
|
||||||
|
graph.invoke(None, config)
|
||||||
|
|
||||||
|
graph.invoke(None, paused)
|
||||||
|
|
||||||
|
state = graph.get_state(config)
|
||||||
|
assert state.next == (), f"replay paused again before {state.next}"
|
||||||
|
assert state.values["log"] == state.values["plain"] == ["in", "a", "b"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("addressed", [False, True])
|
||||||
|
def test_new_input_on_an_interrupted_head_does_not_replay_its_pending_writes(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver, durability: Durability, addressed: bool
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build_parallel_interrupt(sync_checkpointer)
|
||||||
|
graph.invoke(_both("in-1"), config, durability=durability)
|
||||||
|
head = graph.get_state(config).config
|
||||||
|
|
||||||
|
graph.invoke(_both("in-2"), head if addressed else config, durability=durability)
|
||||||
|
|
||||||
|
state = graph.get_state(config)
|
||||||
|
assert state.values["log"] == state.values["plain"] == ["in-1", "in-2", "p"]
|
||||||
|
|
||||||
|
|
||||||
|
def _build_deferred_after_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
|
||||||
|
builder = StateGraph(_State)
|
||||||
|
builder.add_node("a", lambda state: _both("a"))
|
||||||
|
builder.add_node("b", lambda state: _both("b"), defer=True)
|
||||||
|
builder.add_node("c", lambda state: {})
|
||||||
|
builder.add_edge(START, "a")
|
||||||
|
builder.add_edge("a", "b")
|
||||||
|
builder.add_edge("a", "c")
|
||||||
|
return builder.compile(checkpointer=checkpointer, interrupt_after=["a"])
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"durability",
|
||||||
|
[
|
||||||
|
"sync",
|
||||||
|
"async",
|
||||||
|
pytest.param(
|
||||||
|
"exit",
|
||||||
|
marks=pytest.mark.xfail(
|
||||||
|
reason="exit durability stores a resumed run's loaded writes twice",
|
||||||
|
strict=True,
|
||||||
|
),
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_resume_on_an_interrupted_head_consumes_its_writes_without_a_snapshot(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build_parallel_interrupt(sync_checkpointer)
|
||||||
|
graph.invoke(_both("in-1"), config, durability=durability)
|
||||||
|
|
||||||
|
graph.invoke(Command(resume="yes"), config, durability=durability)
|
||||||
|
|
||||||
|
state = graph.get_state(config)
|
||||||
|
assert state.next == ()
|
||||||
|
assert state.values["log"] == state.values["plain"] == ["in-1", "p"]
|
||||||
|
assert not _snapshotted_checkpoints(sync_checkpointer, config)
|
||||||
|
|
||||||
|
|
||||||
|
def test_resume_addressed_at_an_interrupted_head_reruns_its_tasks_once(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build_parallel_interrupt(sync_checkpointer)
|
||||||
|
graph.invoke(_both("in-1"), config, durability=durability)
|
||||||
|
|
||||||
|
graph.invoke(
|
||||||
|
Command(resume="yes"), graph.get_state(config).config, durability=durability
|
||||||
|
)
|
||||||
|
|
||||||
|
state = graph.get_state(config)
|
||||||
|
assert state.next == ()
|
||||||
|
assert state.values["log"] == state.values["plain"] == ["in-1", "p"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_update_state_with_the_head_checkpoint_id_keeps_a_deferred_node(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
graph = _build_deferred_after_interrupt(sync_checkpointer)
|
||||||
|
config = _thread("t")
|
||||||
|
graph.invoke(_both("in"), config)
|
||||||
|
|
||||||
|
graph.update_state(graph.get_state(config).config, _both("u"), as_node="c")
|
||||||
|
graph.invoke(None, config)
|
||||||
|
|
||||||
|
state = graph.get_state(config)
|
||||||
|
assert state.next == (), f"deferred node never ran, still pending: {state.next}"
|
||||||
|
assert state.values["log"] == state.values["plain"] == ["in", "a", "u", "b"]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_aupdate_state_with_the_head_checkpoint_id_keeps_a_deferred_node(
|
||||||
|
async_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
graph = _build_deferred_after_interrupt(async_checkpointer)
|
||||||
|
config = _thread("t")
|
||||||
|
await graph.ainvoke(_both("in"), config)
|
||||||
|
|
||||||
|
await graph.aupdate_state(
|
||||||
|
(await graph.aget_state(config)).config, _both("u"), as_node="c"
|
||||||
|
)
|
||||||
|
await graph.ainvoke(None, config)
|
||||||
|
|
||||||
|
state = await graph.aget_state(config)
|
||||||
|
assert state.next == (), f"deferred node never ran, still pending: {state.next}"
|
||||||
|
assert state.values["log"] == state.values["plain"] == ["in", "a", "u", "b"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_turns_addressed_at_the_head_store_no_snapshot(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
config = _thread("t")
|
||||||
|
graph = _build(sync_checkpointer, "turn")
|
||||||
|
graph.invoke(_both("in-1"), config)
|
||||||
|
for turn in range(2, 5):
|
||||||
|
graph.invoke(_both(f"in-{turn}"), graph.get_state(config).config)
|
||||||
|
|
||||||
|
assert not _snapshotted_checkpoints(sync_checkpointer, config)
|
||||||
|
assert (
|
||||||
|
graph.get_state(config).values["log"] == graph.get_state(config).values["plain"]
|
||||||
|
)
|
||||||
@@ -338,3 +338,29 @@ def test_state_history_chain_after_fresh_update_state_delta_channel() -> None:
|
|||||||
assert update_snapshot.metadata["step"] == 0
|
assert update_snapshot.metadata["step"] == 0
|
||||||
assert update_snapshot.parent_config is None
|
assert update_snapshot.parent_config is None
|
||||||
assert [m.content for m in update_snapshot.values["messages"]] == ["hello"]
|
assert [m.content for m in update_snapshot.values["messages"]] == ["hello"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_update_state_that_snapshots_keeps_a_deferred_node_pending() -> None:
|
||||||
|
channel = DeltaChannel(_messages_delta_reducer, snapshot_frequency=1)
|
||||||
|
|
||||||
|
class State(TypedDict):
|
||||||
|
messages: Annotated[list, channel]
|
||||||
|
|
||||||
|
builder = StateGraph(State)
|
||||||
|
builder.add_node("a", lambda state: {"messages": [HumanMessage("a", id="a")]})
|
||||||
|
builder.add_node(
|
||||||
|
"b", lambda state: {"messages": [HumanMessage("b", id="b")]}, defer=True
|
||||||
|
)
|
||||||
|
builder.add_node("c", lambda state: {})
|
||||||
|
builder.add_edge(START, "a")
|
||||||
|
builder.add_edge("a", "b")
|
||||||
|
builder.add_edge("a", "c")
|
||||||
|
graph = builder.compile(checkpointer=InMemorySaver(), interrupt_after=["a"])
|
||||||
|
config = {"configurable": {"thread_id": "t"}}
|
||||||
|
graph.invoke({"messages": [HumanMessage("s", id="s")]}, config)
|
||||||
|
|
||||||
|
graph.update_state(config, {"messages": [HumanMessage("u", id="u")]}, as_node="c")
|
||||||
|
final = graph.invoke(None, config)
|
||||||
|
|
||||||
|
assert [m.content for m in final["messages"]] == ["s", "a", "u", "b"]
|
||||||
|
assert graph.get_state(config).next == ()
|
||||||
|
|||||||
Reference in New Issue
Block a user