Compare commits

...
Author SHA1 Message Date
Elior Nataf Lackritz 22083959f4 fix(langgraph): keep an update_state on an older checkpoint out of its other branches
update_state stores its writes on the checkpoint it addresses. When that
checkpoint already has other children, a DeltaChannel in each of them
replays those writes, so editing an earlier turn leaks the edit into the
branch it forked away from.

When the addressed checkpoint is not the thread's latest, the new
checkpoint snapshots the delta channels the update writes, and those
channels' writes are no longer stored on the addressed checkpoint.
Updates on the latest checkpoint are unchanged.
2026-10-01 12:20:24 -04:00
Elior Nataf Lackritz 858e55f232 refactor(langgraph): call create_checkpoint directly for the fork seal
create_fork_checkpoint only forwarded to create_checkpoint, and an empty
fork set bumps nothing there either, so the four update_state call sites
pass the set straight through. The fork tests also had two identical
helpers; keep one.
2026-09-30 15:07:17 -04:00
Elior Nataf Lackritz e59ecbc23d revert(langgraph): leave exit-mode resume writes as they are on main
Skipping the loaded writes in the exit accumulator removed the duplicate
replay but reordered it: the loaded write keeps its task id while the
run's later writes get step-prefixed ids that sort first. That ordering
problem belongs to exit mode, not to forking, and exists on main without
a fork, so it gets its own fix. The resume test marks exit durability as
an expected failure until then.
2026-09-29 11:39:36 -04:00
Elior Nataf Lackritz 719a4d71bc fix(langgraph): skip the seal on a resume that reuses the head's writes
A resume that is not replaying reuses the head's pending writes instead of
rerunning their tasks, so none of them can leak; deciding in `_first`,
where that is known, stops a plain `Command(resume=...)` from storing a
snapshot. A replaying resume reruns the tasks, so it still seals.

Exit mode also re-recorded the writes a resume loaded from the head: they
were already stored there, and every later read replayed them twice. The
exit accumulator now skips them.
2026-09-29 09:42:55 -04:00
Elior Nataf Lackritz ddaf708cd0 fix(langgraph): snapshot only what a fork can leak, and hide the bump
A fork used to snapshot every DeltaChannel whenever the caller passed a
checkpoint_id. That fired on every turn a client addresses the head
(storing a full copy of the channel per turn), missed new input sent to an
interrupted head without an id, and its storage-only version bump read as
a real write: interrupt_before fired again on resume, and a replay paused
at a node it had already passed.

Snapshot the delta channels the base checkpoint has pending writes for,
since only those can leak into a branch that does not consume them, and
advance versions_seen past any bump that only stores a snapshot, including
the interrupt tracker. update_state no longer records its narrower
updated_channels when it snapshots, so a deferred node listed in next still
runs on resume (#9089).
2026-09-28 19:54:23 -04:00
Elior Nataf Lackritz 9d5b0f1991 refactor(langgraph): drop the deferred fork-snapshot queue
Now that a forced snapshot mints a version for a never-written channel,
the first checkpoint a forked run writes seals every delta channel, so
_delta_channels_awaiting_fork_snapshot never outlived it. Seed
_delta_channels_forced_snapshot directly at loop construction instead.

Also asserts that forking by invoke leaves the abandoned branch's reads
intact, and trims comments and test docstrings.
2026-09-28 19:54:23 -04:00
Elior Nataf Lackritz 9d16b52955 fix(langgraph): seal a fork whose delta channel has no value yet
A DeltaChannel that was never written on the branch being forked has no
value to snapshot and no entry in channel_versions, so create_checkpoint
skipped it and the fork's first checkpoint recorded no boundary at all.
The walk then ran past the fork into the shared base and collected the
abandoned branch's writes, the same failure this branch already fixes for
channels that do have a value.

Two shapes leaked. A run forking off a checkpoint older than the channel's
first value and never writing that channel returned ['in-1'] where the
plain-channel oracle returned []. A bulk update writing the delta key only
in its second superstep returned ['in-1', 's2'] against ['s2'].

No new blob type is needed. _DeltaSnapshot already carries the value and
is already serialized by every saver, and from_checkpoint turns MISSING
into typ(), so _DeltaSnapshot(typ()) reconstructs to the same empty value
the channel would have had. What was missing is a version: without one,
put drops the blob as not-a-new-version, so mint a first one.

Deferring the seal to a later superstep does not work. That superstep
reconstructs through the still-unsealed checkpoint and would only bake the
corrupted value into its own snapshot.

Checked that minting a version does not fire nodes that subscribe to the
channel: a raw Pregel node subscribed directly to the delta channel stays
silent across the fork.

Reported by the Open SWE review bot on #8548.
2026-09-28 19:54:23 -04:00
Elior Nataf Lackritz 2de0c47c1f test(langgraph): compare read-back checkpoints in the immutability saver
MemorySaverAssertImmutable recorded the checkpoint object handed to put,
then compared it against one read back through get. Those two are not the
same shape: channel_values are stored per (channel, version), so a channel
a step did not write is refilled from the blob its inherited version still
points at.

Every channel except DeltaChannel writes its value into channel_values on
every checkpoint, so the two agreed by accident. A DeltaChannel stores
nothing except at a snapshot, so once one snapshots and a later step does
not write it, the saver reports a checkpoint that changed after it was
written when nothing was mutated.

Reproducible on main with no fork involved: a delta channel with
snapshot_frequency=1 written by the first node and left alone by the next
two trips the assertion. Existing delta tests miss it only because they
all use snapshot_frequency=1000.

Record what the saver reads back instead. Comparing read-back against
read-back still catches a checkpoint whose stored data really changed.
2026-09-28 19:54:23 -04:00
Elior Nataf Lackritz 377083220e fix(langgraph): seal a fork on the first checkpoint it writes
The as_node INPUT, END and __copy__ paths write a checkpoint and return
before create_checkpoint_plan_for_update_state_api runs, so a bulk update
whose first superstep took one of them left the branch unsealed. Only
INPUT actually leaked: END absorbs the base's already-run task writes, so
its delta and plain channels agree.

Sealing on a later superstep does not help. By then that superstep has
reconstructed its value by walking through the unsealed checkpoint into
the shared base, so it snapshots an already-corrupted list. The fork's
first checkpoint is the one that has to carry the blob, which is what
create_fork_checkpoint does.

That snapshot was still being dropped by put: these paths apply writes to
the input channel, not the delta channel, so nothing bumped the delta
channel's version and it never entered new_versions. Pass get_next_version
for the manual bump, the same reason exit mode needs it, and derive
new_versions from the returned checkpoint.

fork_pending tracks what is still owed, mirroring
_delta_channels_awaiting_fork_snapshot in _loop.py.

Caught by the Open SWE review bot on #8548.
2026-09-28 19:54:23 -04:00
Elior Nataf Lackritz 9c5914861b fix(langgraph): only fork on the first superstep of a bulk update
perform_superstep returns the config of the checkpoint it just wrote and
bulk_update_state feeds that back in, so from the second superstep on the
incoming config always names a checkpoint whether or not the caller
addressed one. Deriving the fork flag from it made every superstep after
the first force-snapshot every available DeltaChannel and reset its
cadence, storing the whole growing value once per superstep.

Resolve the flag once from the caller's config and pass it explicitly,
true only for the first superstep. The clear-tasks recursion carries it
through, since the checkpoint written there has no delta snapshot and so
leaves a fork unsealed.

Caught by the Open SWE review bot on #8548.
2026-09-28 19:54:23 -04:00
24cf33f348 fix(langgraph): don't replay an abandoned branch into a DeltaChannel fork
Addressing an older checkpoint creates a fork: the shared base ends up
with two children and keeps the checkpoint_writes of the branch the fork
abandons. Nothing records which child consumed which write, so the
DeltaChannel ancestor walk collected the abandoned branch's writes too.
Live execution was correct; only the reconstruction after a reload was
wrong, and it was wrong on every saver.

Fixed on the write side, so no saver changes are needed. A run launched
against an explicitly addressed checkpoint forces every DeltaChannel to
snapshot into its first checkpoint, terminating the walk inside the fork
instead of at the shared base. This mirrors the existing force-snapshot
for Overwrite writes, hence the rename to _delta_channels_forced_snapshot.
update_state against an older checkpoint takes the same path, for the
same reason is_fresh_thread already does.

A channel with no value at the fork base cannot carry a snapshot blob
yet, so the request stays queued until the first superstep that gives it
one. Cost is one snapshot per addressed run, not per superstep.

Fixes #8443

Co-Authored-By: AnnaSuSu <64579968+AnnaSuSu@users.noreply.github.com>
Co-Authored-By: UditDewan <194863456+UditDewan@users.noreply.github.com>
2026-09-28 19:54:23 -04:00
John KennedyGitHubopen-swe[bot] <open-swe@users.noreply.github.com>
07b33185ea fix: reject credential-bearing Git dependencies (#8542)
## Description
Reject Git HTTP dependency URLs containing userinfo before Docker
generation so credentials cannot persist in Dockerfiles or image layers.
Validation now covers local requirement/package metadata and uv
pyproject/lock inputs while keeping errors token-free.

## Test Plan
- [x] Validate credentialed raw, local-manifest, and uv-managed Git URLs
are rejected without echoing secrets
- [x] Validate credential-free HTTPS and SSH Git URLs remain supported

Made by [Open
SWE](https://openswe.vercel.app/agents/81b07455-ece4-3ddc-9955-d7a5bea78d2c)

---------

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
2026-09-27 21:34:53 +00:00
Hugo DURANDandGitHub 7daa3ab49d feat(cli): place self-hosted deployments on a listener (#9056)
Follow-up to #8482. `langgraph deploy --push-to` can now create a
deployment in a workspace that
deploys through a listener in the customer's own cluster, which is the
hybrid case. Before this,
creation in such a workspace was impossible from the CLI: the control
plane rejected it and the CLI
told the user to go and create the deployment in the UI first.

## Changes
- Smart Auto-Placement: The CLI now proactively checks your workspace.
If you only have one listener and one Kubernetes namespace configured
(and are using the managed cloud control plane), it automatically routes
your deployment there. No extra flags needed.
- New Disambiguation Flags: If your workspace has multiple listeners or
namespaces, the CLI will ask you to choose. You can now pass
--listener-id and --k8s-namespace to tell it exactly where to deploy.
- Failing Fast: The CLI now validates your listener and namespace
choices before it starts building and pushing the heavy Docker image. If
you provide an invalid ID, it stops immediately instead of wasting your
time and bandwidth.
- Fixed a Duplication Bug: Previously, if you had many deployments with
similar names, a pagination issue could hide your existing deployment
from the CLI, causing it to accidentally create a duplicate. The CLI now
queries the server for the exact deployment name to guarantee this
doesn't happen.
- Cleaner Errors: Error messages from the control plane are now stripped
of their clunky HTTP envelopes so you get clear, readable sentences when
something goes wrong.

## Testing

Deployment on 3 paths, hybrid, self-hosted, nominal
2026-09-23 13:56:01 -04:00
Sreekara YachamaneniandGitHub e868c3ccfd feat(cli): clarify agent flags and support env defaults (#9063)
Agent deployment options now print a private-beta notice. Rename
`--environment` to `--agent-environment` and accept `LANGSMITH_AGENT_ID`
/ `LANGSMITH_AGENT_ENVIRONMENT` as process-environment defaults for
deploy and list. Explicit flags take precedence, and the backend payload
is unchanged.

Validation: formatting and lint pass. A local smoke check verified
environment-only deployment, explicit flag precedence, list defaults,
and structured JSON output. Full CLI suite: 411 passed; the two known
Docker failures remain (`test_dockerfile_command_with_docker_compose`
and `test_build_generate_proper_build_context`). No new tests added; the
existing test invocation uses the renamed flag.
2026-09-23 17:53:46 +00:00
Mason DaughertyGitHubMason Daughertyopen-swe[bot] <open-swe@users.noreply.github.com>
bdb85b5aa8 chore: remove Claude-specific instructions (#9058)
Remove the root `CLAUDE.md` while retaining the shared `AGENTS.md`
instructions. No Claude-specific GitHub workflows are present, so
existing workflows remain unchanged.

Made by [Open SWE](https://github.com/langchain-ai/open-swe) · [view
thread](https://openswe.vercel.app/agents/541e1bd2-e302-582d-b6cf-bd1df1aadda7)
· openai:gpt-6-astra (low)

Co-authored-by: Mason Daugherty <mdrxy@users.noreply.github.com>
Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
2026-09-23 00:02:12 -04:00
Sreekara YachamaneniandGitHub 1211af45b1 feat(cli): Update langgraph deploy command to use agent_id and environment args (#9055)
- Accept agent_id and environment args for `lanngraph deploy`
  - Validate both arguments present or none
- If agent arguments present, make sure deployment_id and name are not
present
2026-09-22 16:28:36 -04:00
Randall HidajatGitHubHari Dhanushkodiopen-swe[bot] <open-swe@users.noreply.github.com>Hugo Durand
1afaca35a0 feat(cli): add --image-uri flag for self-hosted deployments (#8482)
Adds `--image-uri <uri>` to `langgraph deploy` so self-hosted LangSmith
customers can build, push, and deploy in one command without needing to
script the three steps manually.

When `--image-uri` is provided the CLI:
- Builds the image tagged to the provided URI (same Docker build path as
the local build flow)
- Pushes using whatever Docker credentials are already in the
environment (e.g. WIF, `aws ecr get-login-password`) — no auth handling
in the CLI
- PATCHes the deployment with `source_revision_config.image_uri` (no
`revision_source`, which the self-hosted control plane rejects for
`external_docker` deployments)

Also fixes two self-hosted-specific issues uncovered during testing:
- `LANGSMITH_ENDPOINT` is now used as a fallback when
`LANGGRAPH_HOST_URL` isn't set — the CLI strips the `/api/v1` path and
appends `/api-host` to reach the control plane
- The httpx client now builds full URLs via string concatenation rather
than relying on httpx base_url merging, which silently dropped the
`/api-host` path prefix when paths started with `/`
- The "Check status at" URL after a deploy now correctly points to the
self-hosted LangSmith UI instead of `smith.langchain.com`

**How did you verify your code works?**
Tested end-to-end against a self-hosted LangSmith instance using ECR as
the registry. `langgraph deploy --image-uri <ecr-uri>` successfully
built, pushed, and triggered a deployment revision. Confirmed the
existing SaaS flow (`langgraph deploy` without `--image-uri`) is
unaffected — the new flag is opt-in and the `LANGSMITH_ENDPOINT`
fallback only activates when `LANGGRAPH_HOST_URL` is unset and
`LANGSMITH_ENDPOINT` is present.

---------

Co-authored-by: Hari Dhanushkodi <hari-dhanushkodi@users.noreply.github.com>
Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
Co-authored-by: Hugo Durand <hugo.durand@langchain.dev>
2026-09-22 11:50:02 -04:00
25 changed files with 4807 additions and 753 deletions
-65
View File
@@ -1,65 +0,0 @@
# AGENTS Instructions
This repository is a monorepo. Each library lives in a subdirectory under `libs/`.
<corridor>
## Corridor security analysis
Before generating or modifying code, create a plan and use Corridor's `analyzePlan` tool to analyze it. Apply the resulting security guidance before writing code.
</corridor>
When you modify code in any library, run the following commands in that library's directory before creating a pull request:
- `make format` – run code formatters
- `make lint` – run the linter
- `make test` – execute the test suite
To run a particular test file or to pass additional pytest options you can specify the `TEST` variable:
```
TEST=path/to/test.py make test
```
Other pytest arguments can also be supplied inside the `TEST` variable.
## Libraries
The repository contains several Python and JavaScript/TypeScript libraries.
Below is a high-level overview:
- **checkpoint** – base interfaces for LangGraph checkpointers.
- **checkpoint-postgres** – Postgres implementation of the checkpoint saver.
- **checkpoint-sqlite** – SQLite implementation of the checkpoint saver.
- **cli** – official command-line interface for LangGraph.
- **langgraph** – core framework for building stateful, multi-actor agents.
- **prebuilt** – high-level APIs for creating and running agents and tools.
- **sdk-js** – JS/TS SDK for interacting with the LangGraph REST API.
- **sdk-py** – Python SDK for the LangGraph Server API.
### Dependency map
The diagram below lists downstream libraries for each production dependency as
declared in that library's `pyproject.toml` (or `package.json`).
```text
checkpoint
├── checkpoint-postgres
├── checkpoint-sqlite
├── prebuilt
└── langgraph
prebuilt
└── langgraph
sdk-py
├── langgraph
└── cli
sdk-js (standalone)
```
Changes to a library may impact all of its dependents shown above.
- Do NOT use Sphinx-style double backtick formatting (` ``code`` `). Use single backticks (`` `code` ``) for inline code references in docstrings and comments.
+2
View File
@@ -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.
## Development
+1 -1
View File
@@ -1 +1 @@
__version__ = "0.4.31"
__version__ = "0.4.32"
+87 -3
View File
@@ -6,6 +6,7 @@ import re
import shlex
import textwrap
from collections import Counter
from collections.abc import Iterable
from typing import Literal, NamedTuple
import click
@@ -36,6 +37,10 @@ DISALLOWED_BUILD_COMMAND_CHARS = [
# This blocks background execution (cmd &) while allowing command
# chaining (cmd1 && cmd2) which is common in build commands.
_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(
r"^(?P<major>\d+)"
r"(?:\.(?P<minor>\d+))?"
@@ -78,6 +83,62 @@ def has_disallowed_build_command_content(command: str) -> bool:
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"
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
def validate_config(config: Config) -> Config:
def validate_config(
config: Config, *, source_path: pathlib.Path | None = None
) -> Config:
"""Validate a configuration dictionary."""
graphs = config.get("graphs", {})
@@ -415,6 +478,15 @@ def validate_config(config: Config) -> Config:
' "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_kind = _get_source_kind(config)
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."""
with open(config_path) as 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
# incompatible Node.js version
if validated.get("node_version"):
@@ -1280,6 +1352,7 @@ def python_config_to_docker(
api_version=api_version,
build_tools_to_uninstall=build_tools_to_uninstall,
)
_validate_local_dependency_files(config_path, config)
if pip_installer == "auto":
if _image_supports_uv(base_image):
pip_installer = "uv"
@@ -1490,7 +1563,18 @@ def node_config_to_docker(
) -> tuple[str, dict[str, str]]:
# Calculate paths for monorepo support
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)
if build_context:
File diff suppressed because it is too large Load Diff
+8 -2
View File
@@ -1,12 +1,18 @@
import asyncio
import signal
import sys
from collections.abc import Callable
from collections.abc import Callable, Coroutine
from contextlib import contextmanager
from typing import cast
from typing import Any, Protocol, TypeVar, cast
import click.exceptions
_T = TypeVar("_T")
class CommandRunner(Protocol):
def run(self, coro: Coroutine[Any, Any, _T]) -> _T: ...
@contextmanager
def Runner():
+180 -32
View File
@@ -2,18 +2,125 @@
from __future__ import annotations
from typing import Any
from dataclasses import dataclass
from typing import Any, Literal
from urllib.parse import urlparse
import click
import httpx
CLOUD_CONTROL_PLANE_URL = "https://api.host.langchain.com"
CLOUD_DASHBOARD_URL = "https://smith.langchain.com"
CLOUD_DOMAIN = "langchain.com"
CLOUD_API_HOST = "api.smith.langchain.com"
CLOUD_CONTROL_PLANE_HOST = "api.host.langchain.com"
CLOUD_DASHBOARD_HOST = "smith.langchain.com"
CONTROL_PLANE_PATH = "/api-host"
LANGSMITH_API_PATHS = ("/api/v1", "/api")
LOCAL_HOSTNAMES = ("localhost", "127.0.0.1")
MAX_PAGE_SIZE = 100
SourceName = Literal["internal_docker", "internal_source", "external_docker"]
@dataclass(frozen=True, slots=True)
class ControlPlaneEndpoints:
control_plane_url: str
dashboard_url: str
@classmethod
def resolve(
cls, host_url: str | None, langsmith_endpoint: str | None
) -> ControlPlaneEndpoints:
if host_url:
return cls.from_control_plane_url(host_url)
if langsmith_endpoint:
return cls.from_langsmith_endpoint(langsmith_endpoint)
return cls(CLOUD_CONTROL_PLANE_URL, CLOUD_DASHBOARD_URL)
@property
def is_cloud(self) -> bool:
hostname = urlparse(self.control_plane_url).hostname or ""
return hostname == CLOUD_CONTROL_PLANE_HOST or hostname.endswith(
f".{CLOUD_CONTROL_PLANE_HOST}"
)
@classmethod
def from_control_plane_url(cls, url: str) -> ControlPlaneEndpoints:
control_plane_url = url.rstrip("/")
hostname = urlparse(control_plane_url).hostname or ""
if control_plane_url.endswith(CONTROL_PLANE_PATH):
return cls(control_plane_url, control_plane_url[: -len(CONTROL_PLANE_PATH)])
if hostname in LOCAL_HOSTNAMES:
return cls(control_plane_url, control_plane_url)
return cls(control_plane_url, _cloud_dashboard_for(hostname))
@classmethod
def from_langsmith_endpoint(cls, endpoint: str) -> ControlPlaneEndpoints:
parsed = urlparse(endpoint.rstrip("/"))
hostname = parsed.hostname or ""
if _is_cloud_host(hostname):
return cls.from_control_plane_url(
f"https://{_cloud_control_plane_host_for(hostname)}"
)
root = f"{parsed.scheme}://{parsed.netloc}{_without_api_path(parsed.path)}"
return cls(f"{root}{CONTROL_PLANE_PATH}", root)
def _is_cloud_host(hostname: str) -> bool:
return hostname == CLOUD_DOMAIN or hostname.endswith(f".{CLOUD_DOMAIN}")
def _cloud_control_plane_host_for(langsmith_api_host: str) -> str:
if langsmith_api_host.endswith(f".{CLOUD_API_HOST}"):
region = langsmith_api_host[: -len(CLOUD_API_HOST)]
return f"{region}{CLOUD_CONTROL_PLANE_HOST}"
return CLOUD_CONTROL_PLANE_HOST
def _cloud_dashboard_for(control_plane_host: str) -> str:
if control_plane_host.endswith(f".{CLOUD_CONTROL_PLANE_HOST}"):
region = control_plane_host[: -len(CLOUD_CONTROL_PLANE_HOST) - 1]
return f"https://{region}.{CLOUD_DASHBOARD_HOST}"
return CLOUD_DASHBOARD_URL
def _without_api_path(path: str) -> str:
for api_path in LANGSMITH_API_PATHS:
if path.endswith(api_path):
return path[: -len(api_path)]
return path
def _resources(payload: object) -> list[dict[str, Any]]:
if not isinstance(payload, dict):
return []
resources = payload.get("resources")
if not isinstance(resources, list):
return []
return [item for item in resources if isinstance(item, dict)]
class HostBackendError(click.ClickException):
"""Raised when the host backend returns an error response."""
def __init__(self, message: str, status_code: int | None = None):
def __init__(
self,
message: str,
status_code: int | None = None,
detail: str | None = None,
):
super().__init__(message)
self.status_code = status_code
self.detail = detail
def _error_detail(response: httpx.Response) -> str | None:
try:
body = response.json()
except ValueError:
return None
detail = body.get("detail") if isinstance(body, dict) else None
return detail if isinstance(detail, str) else None
class HostBackendClient:
@@ -24,24 +131,37 @@ class HostBackendClient:
base_url: str,
api_key: str,
tenant_id: str | None = None,
*,
transport: httpx.BaseTransport | None = None,
):
if not base_url:
raise click.UsageError("Host backend URL is required")
transport = httpx.HTTPTransport(retries=3)
headers: dict[str, str] = {
"X-Api-Key": api_key,
"Accept": "application/json",
}
if tenant_id:
headers["X-Tenant-ID"] = tenant_id
self._base_url = base_url.rstrip("/")
self._endpoints = ControlPlaneEndpoints.from_control_plane_url(base_url)
self._base_url = self._endpoints.control_plane_url
self._client = httpx.Client(
base_url=self._base_url,
headers=headers,
transport=transport,
transport=transport or httpx.HTTPTransport(retries=3),
timeout=30,
)
@property
def base_url(self) -> str:
return self._base_url
@property
def endpoints(self) -> ControlPlaneEndpoints:
return self._endpoints
def set_tenant(self, tenant_id: str) -> None:
self._client.headers["X-Tenant-ID"] = tenant_id
def _request(
self,
method: str,
@@ -53,10 +173,12 @@ class HostBackendClient:
resp = self._client.request(method, path, json=payload, params=params)
resp.raise_for_status()
except httpx.HTTPStatusError as err:
detail = err.response.text or str(err.response.status_code)
detail = _error_detail(err.response)
reason = detail or err.response.text or str(err.response.status_code)
raise HostBackendError(
f"{method} {path} failed with status {err.response.status_code}: {detail}",
f"{method} {path} failed with status {err.response.status_code}: {reason}",
status_code=err.response.status_code,
detail=detail,
) from None
except httpx.TransportError as err:
raise HostBackendError(str(err)) from None
@@ -72,30 +194,52 @@ class HostBackendClient:
def create_deployment(
self,
name: str,
deployment_type: str,
source: str,
config_path: str | None = None,
*,
name: str | None,
source: SourceName,
source_config: dict[str, object],
source_revision_config: dict[str, object],
secrets: list[dict[str, str]] | None = None,
agent: dict[str, str] | None = None,
) -> dict[str, Any]:
"""Create a deployment."""
payload: dict[str, Any] = {
"name": name,
"source": source,
"source_config": {"deployment_type": deployment_type},
"source_revision_config": {},
"source_config": source_config,
"source_revision_config": source_revision_config,
}
if source == "internal_source" and config_path:
payload["source_revision_config"]["langgraph_config_path"] = config_path
if agent is not None:
payload["agent"] = agent
else:
payload["name"] = name
if secrets is not None:
payload["secrets"] = secrets
return self._request("POST", "/v2/deployments", payload)
def list_deployments(self, name_contains: str = "") -> dict[str, Any]:
return self._request(
"GET",
"/v2/deployments",
params={"name_contains": name_contains},
def list_deployments(
self,
*,
name: str | None = None,
name_contains: str | None = None,
limit: int | None = None,
agent_id: str | None = None,
agent_environment: str | None = None,
) -> list[dict[str, Any]]:
given = (
("name", name),
("name_contains", name_contains),
("limit", limit),
("agent_id", agent_id),
("agent_environment", agent_environment),
)
params = {key: value for key, value in given if value is not None}
return _resources(self._request("GET", "/v2/deployments", params=params))
def get_listener(self, listener_id: str) -> dict[str, Any]:
return self._request("GET", f"/v2/listeners/{listener_id}")
def list_listeners(self) -> list[dict[str, Any]]:
return _resources(
self._request("GET", "/v2/listeners", params={"limit": MAX_PAGE_SIZE})
)
def get_deployment(self, deployment_id: str) -> dict[str, Any]:
@@ -121,22 +265,21 @@ class HostBackendClient:
self,
deployment_id: str,
image_uri: str,
*,
revision_source: SourceName | None,
secrets: list[dict[str, str]] | None = None,
tracked_packages: list[str] | None = None,
) -> dict[str, Any]:
payload: dict[str, Any] = {
"revision_source": "internal_docker",
"source_revision_config": {"image_uri": image_uri},
}
if revision_source is not None:
payload["revision_source"] = revision_source
if tracked_packages:
payload["tracked_packages"] = tracked_packages
if secrets is not None:
payload["secrets"] = secrets
return self._request(
"PATCH",
f"/v2/deployments/{deployment_id}",
payload,
)
return self._request("PATCH", f"/v2/deployments/{deployment_id}", payload)
def update_deployment_internal_source(
self,
@@ -171,10 +314,15 @@ class HostBackendClient:
payload["secrets"] = secrets
return self._request("PATCH", f"/v2/deployments/{deployment_id}", payload)
def list_revisions(self, deployment_id: str, limit: int = 1) -> dict[str, Any]:
return self._request(
"GET",
f"/v2/deployments/{deployment_id}/revisions?limit={limit}",
def list_revisions(
self, deployment_id: str, limit: int = 1
) -> list[dict[str, Any]]:
return _resources(
self._request(
"GET",
f"/v2/deployments/{deployment_id}/revisions",
params={"limit": limit},
)
)
def get_revision(self, deployment_id: str, revision_id: str) -> dict[str, Any]:
+35
View File
@@ -0,0 +1,35 @@
from __future__ import annotations
from dataclasses import dataclass, replace
DIGEST_SEPARATOR = "@sha256:"
DIGEST_MARKER = "@"
TAG_SEPARATOR = ":"
PATH_SEPARATOR = "/"
@dataclass(frozen=True, slots=True)
class ImageReference:
repository: str
tag: str | None = None
@classmethod
def parse(cls, reference: str) -> ImageReference:
if DIGEST_MARKER in reference:
raise ValueError(f"{reference!r} carries a digest and cannot be tagged")
path_start = reference.rfind(PATH_SEPARATOR) + 1
name, separator, tag = reference[path_start:].partition(TAG_SEPARATOR)
if not separator:
return cls(reference)
return cls(reference[:path_start] + name, tag)
def with_tag(self, tag: str) -> ImageReference:
return replace(self, tag=tag)
def matches_digest(self, repo_digest: str) -> bool:
return repo_digest.startswith(f"{self.repository}{DIGEST_SEPARATOR}")
def __str__(self) -> str:
if self.tag is None:
return self.repository
return f"{self.repository}{TAG_SEPARATOR}{self.tag}"
+5 -1
View File
@@ -650,7 +650,8 @@ class Config(TypedDict, total=False):
pip_config_file: str | None
"""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.
"""
@@ -689,6 +690,9 @@ class Config(TypedDict, total=False):
- "." or "./src" if you have a local Python package
- str (aka "anthropic") for a PyPI 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.
This field is not supported when `source.kind` is `uv`.
+10
View File
@@ -880,6 +880,7 @@ def python_config_to_docker_uv_lock(
_get_node_pm_install_cmd,
_get_pip_cleanup_lines,
_image_supports_uv,
_validate_git_http_url_userinfo_files,
docker_tag,
)
@@ -890,11 +891,20 @@ def python_config_to_docker_uv_lock(
)
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"
_, global_reqs_pip_install, pip_config_file_str = _build_python_install_commands(
config, install_cmd
)
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)
for section, key in [
+2 -2
View File
@@ -28,7 +28,7 @@
"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": {
"anyOf": [
@@ -270,7 +270,7 @@
"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": {
"anyOf": [
+2 -2
View File
@@ -28,7 +28,7 @@
"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": {
"anyOf": [
@@ -270,7 +270,7 @@
"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": {
"anyOf": [
+27 -31
View File
@@ -382,20 +382,18 @@ def test_deploy_list_command(monkeypatch) -> None:
def list_deployments(self, name_contains: str = ""):
captured["name_contains"] = name_contains
return {
"resources": [
{
"id": "dep-123",
"name": "alpha",
"source_config": {"custom_url": "https://alpha.example.com"},
},
{
"id": "dep-456",
"name": "beta",
"source_config": {"custom_url": "https://beta.example.com"},
},
]
}
return [
{
"id": "dep-123",
"name": "alpha",
"source_config": {"custom_url": "https://alpha.example.com"},
},
{
"id": "dep-456",
"name": "beta",
"source_config": {"custom_url": "https://beta.example.com"},
},
]
monkeypatch.setattr(deploy_module, "HostBackendClient", FakeClient)
@@ -435,7 +433,7 @@ def test_deploy_list_command_no_results(monkeypatch) -> None:
pass
def list_deployments(self, name_contains: str = ""):
return {"resources": []}
return []
monkeypatch.setattr(deploy_module, "HostBackendClient", FakeClient)
@@ -468,20 +466,18 @@ def test_deploy_revisions_list_command(monkeypatch) -> None:
def list_revisions(self, deployment_id: str, limit: int = 1):
captured["deployment_id"] = deployment_id
captured["limit"] = str(limit)
return {
"resources": [
{
"id": "rev-123",
"status": "CREATING",
"created_at": "2023-11-07T05:31:56Z",
},
{
"id": "rev-456",
"status": "DEPLOYED",
"created_at": "2023-11-08T10:00:00Z",
},
]
}
return [
{
"id": "rev-123",
"status": "CREATING",
"created_at": "2023-11-07T05:31:56Z",
},
{
"id": "rev-456",
"status": "DEPLOYED",
"created_at": "2023-11-08T10:00:00Z",
},
]
monkeypatch.setattr(deploy_module, "HostBackendClient", FakeClient)
@@ -522,7 +518,7 @@ def test_deploy_revisions_list_command_no_results(monkeypatch) -> None:
pass
def list_revisions(self, deployment_id: str, limit: int = 1):
return {"resources": []}
return []
monkeypatch.setattr(deploy_module, "HostBackendClient", FakeClient)
@@ -555,7 +551,7 @@ def test_deploy_revisions_list_command_with_explicit_limit(monkeypatch) -> None:
def list_revisions(self, deployment_id: str, limit: int = 1):
captured["deployment_id"] = deployment_id
captured["limit"] = str(limit)
return {"resources": []}
return []
monkeypatch.setattr(deploy_module, "HostBackendClient", FakeClient)
File diff suppressed because it is too large Load Diff
+237
View File
@@ -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():
"""Test validation of image_distro field."""
# Valid image_distro values should work
@@ -0,0 +1,119 @@
import json
from unittest.mock import Mock
import httpx
import pytest
from click.testing import CliRunner
import langgraph_cli.deploy as deploy
from langgraph_cli.cli import cli
from langgraph_cli.host_backend import HostBackendClient
@pytest.fixture
def deployment_api(monkeypatch, tmp_path):
monkeypatch.chdir(tmp_path)
monkeypatch.delenv("LANGSMITH_DEPLOYMENT_NAME", raising=False)
monkeypatch.setattr(deploy, "_emitter", None)
monkeypatch.setattr(deploy, "_no_input", False)
(tmp_path / "langgraph.json").write_text(
json.dumps({"dependencies": ["."], "graphs": {"agent": "./agent.py:graph"}})
)
(tmp_path / ".env").write_text("LANGSMITH_DEPLOYMENT_NAME=legacy\n")
requests = []
state = {"enabled": True, "resources": []}
def handler(request):
requests.append(request)
assert request.url.path == "/v2/deployments"
if request.method == "GET":
if not state["enabled"] and (
"agent_id" in request.url.params
or "agent_environment" in request.url.params
):
return httpx.Response(
400, text="Agent filters are not available for this tenant."
)
return httpx.Response(200, json={"resources": state["resources"]})
assert request.method == "POST"
return httpx.Response(200, json={"id": "runtime-id", "name": "server-name"})
client = HostBackendClient("https://api.example.com", "test-key")
client._client.close()
client._client = httpx.Client(
base_url="https://api.example.com",
transport=httpx.MockTransport(handler),
headers={"X-Api-Key": "test-key"},
)
monkeypatch.setattr(deploy, "_create_host_backend_client", lambda *a, **kw: client)
monkeypatch.setattr(deploy, "find_tracked_packages", lambda *a: [])
remote_build = Mock(return_value=deploy.BuildResult())
monkeypatch.setattr(deploy, "_run_remote_build", remote_build)
monkeypatch.setattr(deploy, "_resolve_build_mode", lambda flag, **kw: (flag, None))
yield state, requests, remote_build
client._client.close()
AGENT_ARGS = [
"deploy",
"--agent-id",
"customer-support",
"--agent-environment",
"staging",
"--remote",
"--no-wait",
"--no-input",
]
def test_agent_create(deployment_api, tmp_path, monkeypatch):
monkeypatch.setenv("LANGSMITH_DEPLOYMENT_NAME", "legacy")
_, requests, build = deployment_api
result = CliRunner().invoke(cli, AGENT_ARGS)
assert result.exit_code == 0, result.output
assert dict(requests[0].url.params) == {
"agent_id": "customer-support",
"agent_environment": "staging",
"limit": "100",
}
payload = json.loads(requests[1].content)
assert payload["agent"] == {
"agent_id": "customer-support",
"environment": "staging",
}
assert "name" not in payload
assert build.call_args.kwargs["deployment_id"] == "runtime-id"
assert "server-name" in result.output
assert (tmp_path / ".env").read_text() == "LANGSMITH_DEPLOYMENT_NAME=legacy\n"
def test_agent_update(deployment_api):
state, requests, build = deployment_api
state["resources"] = [{"id": "existing-id", "is_preview": False}]
result = CliRunner().invoke(cli, AGENT_ARGS)
assert result.exit_code == 0, result.output
assert len(requests) == 1
assert build.call_args.kwargs["deployment_id"] == "existing-id"
def test_agent_rejects_explicit_name(deployment_api, monkeypatch):
monkeypatch.setenv("LANGSMITH_DEPLOYMENT_NAME", "legacy")
_, requests, _ = deployment_api
result = CliRunner().invoke(cli, [*AGENT_ARGS, "--name", "legacy"])
assert result.exit_code == 2
assert "cannot be combined" in result.output
assert not requests
def test_agent_lookup_refuses_a_control_plane_that_ignores_the_filter(deployment_api):
state, requests, _ = deployment_api
state["resources"] = [
{"id": "someone-elses", "is_preview": False},
{"id": "another", "is_preview": False},
]
result = CliRunner().invoke(cli, AGENT_ARGS)
assert result.exit_code != 0
assert "does not filter deployments by agent" in result.output
assert len(requests) == 1
+529 -57
View File
@@ -13,6 +13,17 @@ import pytest
import langgraph_cli.deploy as deploy_mod
from langgraph_cli.deploy import (
ById,
ByName,
CustomerRegistrySource,
DockerBuildCommand,
ExistingDeployment,
Listener,
ManagedRegistrySource,
OnListener,
RemoteBuildSource,
RequestedPlacement,
Unplaced,
_call_host_backend_with_optional_tenant,
_create_host_backend_client,
_docker_config_for_token,
@@ -21,12 +32,14 @@ from langgraph_cli.deploy import (
_parse_env_from_config,
_resolve_env_path,
_resolve_pushed_image_digest,
_smith_dashboard_base_url,
_select_source,
_validate_prebuilt_image,
find_deployment_by_name,
normalize_image_tag,
normalize_name,
)
from langgraph_cli.host_backend import HostBackendClient, HostBackendError
from langgraph_cli.image_reference import ImageReference
class TestDockerConfigForToken:
@@ -259,31 +272,29 @@ class TestEnvWithoutDeploymentName:
class TestCallHostBackendWithOptionalTenant:
def _make_client(self, handler):
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
c = HostBackendClient(
"https://api.example.com",
"test-key",
transport=httpx.MockTransport(handler),
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
return c
def _make_eu_client(self, handler):
c = HostBackendClient("https://eu.api.host.langchain.com", "test-key")
c._client = httpx.Client(
base_url="https://eu.api.host.langchain.com",
c = HostBackendClient(
"https://eu.api.host.langchain.com",
"test-key",
transport=httpx.MockTransport(handler),
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
return c
def test_success_passes_through(self):
client = self._make_client(lambda req: httpx.Response(200, json={"ok": True}))
client = self._make_client(
lambda req: httpx.Response(200, json={"resources": [{"id": "dep-1"}]})
)
result = _call_host_backend_with_optional_tenant(
client, lambda c: c.list_deployments()
)
assert result == {"ok": True}
assert result == [{"id": "dep-1"}]
def test_403_not_enabled_gives_actionable_error(self):
detail = (
@@ -334,7 +345,6 @@ class TestCallHostBackendWithOptionalTenant:
assert exc_info.value.status_code == 403
assert "smith.langchain.com" in exc_info.value.message
assert seen_tenant_ids == [None, "workspace-123"]
assert client._client.headers["X-Tenant-ID"] == "workspace-123"
def test_other_403_re_raises_original(self):
client = self._make_client(
@@ -540,60 +550,226 @@ class TestCreateHostBackendClientNoInput:
assert client is not None
class TestSmithDashboardBaseUrl:
def test_none_returns_default(self):
assert _smith_dashboard_base_url(None) == "https://smith.langchain.com"
class TestCreateHostBackendClientEndpoint:
def test_langsmith_endpoint_from_project_env_selects_self_hosted_control_plane(
self, monkeypatch
):
monkeypatch.setenv("LANGSMITH_API_KEY", "lsv2_test")
monkeypatch.delenv("LANGSMITH_ENDPOINT", raising=False)
def test_empty_returns_default(self):
assert _smith_dashboard_base_url("") == "https://smith.langchain.com"
def test_prod_host_url(self):
assert (
_smith_dashboard_base_url("https://api.host.langchain.com")
== "https://smith.langchain.com"
client = _create_host_backend_client(
host_url=None,
api_key=None,
env_vars={"LANGSMITH_ENDPOINT": "https://smith.example.com/api/v1"},
)
def test_dev_host_url(self):
assert (
_smith_dashboard_base_url("https://dev.api.host.langchain.com")
== "https://dev.smith.langchain.com"
assert client.base_url == "https://smith.example.com/api-host"
def test_explicit_host_url_wins_over_langsmith_endpoint(self, monkeypatch):
monkeypatch.setenv("LANGSMITH_API_KEY", "lsv2_test")
monkeypatch.setenv("LANGSMITH_ENDPOINT", "https://smith.example.com/api/v1")
client = _create_host_backend_client(
host_url="https://custom.host.com", api_key=None, env_vars={}
)
def test_eu_host_url(self):
assert (
_smith_dashboard_base_url("https://eu.api.host.langchain.com")
== "https://eu.smith.langchain.com"
assert client.base_url == "https://custom.host.com"
class TestDockerBuildCommand:
@pytest.mark.parametrize(
("machine", "verbose", "expected"),
[
pytest.param(
"x86_64",
False,
DockerBuildCommand(("docker", "build"), ()),
id="amd64_host_builds_natively",
),
pytest.param(
"arm64",
False,
DockerBuildCommand(
("docker", "buildx", "build"),
("--platform", "linux/amd64", "--load", "--progress=quiet"),
),
id="other_hosts_cross_build_quietly",
),
pytest.param(
"arm64",
True,
DockerBuildCommand(
("docker", "buildx", "build"),
("--platform", "linux/amd64", "--load"),
),
id="verbose_cross_build_keeps_progress_output",
),
],
)
def test_for_host_targets_the_deployment_platform(self, machine, verbose, expected):
assert DockerBuildCommand.for_host(machine, verbose=verbose) == expected
class TestSelectSource:
OPTIONS = {
"push_to": None,
"image": None,
"image_name": None,
"tag": None,
"remote_build_flag": None,
"placement": RequestedPlacement(),
"selector": ByName("my-app"),
}
REPOSITORY = "registry.example.com/app"
@pytest.mark.parametrize(
("flags", "docker_available", "expected"),
[
pytest.param(
{"push_to": REPOSITORY},
True,
CustomerRegistrySource(
reference=ImageReference(REPOSITORY, "latest"),
prebuilt_image=None,
requested_placement=RequestedPlacement(),
),
id="push_to_selects_the_external_source_with_the_default_tag",
),
pytest.param(
{"push_to": f"{REPOSITORY}:v2"},
True,
CustomerRegistrySource(
reference=ImageReference(REPOSITORY, "v2"),
prebuilt_image=None,
requested_placement=RequestedPlacement(),
),
id="push_to_keeps_a_tag_given_in_the_reference",
),
pytest.param(
{"push_to": REPOSITORY, "tag": "v3"},
True,
CustomerRegistrySource(
reference=ImageReference(REPOSITORY, "v3"),
prebuilt_image=None,
requested_placement=RequestedPlacement(),
),
id="tag_flag_composes_with_push_to",
),
pytest.param(
{"push_to": REPOSITORY, "image": "app:dev"},
False,
CustomerRegistrySource(
reference=ImageReference(REPOSITORY, "latest"),
prebuilt_image="app:dev",
requested_placement=RequestedPlacement(),
),
id="prebuilt_image_is_retagged_for_push_to_without_docker_checks",
),
pytest.param(
{
"push_to": REPOSITORY,
"placement": RequestedPlacement("listener-1", "agents"),
},
True,
CustomerRegistrySource(
reference=ImageReference(REPOSITORY, "latest"),
prebuilt_image=None,
requested_placement=RequestedPlacement("listener-1", "agents"),
),
id="push_to_carries_the_requested_placement",
),
pytest.param(
{"remote_build_flag": True},
True,
RemoteBuildSource(),
id="remote_flag_selects_the_source_upload",
),
pytest.param(
{},
False,
RemoteBuildSource(),
id="no_local_docker_falls_back_to_the_source_upload",
),
pytest.param(
{},
True,
ManagedRegistrySource(
prebuilt_image=None, image_name=None, tag="latest"
),
id="local_docker_selects_the_internal_docker_source",
),
pytest.param(
{"image": "app:dev", "tag": "v1"},
False,
ManagedRegistrySource(
prebuilt_image="app:dev", image_name=None, tag="v1"
),
id="prebuilt_image_forces_the_internal_docker_source",
),
],
)
def test_flags_select_one_source(
self, monkeypatch, mocker, flags, docker_available, expected
):
mocker.patch(
"langgraph_cli.deploy._get_emitter", return_value=mocker.MagicMock()
)
monkeypatch.setattr(
deploy_mod,
"can_build_locally",
lambda: (True, None) if docker_available else (False, "Docker is required"),
)
def test_staging_host_url(self):
assert (
_smith_dashboard_base_url("https://staging.api.host.langchain.com")
== "https://staging.smith.langchain.com"
assert _select_source(**{**self.OPTIONS, **flags}) == expected
def test_push_to_build_requires_local_docker(self, monkeypatch):
monkeypatch.setattr(
deploy_mod, "can_build_locally", lambda: (False, "Docker is required")
)
def test_localhost(self):
assert (
_smith_dashboard_base_url("http://localhost:8080")
== "http://localhost:8080"
)
with pytest.raises(click.UsageError, match="Docker is required"):
_select_source(**{**self.OPTIONS, "push_to": self.REPOSITORY})
def test_localhost_trailing_slash(self):
assert (
_smith_dashboard_base_url("http://localhost:8080/")
== "http://localhost:8080"
)
@pytest.mark.parametrize(
("flags", "message"),
[
pytest.param(
{"push_to": REPOSITORY, "remote_build_flag": True},
"--push-to cannot be combined with --remote.",
id="push_to_with_remote",
),
pytest.param(
{"push_to": f"{REPOSITORY}:v1", "tag": "v2"},
"already includes a tag",
id="push_to_with_a_tag_and_the_tag_flag",
),
pytest.param(
{"push_to": f"{REPOSITORY}@sha256:abc"},
"not a digest",
id="push_to_with_a_digest",
),
pytest.param(
{"image": "app:dev", "remote_build_flag": True},
"--image cannot be combined with --remote builds.",
id="image_with_remote",
),
pytest.param(
{"placement": RequestedPlacement(listener_id="listener-1")},
"only apply when creating a deployment with --push-to",
id="listener_without_push_to",
),
pytest.param(
{"placement": RequestedPlacement(k8s_namespace="agents")},
"only apply when creating a deployment with --push-to",
id="namespace_without_push_to",
),
],
)
def test_conflicting_flags_are_rejected(self, monkeypatch, flags, message):
monkeypatch.setattr(deploy_mod, "can_build_locally", lambda: (True, None))
def test_127_0_0_1(self):
assert (
_smith_dashboard_base_url("http://127.0.0.1:3000")
== "http://127.0.0.1:3000"
)
def test_unknown_domain_returns_default(self):
assert (
_smith_dashboard_base_url("https://custom.example.com")
== "https://smith.langchain.com"
)
with pytest.raises(click.UsageError, match=message):
_select_source(**{**self.OPTIONS, **flags})
class TestResolvePushedImageDigest:
@@ -644,6 +820,16 @@ class TestResolvePushedImageDigest:
)
assert out == "us-central1-docker.pkg.dev/proj/repo@sha256:abc123"
def test_registry_port_without_tag_still_resolves_the_digest(self):
runner = self._runner('["localhost:5000/repo@sha256:abc123"]')
out = _resolve_pushed_image_digest(
runner,
remote_image="localhost:5000/repo",
docker_config_dir=None,
verbose=False,
)
assert out == "localhost:5000/repo@sha256:abc123"
def test_empty_repodigests_falls_back_with_warning(self, mocker):
emitter = mocker.MagicMock()
mocker.patch("langgraph_cli.deploy._get_emitter", return_value=emitter)
@@ -747,3 +933,289 @@ class TestResolvePushedImageDigest:
frame_locals = captured["coro"].cr_frame.f_locals
assert "--config" not in frame_locals["args"]
captured["coro"].close()
class TestListener:
@pytest.mark.parametrize(
("resource", "expected"),
[
pytest.param(
{
"id": "listener-1",
"compute_id": "prod-cluster",
"compute_config": {"k8s_namespaces": ["agents", "agents-staging"]},
},
Listener("listener-1", "prod-cluster", ("agents", "agents-staging")),
id="reads_id_cluster_and_namespaces",
),
pytest.param(
{"id": "listener-1", "compute_id": "c", "compute_config": {}},
Listener("listener-1", "c", ()),
id="missing_namespaces",
),
pytest.param(
{"id": "listener-1", "compute_id": "c", "compute_config": None},
Listener("listener-1", "c", ()),
id="null_compute_config",
),
pytest.param(
{"id": "listener-1"},
Listener("listener-1", "", ()),
id="only_an_id",
),
],
)
def test_from_resource_reads_the_control_plane_shape(self, resource, expected):
assert Listener.from_resource(resource) == expected
ONE_NAMESPACE = Listener("listener-1", "prod-cluster", ("agents",))
TWO_NAMESPACES = Listener("listener-2", "multi-cluster", ("agents", "agents-staging"))
NO_NAMESPACE = Listener("listener-3", "broken-cluster", ())
class TestRequestedPlacement:
@pytest.mark.parametrize(
("request_", "listeners", "expected"),
[
pytest.param(
RequestedPlacement(), (), Unplaced(), id="no_listeners_no_request"
),
pytest.param(
RequestedPlacement(),
(ONE_NAMESPACE,),
OnListener("listener-1", "agents"),
id="uses_the_only_possible_answer",
),
pytest.param(
RequestedPlacement(k8s_namespace="agents-staging"),
(TWO_NAMESPACES,),
OnListener("listener-2", "agents-staging"),
id="namespace_alone_picks_the_only_listener",
),
],
)
def test_resolves_to_a_placement(self, request_, listeners, expected):
assert request_.among(listeners) == expected
@pytest.mark.parametrize(
("request_", "listeners", "message"),
[
pytest.param(
RequestedPlacement(listener_id="listener-1"),
(),
"no listeners",
id="workspace_has_no_listeners",
),
pytest.param(
RequestedPlacement(),
(ONE_NAMESPACE, TWO_NAMESPACES),
"--listener-id",
id="several_listeners_need_a_choice",
),
pytest.param(
RequestedPlacement(k8s_namespace="agents"),
(ONE_NAMESPACE, TWO_NAMESPACES),
"--listener-id",
id="namespace_alone_is_ambiguous_with_several_listeners",
),
pytest.param(
RequestedPlacement(k8s_namespace="agents"),
(),
"no listeners",
id="namespace_without_any_listener",
),
pytest.param(
RequestedPlacement(),
(TWO_NAMESPACES,),
"--k8s-namespace",
id="several_namespaces_need_a_choice",
),
],
)
def test_refuses_and_names_the_choices(self, request_, listeners, message):
with pytest.raises(click.UsageError, match=message):
request_.among(listeners)
def test_the_error_lists_every_listener_with_its_cluster_and_namespaces(self):
with pytest.raises(click.UsageError) as error:
RequestedPlacement().among((ONE_NAMESPACE, TWO_NAMESPACES))
assert "listener-1" in error.value.message
assert "prod-cluster" in error.value.message
assert "agents-staging" in error.value.message
@pytest.mark.parametrize(
("placement", "expected"),
[
pytest.param(Unplaced(), {}, id="unplaced_adds_nothing"),
pytest.param(
OnListener("listener-1", "agents"),
{
"listener_id": "listener-1",
"listener_config": {"k8s_namespace": "agents"},
},
id="placed_carries_listener_and_namespace",
),
],
)
def test_source_config_matches_the_control_plane_shape(self, placement, expected):
assert placement.source_config() == expected
def test_finding_a_deployment_by_name_narrows_the_search_for_every_server_version():
seen: dict = {}
def handler(req: httpx.Request) -> httpx.Response:
seen["params"] = dict(req.url.params)
return httpx.Response(
200,
json={"resources": [{"id": "dep-1", "name": "agent", "source": "github"}]},
)
client = HostBackendClient(
"https://api.example.com", "key", transport=httpx.MockTransport(handler)
)
found = find_deployment_by_name(client, "agent")
assert seen["params"] == {
"name": "agent",
"name_contains": "agent",
"limit": "100",
}
assert found == ExistingDeployment("dep-1", "github")
def test_a_server_that_ignores_the_exact_name_filter_never_matches_another_deployment():
client = HostBackendClient(
"https://api.example.com",
"key",
transport=httpx.MockTransport(
lambda req: httpx.Response(
200,
json={
"resources": [
{
"id": "dep-other",
"name": "another-teams-agent",
"source": "external_docker",
}
]
},
)
),
)
assert find_deployment_by_name(client, "brand-new-agent") is None
def test_a_full_page_without_a_match_refuses_to_claim_the_name_is_free():
page = [
{"id": f"dep-{index}", "name": f"other-agent-{index}"} for index in range(100)
]
client = HostBackendClient(
"https://api.example.com",
"key",
transport=httpx.MockTransport(
lambda req: httpx.Response(200, json={"resources": page})
),
)
with pytest.raises(click.ClickException, match="--deployment-id"):
find_deployment_by_name(client, "brand-new-agent")
def test_a_partial_page_without_a_match_means_the_name_is_free():
client = HostBackendClient(
"https://api.example.com",
"key",
transport=httpx.MockTransport(
lambda req: httpx.Response(
200, json={"resources": [{"id": "dep-1", "name": "other"}]}
)
),
)
assert find_deployment_by_name(client, "brand-new-agent") is None
@pytest.mark.parametrize(
"resource",
[
pytest.param({"compute_id": "c"}, id="no_id"),
pytest.param({"id": ""}, id="empty_id"),
],
)
def test_a_listener_without_an_id_is_refused(resource):
with pytest.raises(HostBackendError, match="without an id"):
Listener.from_resource(resource)
def test_a_deployment_id_with_listener_flags_is_refused_without_probing_docker(
monkeypatch,
):
def explode() -> tuple[bool, str | None]:
raise AssertionError("docker must not be probed for an argv-only conflict")
monkeypatch.setattr(deploy_mod, "can_build_locally", explode)
with pytest.raises(click.UsageError, match="--deployment-id"):
_select_source(
push_to="registry.example.com/app",
image=None,
image_name=None,
tag=None,
remote_build_flag=None,
placement=RequestedPlacement(listener_id="listener-1"),
selector=ById("dep-1"),
)
class TestPlacementOnAKnownListener:
@pytest.mark.parametrize(
("request_", "listener", "expected"),
[
pytest.param(
RequestedPlacement(listener_id="listener-1"),
ONE_NAMESPACE,
OnListener("listener-1", "agents"),
id="the_only_namespace_is_used",
),
pytest.param(
RequestedPlacement(listener_id="listener-2", k8s_namespace="agents"),
TWO_NAMESPACES,
OnListener("listener-2", "agents"),
id="the_chosen_namespace_is_used",
),
],
)
def test_places_on_the_listener(self, request_, listener, expected):
assert request_.on(listener) == expected
@pytest.mark.parametrize(
("request_", "listener", "message"),
[
pytest.param(
RequestedPlacement(listener_id="listener-2"),
TWO_NAMESPACES,
"--k8s-namespace",
id="several_namespaces_need_a_choice",
),
pytest.param(
RequestedPlacement(listener_id="listener-2", k8s_namespace="nope"),
TWO_NAMESPACES,
"does not serve namespace",
id="unknown_namespace",
),
pytest.param(
RequestedPlacement(listener_id="listener-3"),
NO_NAMESPACE,
"serves no namespaces",
id="listener_without_namespaces",
),
],
)
def test_refuses_and_names_the_namespaces(self, request_, listener, message):
with pytest.raises(click.UsageError, match=message):
request_.on(listener)
+535 -148
View File
@@ -3,29 +3,16 @@ import json
import httpx
import pytest
from langgraph_cli.host_backend import HostBackendClient, HostBackendError
@pytest.fixture
def mock_transport():
return httpx.MockTransport(lambda req: httpx.Response(200, json={"ok": True}))
@pytest.fixture
def client(mock_transport):
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=mock_transport,
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
return c
from langgraph_cli.host_backend import (
ControlPlaneEndpoints,
HostBackendClient,
HostBackendError,
)
def test_constructor_strips_trailing_slash():
c = HostBackendClient("https://api.example.com/", "key")
assert str(c._client.base_url) == "https://api.example.com"
assert c.base_url == "https://api.example.com"
def test_constructor_empty_url_raises():
@@ -39,12 +26,8 @@ def test_request_sends_headers():
assert req.headers["accept"] == "application/json"
return httpx.Response(200, json={"ok": True})
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=httpx.MockTransport(handler),
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
c = HostBackendClient(
"https://api.example.com", "test-key", transport=httpx.MockTransport(handler)
)
result = c._request("GET", "/test")
assert result == {"ok": True}
@@ -56,12 +39,8 @@ def test_request_sends_json_payload():
assert req.content == b'{"key":"value"}'
return httpx.Response(200, json={"created": True})
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=httpx.MockTransport(handler),
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
c = HostBackendClient(
"https://api.example.com", "test-key", transport=httpx.MockTransport(handler)
)
result = c._request("POST", "/test", {"key": "value"})
assert result == {"created": True}
@@ -69,25 +48,13 @@ def test_request_sends_json_payload():
def test_request_empty_body_returns_none():
transport = httpx.MockTransport(lambda req: httpx.Response(200, content=b""))
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=transport,
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
c = HostBackendClient("https://api.example.com", "test-key", transport=transport)
assert c._request("DELETE", "/test") is None
def test_request_http_error_raises():
transport = httpx.MockTransport(lambda req: httpx.Response(404, text="not found"))
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=transport,
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
c = HostBackendClient("https://api.example.com", "test-key", transport=transport)
with pytest.raises(HostBackendError, match="404"):
c._request("GET", "/missing")
@@ -96,13 +63,7 @@ def test_request_invalid_json_raises():
transport = httpx.MockTransport(
lambda req: httpx.Response(200, content=b"not json")
)
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=transport,
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
c = HostBackendClient("https://api.example.com", "test-key", transport=transport)
with pytest.raises(HostBackendError, match="Failed to decode"):
c._request("GET", "/bad-json")
@@ -111,84 +72,20 @@ def test_request_transport_error_raises():
def handler(req: httpx.Request) -> httpx.Response:
raise httpx.ConnectError("connection refused")
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=httpx.MockTransport(handler),
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
c = HostBackendClient(
"https://api.example.com", "test-key", transport=httpx.MockTransport(handler)
)
with pytest.raises(HostBackendError, match="connection refused"):
c._request("GET", "/test")
def test_create_deployment(client):
result = client.create_deployment(
name="my-deploy", deployment_type="dev", source="internal_docker"
)
assert result == {"ok": True}
def test_get_deployment(client):
result = client.get_deployment("dep-123")
assert result == {"ok": True}
def test_list_deployments(client):
result = client.list_deployments("my-app")
assert result == {"ok": True}
def test_list_deployments_sends_query_params():
def handler(req: httpx.Request) -> httpx.Response:
assert req.url.path == "/v2/deployments"
assert req.url.params["name_contains"] == "my app"
return httpx.Response(200, json={"ok": True})
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=httpx.MockTransport(handler),
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
result = c.list_deployments("my app")
assert result == {"ok": True}
def test_delete_deployment(client):
result = client.delete_deployment("dep-123")
assert result == {"ok": True}
def test_request_push_token(client):
result = client.request_push_token("dep-123")
assert result == {"ok": True}
def test_update_deployment(client):
result = client.update_deployment(
"dep-123", "image:latest", secrets=[{"name": "KEY", "value": "val"}]
)
assert result == {"ok": True}
def test_update_deployment_no_secrets(client):
result = client.update_deployment("dep-123", "image:latest")
assert result == {"ok": True}
def _capturing_client(captured: dict) -> HostBackendClient:
def handler(req: httpx.Request) -> httpx.Response:
captured["body"] = req.read()
return httpx.Response(200, json={"ok": True})
c = HostBackendClient("https://api.example.com", "key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=httpx.MockTransport(handler),
headers={"X-Api-Key": "key", "Accept": "application/json"},
timeout=30,
c = HostBackendClient(
"https://api.example.com", "key", transport=httpx.MockTransport(handler)
)
return c
@@ -199,6 +96,7 @@ def test_update_deployment_forwards_tracked_packages():
c.update_deployment(
"dep-123",
"image:latest",
revision_source="internal_docker",
tracked_packages=["google-adk:1.0.0"],
)
body = json.loads(captured["body"])
@@ -209,7 +107,7 @@ def test_update_deployment_forwards_tracked_packages():
def test_update_deployment_omits_tracked_packages_when_absent():
captured: dict = {}
c = _capturing_client(captured)
c.update_deployment("dep-123", "image:latest")
c.update_deployment("dep-123", "image:latest", revision_source="internal_docker")
body = json.loads(captured["body"])
assert "tracked_packages" not in body
@@ -241,33 +139,14 @@ def test_update_deployment_internal_source_omits_tracked_packages_when_absent():
assert "tracked_packages" not in body
def test_list_revisions(client):
result = client.list_revisions("dep-123", limit=5)
assert result == {"ok": True}
def test_get_revision(client):
result = client.get_revision("dep-123", "rev-456")
assert result == {"ok": True}
def test_get_build_logs(client):
result = client.get_build_logs("proj-1", "rev-1", {"limit": 10})
assert result == {"ok": True}
def test_get_deploy_logs_all_revisions():
def handler(req: httpx.Request) -> httpx.Response:
assert "/v1/projects/proj-1/deploy_logs" in str(req.url)
assert "/revisions/" not in str(req.url)
return httpx.Response(200, json={"logs": [{"message": "running"}]})
c = HostBackendClient("https://api.example.com", "key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=httpx.MockTransport(handler),
headers={"X-Api-Key": "key", "Accept": "application/json"},
timeout=30,
c = HostBackendClient(
"https://api.example.com", "key", transport=httpx.MockTransport(handler)
)
result = c.get_deploy_logs("proj-1", {"limit": 10})
assert result == {"logs": [{"message": "running"}]}
@@ -278,12 +157,520 @@ def test_get_deploy_logs_specific_revision():
assert "/v1/projects/proj-1/revisions/rev-2/deploy_logs" in str(req.url)
return httpx.Response(200, json={"logs": []})
c = HostBackendClient("https://api.example.com", "key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=httpx.MockTransport(handler),
headers={"X-Api-Key": "key", "Accept": "application/json"},
timeout=30,
c = HostBackendClient(
"https://api.example.com", "key", transport=httpx.MockTransport(handler)
)
result = c.get_deploy_logs("proj-1", {"limit": 10}, revision_id="rev-2")
assert result == {"logs": []}
def _routing_client(seen: dict) -> HostBackendClient:
def handler(req: httpx.Request) -> httpx.Response:
seen["method"] = req.method
seen["url"] = str(req.url)
return httpx.Response(200, json={"ok": True})
c = HostBackendClient(
"https://api.example.com/prefix", "key", transport=httpx.MockTransport(handler)
)
return c
@pytest.mark.parametrize(
("call", "expected_body"),
[
pytest.param(
lambda c: c.create_deployment(
name="my-deploy",
source="internal_docker",
source_config={"deployment_type": "dev"},
source_revision_config={},
),
{
"name": "my-deploy",
"source": "internal_docker",
"source_config": {"deployment_type": "dev"},
"source_revision_config": {},
},
id="internal_docker_create_omits_secrets_key_when_not_given",
),
pytest.param(
lambda c: c.create_deployment(
name="my-deploy",
source="internal_docker",
source_config={"deployment_type": "prod"},
source_revision_config={},
secrets=[{"name": "KEY", "value": "val"}],
),
{
"name": "my-deploy",
"source": "internal_docker",
"source_config": {"deployment_type": "prod"},
"source_revision_config": {},
"secrets": [{"name": "KEY", "value": "val"}],
},
id="internal_docker_create_forwards_secrets",
),
pytest.param(
lambda c: c.update_deployment(
"dep-123",
"registry.example.com/app@sha256:abc",
revision_source="internal_docker",
secrets=[{"name": "KEY", "value": "val"}],
),
{
"revision_source": "internal_docker",
"source_revision_config": {
"image_uri": "registry.example.com/app@sha256:abc"
},
"secrets": [{"name": "KEY", "value": "val"}],
},
id="internal_docker_revision_names_its_source",
),
pytest.param(
lambda c: c.update_deployment_internal_source(
"dep-123",
source_tarball_path="tarballs/src.tgz",
config_path="langgraph.json",
secrets=[],
install_command="yarn install",
build_command="yarn build",
),
{
"revision_source": "internal_source",
"source_revision_config": {
"source_tarball_path": "tarballs/src.tgz",
"langgraph_config_path": "langgraph.json",
},
"source_config": {
"install_command": "yarn install",
"build_command": "yarn build",
},
"secrets": [],
},
id="internal_source_revision_sends_js_build_commands",
),
pytest.param(
lambda c: c.update_deployment_internal_source(
"dep-123",
source_tarball_path="tarballs/src.tgz",
config_path="langgraph.json",
),
{
"revision_source": "internal_source",
"source_revision_config": {
"source_tarball_path": "tarballs/src.tgz",
"langgraph_config_path": "langgraph.json",
},
},
id="internal_source_revision_omits_source_config_without_commands",
),
pytest.param(
lambda c: c.create_deployment(
name="agent",
source="external_docker",
source_config={"resource_spec": {}},
source_revision_config={
"image_uri": "registry.example.com/agent@sha256:1"
},
secrets=[],
),
{
"name": "agent",
"source": "external_docker",
"source_config": {"resource_spec": {}},
"source_revision_config": {
"image_uri": "registry.example.com/agent@sha256:1"
},
"secrets": [],
},
id="create_sends_the_source_configs_as_given",
),
pytest.param(
lambda c: c.update_deployment(
"dep-1", "registry.example.com/agent@sha256:2", revision_source=None
),
{
"source_revision_config": {
"image_uri": "registry.example.com/agent@sha256:2"
}
},
id="revision_without_source_override_omits_revision_source",
),
pytest.param(
lambda c: c.update_deployment(
"dep-1",
"registry.example.com/agent@sha256:2",
revision_source="internal_docker",
tracked_packages=["langgraph:1.0.0"],
),
{
"revision_source": "internal_docker",
"source_revision_config": {
"image_uri": "registry.example.com/agent@sha256:2"
},
"tracked_packages": ["langgraph:1.0.0"],
},
id="revision_with_source_override_names_it",
),
],
)
def test_request_body_matches_control_plane_contract(call, expected_body):
captured: dict = {}
call(_capturing_client(captured))
assert json.loads(captured["body"]) == expected_body
@pytest.mark.parametrize(
("call", "method", "route"),
[
pytest.param(
lambda c: c.create_deployment(
name="n",
source="internal_docker",
source_config={"deployment_type": "dev"},
source_revision_config={},
),
"POST",
"/v2/deployments",
id="create_deployment",
),
pytest.param(
lambda c: c.get_deployment("dep-1"),
"GET",
"/v2/deployments/dep-1",
id="get_deployment",
),
pytest.param(
lambda c: c.delete_deployment("dep-1"),
"DELETE",
"/v2/deployments/dep-1",
id="delete_deployment",
),
pytest.param(
lambda c: c.update_deployment("dep-1", "img", revision_source=None),
"PATCH",
"/v2/deployments/dep-1",
id="patch_deployment",
),
pytest.param(
lambda c: c.request_push_token("dep-1"),
"POST",
"/v2/deployments/dep-1/push-token",
id="push_token",
),
pytest.param(
lambda c: c.request_upload_url("dep-1"),
"POST",
"/v2/deployments/dep-1/upload-url",
id="upload_url",
),
pytest.param(
lambda c: c.list_revisions("dep-1", limit=5),
"GET",
"/v2/deployments/dep-1/revisions?limit=5",
id="list_revisions_puts_limit_in_query",
),
pytest.param(
lambda c: c.get_revision("dep-1", "rev-2"),
"GET",
"/v2/deployments/dep-1/revisions/rev-2",
id="get_revision",
),
pytest.param(
lambda c: c.get_build_logs("dep-1", "rev-2", {"limit": 10}),
"POST",
"/v1/projects/dep-1/revisions/rev-2/build_logs",
id="build_logs",
),
],
)
def test_request_targets_control_plane_route_under_base_url(call, method, route):
seen: dict = {}
call(_routing_client(seen))
assert (seen["method"], seen["url"]) == (
method,
f"https://api.example.com/prefix{route}",
)
def test_injected_transport_receives_requests_under_the_prefixed_base_url():
seen: dict = {}
def handler(req: httpx.Request) -> httpx.Response:
seen["url"] = str(req.url)
seen["api_key"] = req.headers["x-api-key"]
return httpx.Response(200, json={"ok": True})
c = HostBackendClient(
"https://smith.example.com/api-host",
"key",
transport=httpx.MockTransport(handler),
)
assert c.list_revisions("dep-1", limit=2) == []
assert seen == {
"url": "https://smith.example.com/api-host/v2/deployments/dep-1/revisions?limit=2",
"api_key": "key",
}
CLOUD = ("https://api.host.langchain.com", "https://smith.langchain.com")
@pytest.mark.parametrize(
("host_url", "langsmith_endpoint", "expected"),
[
pytest.param(None, None, CLOUD, id="nothing_configured_targets_cloud"),
pytest.param(
None, "https://api.smith.langchain.com", CLOUD, id="cloud_langsmith_api"
),
pytest.param(
None,
"https://api.smith.langchain.com/api/v1",
CLOUD,
id="cloud_langsmith_api_with_versioned_path",
),
pytest.param(
None, "https://api.langchain.com", CLOUD, id="cloud_langchain_api_alias"
),
pytest.param(
None,
"https://xapi.smith.langchain.com",
CLOUD,
id="lookalike_cloud_host_is_not_rewritten_into_a_control_plane",
),
pytest.param(
None,
"https://eu.api.smith.langchain.com",
("https://eu.api.host.langchain.com", "https://eu.smith.langchain.com"),
id="eu_cloud_maps_to_eu_control_plane",
),
pytest.param(
None,
"https://dev.api.smith.langchain.com",
("https://dev.api.host.langchain.com", "https://dev.smith.langchain.com"),
id="dev_cloud_maps_to_dev_control_plane",
),
pytest.param(
None,
"https://aks.smith.langchain.dev/api",
(
"https://aks.smith.langchain.dev/api-host",
"https://aks.smith.langchain.dev",
),
id="self_hosted_api_path_becomes_api_host",
),
pytest.param(
None,
"https://smith.example.com/api/v1",
("https://smith.example.com/api-host", "https://smith.example.com"),
id="self_hosted_versioned_api_path_becomes_api_host",
),
pytest.param(
None,
"https://smith.example.com",
("https://smith.example.com/api-host", "https://smith.example.com"),
id="self_hosted_origin_gets_api_host_appended",
),
pytest.param(
None,
"https://corp.example.com/langsmith/api/v1",
(
"https://corp.example.com/langsmith/api-host",
"https://corp.example.com/langsmith",
),
id="self_hosted_path_prefix_is_kept",
),
pytest.param(
"https://custom.host.example",
"https://aks.smith.langchain.dev/api",
("https://custom.host.example", "https://smith.langchain.com"),
id="explicit_host_url_beats_langsmith_endpoint",
),
pytest.param(
"https://api.host.langchain.com",
"https://aks.smith.langchain.dev/api",
CLOUD,
id="explicit_cloud_host_url_beats_self_hosted_endpoint",
),
pytest.param(
"https://smith.example.com/api-host/",
None,
("https://smith.example.com/api-host", "https://smith.example.com"),
id="explicit_api_host_url_derives_dashboard_root",
),
pytest.param(
"https://corp.example.com/langsmith/api-host",
None,
(
"https://corp.example.com/langsmith/api-host",
"https://corp.example.com/langsmith",
),
id="explicit_api_host_url_keeps_path_prefix_in_dashboard",
),
pytest.param(
"http://localhost:8080",
None,
("http://localhost:8080", "http://localhost:8080"),
id="localhost_dashboard_is_the_same_origin",
),
pytest.param(
"http://localhost:8080/api-host",
None,
("http://localhost:8080/api-host", "http://localhost:8080"),
id="localhost_api_host_dashboard_is_the_origin",
),
pytest.param(
"https://eu.api.host.langchain.com",
None,
("https://eu.api.host.langchain.com", "https://eu.smith.langchain.com"),
id="regional_control_plane_maps_to_regional_dashboard",
),
],
)
def test_control_plane_endpoints_resolve(host_url, langsmith_endpoint, expected):
endpoints = ControlPlaneEndpoints.resolve(host_url, langsmith_endpoint)
assert (endpoints.control_plane_url, endpoints.dashboard_url) == expected
@pytest.mark.parametrize(
("payload", "expected"),
[
pytest.param(
{"resources": [{"id": "a"}, {"id": "b"}]},
[{"id": "a"}, {"id": "b"}],
id="list_returns_the_resources",
),
pytest.param({"resources": []}, [], id="empty_list"),
pytest.param({}, [], id="missing_key"),
pytest.param({"resources": None}, [], id="null_resources"),
pytest.param(
{"resources": ["nope", {"id": "a"}]}, [{"id": "a"}], id="skips_non_objects"
),
pytest.param([], [], id="unexpected_envelope"),
],
)
def test_list_endpoints_return_resource_objects(payload, expected):
def handler(req: httpx.Request) -> httpx.Response:
return httpx.Response(200, json=payload)
c = HostBackendClient(
"https://api.example.com", "key", transport=httpx.MockTransport(handler)
)
assert c.list_deployments() == expected
def test_list_listeners_asks_for_a_full_page():
seen: dict = {}
def handler(req: httpx.Request) -> httpx.Response:
seen["url"] = str(req.url)
return httpx.Response(200, json={"resources": [{"id": "listener-1"}]})
c = HostBackendClient(
"https://api.example.com", "key", transport=httpx.MockTransport(handler)
)
assert c.list_listeners() == [{"id": "listener-1"}]
assert seen["url"] == "https://api.example.com/v2/listeners?limit=100"
@pytest.mark.parametrize(
("control_plane_url", "expected"),
[
pytest.param("https://api.host.langchain.com", True, id="cloud"),
pytest.param("https://eu.api.host.langchain.com", True, id="cloud_region"),
pytest.param("https://dev.api.host.langchain.com", True, id="cloud_dev"),
pytest.param("https://smith.example.com/api-host", False, id="self_hosted"),
pytest.param(
"https://corp.example.com/langsmith/api-host",
False,
id="self_hosted_prefix",
),
pytest.param("http://localhost:8080/api-host", False, id="local"),
pytest.param(
"https://evil-api.host.langchain.com", False, id="lookalike_needs_a_dot"
),
],
)
def test_is_cloud_recognises_the_managed_control_plane(control_plane_url, expected):
endpoints = ControlPlaneEndpoints.from_control_plane_url(control_plane_url)
assert endpoints.is_cloud is expected
@pytest.mark.parametrize(
("call", "expected_params"),
[
pytest.param(
lambda c: c.list_deployments(name="agent"),
{"name": "agent"},
id="exact_name_filters_server_side",
),
pytest.param(
lambda c: c.list_deployments(name_contains="age"),
{"name_contains": "age"},
id="substring_search_keeps_its_own_parameter",
),
pytest.param(
lambda c: c.list_deployments(),
{},
id="no_filter_sends_no_parameters",
),
pytest.param(
lambda c: c.list_deployments(
name="agent", name_contains="agent", limit=100
),
{"name": "agent", "name_contains": "agent", "limit": "100"},
id="both_filters_travel_together_for_older_servers",
),
],
)
def test_list_deployments_sends_one_name_filter(call, expected_params):
seen: dict = {}
def handler(req: httpx.Request) -> httpx.Response:
seen.update(dict(req.url.params))
return httpx.Response(200, json={"resources": []})
call(
HostBackendClient(
"https://api.example.com", "key", transport=httpx.MockTransport(handler)
)
)
assert seen == expected_params
@pytest.mark.parametrize(
("body", "expected"),
[
pytest.param(
{"detail": "Source configuration error: bad listener"},
"Source configuration error: bad listener",
id="fastapi_detail_is_unwrapped",
),
pytest.param(
{"detail": {"loc": ["body"], "msg": "nope"}},
None,
id="a_structured_detail_is_left_alone",
),
pytest.param({"other": "shape"}, None, id="an_unknown_shape_is_left_alone"),
],
)
def test_error_detail_is_readable(body, expected):
c = HostBackendClient(
"https://api.example.com",
"key",
transport=httpx.MockTransport(lambda req: httpx.Response(400, json=body)),
)
with pytest.raises(HostBackendError) as error:
c.get_deployment("dep-1")
assert error.value.detail == expected
if expected is not None:
assert error.value.message.endswith(expected)
@@ -0,0 +1,71 @@
import pytest
from langgraph_cli.image_reference import ImageReference
@pytest.mark.parametrize(
("reference", "repository", "tag"),
[
pytest.param(
"registry.example.com/team/app:v1",
"registry.example.com/team/app",
"v1",
id="tag_after_last_slash",
),
pytest.param(
"registry.example.com/team/app",
"registry.example.com/team/app",
None,
id="no_tag",
),
pytest.param(
"localhost:5000/app",
"localhost:5000/app",
None,
id="registry_port_is_not_a_tag",
),
pytest.param(
"localhost:5000/app:latest",
"localhost:5000/app",
"latest",
id="registry_port_with_tag",
),
pytest.param("app:dev", "app", "dev", id="bare_name_with_tag"),
],
)
def test_parse_splits_repository_and_tag(reference, repository, tag):
assert ImageReference.parse(reference) == ImageReference(repository, tag)
def test_with_tag_replaces_the_tag():
assert ImageReference("r/app", "v1").with_tag("v2") == ImageReference("r/app", "v2")
@pytest.mark.parametrize(
("reference", "expected"),
[
pytest.param(ImageReference("r/app", "v1"), "r/app:v1", id="tagged"),
pytest.param(ImageReference("r/app"), "r/app", id="untagged"),
],
)
def test_str_renders_the_docker_reference(reference, expected):
assert str(reference) == expected
@pytest.mark.parametrize(
("repo_digest", "expected"),
[
pytest.param("localhost:5000/app@sha256:abc", True, id="same_repository"),
pytest.param("localhost:5000/app-2@sha256:abc", False, id="other_repository"),
pytest.param("mirror.example.com/app@sha256:abc", False, id="other_registry"),
],
)
def test_matches_digest_only_for_the_same_repository(repo_digest, expected):
assert ImageReference("localhost:5000/app", "v1").matches_digest(repo_digest) is (
expected
)
def test_parse_rejects_a_digest_reference():
with pytest.raises(ValueError, match="digest"):
ImageReference.parse("registry.example.com/app@sha256:abc")
+63 -22
View File
@@ -8,13 +8,15 @@ from typing import Any, cast
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import (
BaseCheckpointSaver,
ChannelVersions,
Checkpoint,
PendingWrite,
)
from langgraph.checkpoint.base.id import uuid6
from langgraph.checkpoint.serde.types import _DeltaSnapshot
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.channels.base import BaseChannel
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(
channels: Mapping[str, BaseChannel],
updated_channels: set[str],
@@ -122,6 +141,7 @@ def create_checkpoint_plan_for_update_state_api(
parents: dict[str, Any],
saved_metadata: Mapping[str, Any] | None,
is_fresh_thread: bool,
fork_channels: set[str],
) -> tuple[set[str], dict[str, Any]]:
"""Return ``(channels_to_snapshot, metadata)`` for an update_state head."""
metadata: dict[str, Any] = {
@@ -137,7 +157,9 @@ def create_checkpoint_plan_for_update_state_api(
updated_channels,
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:
new_counters[k] = (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()
channels_to_snapshot = channels_to_snapshot or set()
bumped: dict[str, tuple[Any, Any]] = {}
if channels is None:
values = checkpoint["channel_values"]
channel_versions = checkpoint["channel_versions"]
@@ -174,30 +197,29 @@ def create_checkpoint(
values = {}
channel_versions = dict(checkpoint["channel_versions"])
for k in channels:
if k not in channel_versions:
continue
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:
# Callers force a full snapshot blob here: exit mode when a
# delta channel reaches its snapshot cadence, and update_state
# on a fresh thread (no ancestor to replay writes from). The
# manual version-bump below only applies to the exit-mode case.
#
# 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).
# `put` only stores a blob for a channel whose version moved,
# so snapshotting a channel this step did not write needs a
# bump: exit mode reaching the cadence on a superstep that
# skipped the channel, and a fork's first checkpoint.
if get_next_version is not None and (
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())
else:
v = ch.checkpoint()
@@ -209,11 +231,30 @@ def create_checkpoint(
id=id or str(uuid6(clock_seq=step)),
channel_values=values,
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),
)
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:
"""True if `spec` is a `DeltaChannel` and no value is stored at this
checkpoint, requiring an ancestor walk to reconstruct.
+22 -11
View File
@@ -102,6 +102,7 @@ from langgraph.pregel._checkpoint import (
copy_checkpoint,
create_checkpoint,
delta_channels_to_snapshot,
delta_channels_with_pending_writes,
empty_checkpoint,
exit_delta_task_id,
)
@@ -222,10 +223,13 @@ class PregelLoop:
# under the saver's `ORDER BY task_id, idx` sorting.
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
# Delta channels that saw an Overwrite since the last checkpoint. These
# channels must snapshot after live update applies overwrite semantics so
# sparse replay starts from the same post-overwrite value.
_delta_channels_with_overwrite: set[str]
# Delta channels that must snapshot at the next checkpoint, whatever their
# cadence counters say:
# * an Overwrite arrived since the last checkpoint, so sparse replay has to
# 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__`
# (or the synthetic-empty checkpoint, on first run). We capture it
@@ -683,7 +687,7 @@ class PregelLoop:
def after_tick(self) -> None:
# finish superstep
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
for ch, v in writes
if isinstance(self.specs.get(ch), DeltaChannel) and _get_overwrite(v)[0]
@@ -898,6 +902,15 @@ class PregelLoop:
self.checkpoint_pending_writes = [
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
if input_is_command:
@@ -991,7 +1004,7 @@ class PregelLoop:
manager=None,
updated_channels=updated_channels,
)
self._delta_channels_with_overwrite.update(
self._delta_channels_forced_snapshot.update(
c
for c, v in input_writes
if isinstance(self.specs.get(c), DeltaChannel) and _get_overwrite(v)[0]
@@ -1136,7 +1149,7 @@ class PregelLoop:
# create new checkpoint
channels_to_snapshot = (
delta_channels_to_snapshot(self.channels, new_counters)
| self._delta_channels_with_overwrite
| self._delta_channels_forced_snapshot
if do_checkpoint
else set()
)
@@ -1154,7 +1167,7 @@ class PregelLoop:
for k in channels_to_snapshot:
new_counters[k] = (0, 0)
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)}
if non_zero:
self.checkpoint_metadata["counters_since_delta_snapshot"] = non_zero
@@ -1239,7 +1252,7 @@ class PregelLoop:
)
channels_to_snapshot = (
delta_channels_to_snapshot(self.channels, counters)
| self._delta_channels_with_overwrite
| self._delta_channels_forced_snapshot
)
pending = [
@@ -1684,7 +1697,6 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
)
self._delta_write_futs = []
self._error_handler_write_futs = []
self._delta_channels_with_overwrite = set()
self._exit_delta_writes = (
[] 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._error_handler_write_futs = []
self._delta_channels_with_overwrite = set()
self._exit_delta_writes = (
[] if self.durability == "exit" and self.checkpointer is not None else None
)
+141 -24
View File
@@ -108,6 +108,7 @@ from langgraph.callbacks import (
get_sync_graph_callback_manager_for_config,
)
from langgraph.channels.base import BaseChannel
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.topic import Topic
from langgraph.config import get_config
from langgraph.constants import END
@@ -133,6 +134,7 @@ from langgraph.pregel._checkpoint import (
copy_checkpoint,
create_checkpoint,
create_checkpoint_plan_for_update_state_api,
delta_channels_with_pending_writes,
empty_checkpoint,
get_updated_channels_from_tasks,
)
@@ -1637,12 +1639,22 @@ class Pregel(
else:
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(
input_config: RunnableConfig, updates: Sequence[StateUpdate]
) -> RunnableConfig:
nonlocal fork_pending
# get last checkpoint
config = ensure_config(self.config, input_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:
self._migrate_checkpoint(saved.checkpoint)
checkpoint = (
@@ -1726,9 +1738,17 @@ class Pregel(
self.trigger_to_nodes,
)
# 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(
checkpoint_config,
create_checkpoint(checkpoint, channels, step),
next_checkpoint,
{
"source": "update",
"step": step + 1,
@@ -1736,7 +1756,7 @@ class Pregel(
},
get_new_channel_versions(
checkpoint_previous_versions,
checkpoint["channel_versions"],
next_checkpoint["channel_versions"],
),
)
return patch_checkpoint_map(
@@ -1765,9 +1785,17 @@ class Pregel(
if saved and saved.metadata.get("step") is not None
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(
checkpoint_config,
create_checkpoint(checkpoint, channels, next_step),
next_checkpoint,
{
"source": "input",
"step": next_step,
@@ -1777,7 +1805,7 @@ class Pregel(
},
get_new_channel_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)
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]
if channel_writes:
checkpointer.put_writes(
checkpoint_config, channel_writes, task_id
)
edited_delta_channels = {
ch
for ch in updated_channels
if isinstance(self.channels.get(ch), DeltaChannel)
}
# 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(
checkpoint,
channels,
@@ -2020,18 +2056,29 @@ class Pregel(
parents=saved.metadata.get("parents", {}) if saved else {},
saved_metadata=saved.metadata if saved else None,
is_fresh_thread=saved is None,
fork_channels=fork_pending,
)
)
checkpoint = create_checkpoint(
checkpoint,
channels,
step + 1,
updated_channels=updated_channels if channels_to_snapshot else None,
get_next_version=checkpointer.get_next_version
if channels_to_snapshot
else None,
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(
checkpoint_config,
checkpoint,
@@ -2103,12 +2150,22 @@ class Pregel(
else:
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(
input_config: RunnableConfig, updates: Sequence[StateUpdate]
) -> RunnableConfig:
nonlocal fork_pending
# get last checkpoint
config = ensure_config(self.config, input_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:
self._migrate_checkpoint(saved.checkpoint)
checkpoint = (
@@ -2190,16 +2247,25 @@ class Pregel(
self.trigger_to_nodes,
)
# 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(
checkpoint_config,
create_checkpoint(checkpoint, channels, step),
next_checkpoint,
{
"source": "update",
"step": step + 1,
"parents": saved.metadata.get("parents", {}) if saved else {},
},
get_new_channel_versions(
checkpoint_previous_versions, checkpoint["channel_versions"]
checkpoint_previous_versions,
next_checkpoint["channel_versions"],
),
)
return patch_checkpoint_map(
@@ -2228,9 +2294,17 @@ class Pregel(
if saved and saved.metadata.get("step") is not None
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(
checkpoint_config,
create_checkpoint(checkpoint, channels, next_step),
next_checkpoint,
{
"source": "input",
"step": next_step,
@@ -2240,7 +2314,7 @@ class Pregel(
},
get_new_channel_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)
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]
if channel_writes:
await checkpointer.aput_writes(
checkpoint_config, channel_writes, task_id
)
edited_delta_channels = {
ch
for ch in updated_channels
if isinstance(self.channels.get(ch), DeltaChannel)
}
# 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(
checkpoint,
channels,
@@ -2480,18 +2562,29 @@ class Pregel(
parents=saved.metadata.get("parents", {}) if saved else {},
saved_metadata=saved.metadata if saved else None,
is_fresh_thread=saved is None,
fork_channels=fork_pending,
)
)
checkpoint = create_checkpoint(
checkpoint,
channels,
step + 1,
updated_channels=updated_channels if channels_to_snapshot else None,
get_next_version=checkpointer.get_next_version
if channels_to_snapshot
else None,
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(
checkpoint_config,
checkpoint,
@@ -4181,6 +4274,30 @@ def _trigger_to_nodes(nodes: dict[str, PregelNode]) -> Mapping[str, Sequence[str
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(
stream_mode: StreamMode | Sequence[StreamMode],
print_mode: StreamMode | Sequence[StreamMode],
+5 -3
View File
@@ -85,11 +85,13 @@ class MemorySaverAssertImmutable(InMemorySaver):
)
== saved
), 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.serde.dumps_typed(checkpoint)
self.serde.dumps_typed(super().get(next_config))
)
# call super to write checkpoint
return super().put(config, checkpoint, metadata, new_versions)
return next_config
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.parent_config is None
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 == ()