mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-03 15:05:06 +02:00
Compare commits
17
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
741c6f8d50 | ||
|
|
37a5504433 | ||
|
|
f6286dce38 | ||
|
|
93a0dfec08 | ||
|
|
56834787eb | ||
|
|
96199c4fee | ||
|
|
682814e944 | ||
|
|
210c4b3877 | ||
|
|
81489ab858 | ||
|
|
b7b052e66c | ||
|
|
d595f11b43 | ||
|
|
ed540155e3 | ||
|
|
f78892d462 | ||
|
|
7488cf2448 | ||
|
|
e77201cbb1 | ||
|
|
acae5e23b0 | ||
|
|
14ce607111 |
+482
-121
@@ -13,6 +13,7 @@ import tempfile
|
||||
import time
|
||||
from collections.abc import Callable, Sequence
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import click
|
||||
import click.exceptions
|
||||
@@ -26,10 +27,11 @@ from langgraph_cli.config import Config
|
||||
from langgraph_cli.constants import DEFAULT_CONFIG, DEFAULT_PORT
|
||||
from langgraph_cli.docker import DockerCapabilities
|
||||
from langgraph_cli.exec import Runner, subp_exec
|
||||
from langgraph_cli.helpers import format_log_entry, level_fg, resolve_deployment_id
|
||||
from langgraph_cli.host_backend import HostBackendClient, HostBackendError
|
||||
from langgraph_cli.progress import Progress
|
||||
from langgraph_cli.templates import TEMPLATE_HELP_STRING, create_new
|
||||
from langgraph_cli.util import warn_non_wolfi_distro
|
||||
from langgraph_cli.util import format_deployments_table, warn_non_wolfi_distro
|
||||
from langgraph_cli.version import __version__
|
||||
|
||||
RESERVED_ENV_VARS = frozenset(
|
||||
@@ -287,6 +289,33 @@ OPT_API_VERSION = click.option(
|
||||
help="API server version to use for the base image. If unspecified, the latest version will be used.",
|
||||
)
|
||||
|
||||
OPT_HOST_API_KEY = click.option(
|
||||
"--api-key",
|
||||
envvar="LANGGRAPH_HOST_API_KEY",
|
||||
help=(
|
||||
"API key. Can also be set via LANGGRAPH_HOST_API_KEY, "
|
||||
"LANGSMITH_API_KEY, or LANGCHAIN_API_KEY environment variable or .env file."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
OPT_HOST_DEPLOYMENT_NAME = click.option(
|
||||
"--name",
|
||||
envvar=_DEPLOYMENT_NAME_ENV,
|
||||
help=(
|
||||
"Deployment name. Can also be set via LANGSMITH_DEPLOYMENT_NAME "
|
||||
"environment variable or .env file. Defaults to current directory name "
|
||||
"if --deployment-id is not provided."
|
||||
),
|
||||
)
|
||||
|
||||
OPT_HOST_URL = click.option(
|
||||
"--host-url",
|
||||
envvar="LANGGRAPH_HOST_URL",
|
||||
default="https://api.host.langchain.com",
|
||||
hidden=True,
|
||||
)
|
||||
|
||||
OPT_ENGINE_RUNTIME_MODE = click.option(
|
||||
"--engine-runtime-mode",
|
||||
type=click.Choice(["combined_queue_worker", "distributed"]),
|
||||
@@ -295,7 +324,67 @@ OPT_ENGINE_RUNTIME_MODE = click.option(
|
||||
)
|
||||
|
||||
|
||||
@click.group()
|
||||
class NestedHelpGroup(click.Group):
|
||||
"""Click group that shows one level of nested subcommands in top-level help."""
|
||||
|
||||
def format_commands(
|
||||
self, ctx: click.Context, formatter: click.HelpFormatter
|
||||
) -> None:
|
||||
command_entries: list[tuple[str, click.Command]] = []
|
||||
# Collect the top-level commands first, then append one level of nested
|
||||
# subcommands using names like "deploy list" so they show up in the
|
||||
# top-level help output.
|
||||
for command_name in self.list_commands(ctx):
|
||||
command = self.get_command(ctx, command_name)
|
||||
if command is None or command.hidden:
|
||||
continue
|
||||
command_entries.append((command_name, command))
|
||||
if isinstance(command, click.Group):
|
||||
# Build a child context so Click resolves the subcommands the same
|
||||
# way it would for the nested group itself.
|
||||
sub_ctx = click.Context(command, info_name=command_name, parent=ctx)
|
||||
for subcommand_name in command.list_commands(sub_ctx):
|
||||
subcommand = command.get_command(sub_ctx, subcommand_name)
|
||||
if subcommand is None or subcommand.hidden:
|
||||
continue
|
||||
command_entries.append(
|
||||
(f"{command_name} {subcommand_name}", subcommand)
|
||||
)
|
||||
|
||||
# Compute the available width for help text up front so we can truncate
|
||||
# descriptions before handing them to Click. That keeps each command on
|
||||
# a single line instead of allowing wrapped descriptions.
|
||||
command_width = max((len(name) for name, _ in command_entries), default=0)
|
||||
help_width = max(formatter.width - command_width - 6, 10)
|
||||
rows = [
|
||||
(name, command.get_short_help_str(help_width))
|
||||
for name, command in command_entries
|
||||
]
|
||||
|
||||
if rows:
|
||||
# Render the flattened command list using Click's standard
|
||||
# definition-list formatter so alignment stays consistent with the
|
||||
# rest of the CLI help output.
|
||||
with formatter.section("Commands"):
|
||||
formatter.write_dl(rows)
|
||||
|
||||
|
||||
class DeployGroup(NestedHelpGroup):
|
||||
"""Group that treats leading '-' args as passthrough docker flags."""
|
||||
|
||||
def parse_args(self, ctx: click.Context, args: list[str]) -> list[str]:
|
||||
result = super().parse_args(ctx, args)
|
||||
if ctx._protected_args and ctx._protected_args[0].startswith("-"):
|
||||
# Click stores the would-be subcommand in _protected_args; if it looks
|
||||
# like an option (e.g. --build-arg) treat it as passthrough docker
|
||||
# args instead of insisting on a nested command.
|
||||
ctx.args = [*ctx._protected_args, *ctx.args]
|
||||
ctx._protected_args = []
|
||||
return ctx.args
|
||||
return result
|
||||
|
||||
|
||||
@click.group(cls=NestedHelpGroup)
|
||||
@click.version_option(version=__version__, prog_name="LangGraph CLI")
|
||||
def cli():
|
||||
pass
|
||||
@@ -593,84 +682,109 @@ def build(
|
||||
)
|
||||
|
||||
|
||||
@click.option(
|
||||
"--api-key",
|
||||
envvar="LANGGRAPH_HOST_API_KEY",
|
||||
help=(
|
||||
"API key. Can also be set via LANGGRAPH_HOST_API_KEY, "
|
||||
"LANGSMITH_API_KEY, or LANGCHAIN_API_KEY environment variable or .env file."
|
||||
),
|
||||
)
|
||||
@click.option(
|
||||
"--name",
|
||||
envvar="LANGSMITH_DEPLOYMENT_NAME",
|
||||
help=(
|
||||
"Deployment name. Can also be set via LANGSMITH_DEPLOYMENT_NAME "
|
||||
"environment variable or .env file. Defaults to current directory name "
|
||||
"if --deployment-id is not provided."
|
||||
),
|
||||
)
|
||||
@click.option(
|
||||
"--deployment-id",
|
||||
help=(
|
||||
"ID of an existing deployment to update. If omitted, "
|
||||
"--name is used to find or create the deployment."
|
||||
),
|
||||
)
|
||||
@click.option(
|
||||
"--deployment-type",
|
||||
type=click.Choice(["dev", "prod"]),
|
||||
default="dev",
|
||||
show_default=True,
|
||||
help="Deployment type (used when creating a new deployment).",
|
||||
)
|
||||
@click.option(
|
||||
"--no-wait",
|
||||
is_flag=True,
|
||||
default=False,
|
||||
help="Skip waiting for deployment status.",
|
||||
)
|
||||
@OPT_VERBOSE
|
||||
@click.option(
|
||||
"--host-url",
|
||||
envvar="LANGGRAPH_HOST_URL",
|
||||
default="https://api.host.langchain.com",
|
||||
hidden=True,
|
||||
)
|
||||
@click.option("--image-name", hidden=True)
|
||||
@click.option("--image-tag", default="latest", hidden=True)
|
||||
@click.option(
|
||||
"--config",
|
||||
"-c",
|
||||
default=DEFAULT_CONFIG,
|
||||
hidden=True,
|
||||
type=click.Path(
|
||||
exists=True,
|
||||
file_okay=True,
|
||||
dir_okay=False,
|
||||
resolve_path=True,
|
||||
path_type=pathlib.Path,
|
||||
),
|
||||
)
|
||||
@click.option("--pull/--no-pull", default=True, hidden=True)
|
||||
@click.option("--base-image", hidden=True)
|
||||
@click.option("--install-command", hidden=True)
|
||||
@click.option("--build-command", hidden=True)
|
||||
@click.option("--api-version", type=str, hidden=True)
|
||||
@click.argument("docker_build_args", nargs=-1, type=click.UNPROCESSED)
|
||||
@cli.command(
|
||||
def _deploy_base_options(
|
||||
func: Callable | None = None,
|
||||
*,
|
||||
include_docker_args: bool = True,
|
||||
validate_config_path: bool = True,
|
||||
):
|
||||
"""Apply shared deploy flags.
|
||||
|
||||
The group shares most options but should not consume subcommands, so the
|
||||
docker build args are only attached when requested.
|
||||
"""
|
||||
|
||||
def _apply(target: Callable) -> Callable:
|
||||
decorators = [
|
||||
OPT_HOST_API_KEY,
|
||||
OPT_HOST_DEPLOYMENT_NAME,
|
||||
click.option(
|
||||
"--deployment-id",
|
||||
help=(
|
||||
"ID of an existing deployment to update. If omitted, "
|
||||
"--name is used to find or create the deployment."
|
||||
),
|
||||
),
|
||||
click.option(
|
||||
"--deployment-type",
|
||||
type=click.Choice(["dev", "prod"]),
|
||||
default="dev",
|
||||
show_default=True,
|
||||
help="Deployment type (used when creating a new deployment).",
|
||||
),
|
||||
click.option(
|
||||
"--no-wait",
|
||||
is_flag=True,
|
||||
default=False,
|
||||
help="Skip waiting for deployment status.",
|
||||
),
|
||||
OPT_VERBOSE,
|
||||
OPT_HOST_URL,
|
||||
click.option("--image-name", hidden=True),
|
||||
click.option("--image-tag", default="latest", hidden=True),
|
||||
click.option(
|
||||
"--config",
|
||||
"-c",
|
||||
default=DEFAULT_CONFIG,
|
||||
hidden=True,
|
||||
type=click.Path(
|
||||
exists=validate_config_path,
|
||||
file_okay=True,
|
||||
dir_okay=False,
|
||||
resolve_path=True,
|
||||
path_type=pathlib.Path,
|
||||
),
|
||||
),
|
||||
click.option("--pull/--no-pull", default=True, hidden=True),
|
||||
click.option("--base-image", hidden=True),
|
||||
click.option("--install-command", hidden=True),
|
||||
click.option("--build-command", hidden=True),
|
||||
click.option("--api-version", type=str, hidden=True),
|
||||
]
|
||||
if include_docker_args:
|
||||
# Only attach build args to the default command; on the group they
|
||||
# would capture subcommand names like `list` before Click resolves
|
||||
# them, making those subcommands unreachable.
|
||||
decorators.append(
|
||||
click.argument("docker_build_args", nargs=-1, type=click.UNPROCESSED)
|
||||
)
|
||||
for decorator in reversed(decorators):
|
||||
target = decorator(target)
|
||||
return target
|
||||
|
||||
return _apply(func) if func is not None else _apply
|
||||
|
||||
|
||||
@cli.group(
|
||||
cls=DeployGroup,
|
||||
help=(
|
||||
"[Beta] Build and deploy a LangGraph image to LangSmith Deployments.\n\n"
|
||||
"This command is in beta and under active development. "
|
||||
"Expect frequent updates and improvements.\n\n"
|
||||
"Run from the root of your LangGraph project (where langgraph.json "
|
||||
"is located). This command also accepts build flags (--base-image, "
|
||||
"--pull, etc.). See 'langgraph build --help' for details."
|
||||
"--config, --pull, etc.). See 'langgraph build --help' for details."
|
||||
),
|
||||
context_settings=dict(ignore_unknown_options=True),
|
||||
context_settings=dict(ignore_unknown_options=True, allow_extra_args=True),
|
||||
invoke_without_command=True, # allow `deploy` click group to execute without command
|
||||
)
|
||||
@_deploy_base_options(include_docker_args=False, validate_config_path=False)
|
||||
@click.pass_context
|
||||
@log_command
|
||||
def deploy(
|
||||
def deploy(ctx: click.Context, **_: object):
|
||||
# We register deploy as both a group and a command here.
|
||||
# if we detect no subcommand, we run _deploy (basically run langgraph deploy as a top level command)
|
||||
# otherwise, we return None here and click will proceed to actually run the subcommand (list or delete)
|
||||
if ctx.invoked_subcommand is not None:
|
||||
return
|
||||
docker_build_args = tuple(ctx.args)
|
||||
ctx.args = [] # Prevent Click from re-processing passthrough args later.
|
||||
return ctx.forward(_deploy, docker_build_args=docker_build_args)
|
||||
|
||||
|
||||
@_deploy_base_options()
|
||||
@click.command(context_settings=dict(ignore_unknown_options=True))
|
||||
def _deploy(
|
||||
config: pathlib.Path,
|
||||
pull: bool,
|
||||
verbose: bool,
|
||||
@@ -698,15 +812,6 @@ def deploy(
|
||||
|
||||
env_vars = _parse_env_from_config(config_json, config)
|
||||
|
||||
if not api_key:
|
||||
for key_name in _API_KEY_ENV_NAMES:
|
||||
val = env_vars.get(key_name) or os.environ.get(key_name)
|
||||
if val:
|
||||
api_key = val
|
||||
break
|
||||
if not api_key:
|
||||
api_key = click.prompt("Host API key", hide_input=True)
|
||||
|
||||
if not deployment_id and not name:
|
||||
name = env_vars.get(_DEPLOYMENT_NAME_ENV)
|
||||
if not deployment_id and not name:
|
||||
@@ -743,55 +848,21 @@ def deploy(
|
||||
def log_step(message: str) -> None:
|
||||
click.secho(message, fg="cyan")
|
||||
|
||||
client = HostBackendClient(host_url, api_key)
|
||||
client = _create_host_backend_client(host_url, api_key, env_vars=env_vars)
|
||||
step = 1
|
||||
needs_creation = False
|
||||
|
||||
if deployment_id:
|
||||
log_step(f"{step}. Using deployment {deployment_id}")
|
||||
try:
|
||||
client.get_deployment(deployment_id)
|
||||
except HostBackendError as err:
|
||||
if (
|
||||
err.status_code == 403
|
||||
and "requires workspace specification" in err.message
|
||||
):
|
||||
click.secho(
|
||||
"Your API key is org-scoped and requires a workspace ID.",
|
||||
fg="yellow",
|
||||
)
|
||||
click.secho(
|
||||
"Find your workspace ID in LangSmith under Settings > Workspaces.",
|
||||
fg="yellow",
|
||||
)
|
||||
tenant_id = click.prompt("Workspace ID")
|
||||
client = HostBackendClient(host_url, api_key, tenant_id=tenant_id)
|
||||
client.get_deployment(deployment_id)
|
||||
else:
|
||||
raise
|
||||
_call_host_backend_with_optional_tenant(
|
||||
client, lambda c: c.get_deployment(deployment_id)
|
||||
)
|
||||
step += 1
|
||||
else:
|
||||
log_step(f"{step}. Looking up deployment '{name}'")
|
||||
try:
|
||||
existing = client.list_deployments(name_contains=name)
|
||||
except HostBackendError as err:
|
||||
if (
|
||||
err.status_code == 403
|
||||
and "requires workspace specification" in err.message
|
||||
):
|
||||
click.secho(
|
||||
"Your API key is org-scoped and requires a workspace ID.",
|
||||
fg="yellow",
|
||||
)
|
||||
click.secho(
|
||||
"Find your workspace ID in LangSmith under Settings > Workspaces.",
|
||||
fg="yellow",
|
||||
)
|
||||
tenant_id = click.prompt("Workspace ID")
|
||||
client = HostBackendClient(host_url, api_key, tenant_id=tenant_id)
|
||||
existing = client.list_deployments(name_contains=name)
|
||||
else:
|
||||
raise
|
||||
existing = _call_host_backend_with_optional_tenant(
|
||||
client, lambda c: c.list_deployments(name_contains=name)
|
||||
)
|
||||
found_id = None
|
||||
if isinstance(existing, dict):
|
||||
for dep in existing.get("resources", []):
|
||||
@@ -1050,6 +1121,122 @@ def deploy(
|
||||
)
|
||||
|
||||
|
||||
def _create_host_backend_client(
|
||||
host_url: str | None,
|
||||
api_key: str | None,
|
||||
env_vars: dict[str, str] | None = None,
|
||||
) -> HostBackendClient:
|
||||
if env_vars is None:
|
||||
env_vars = _parse_env_from_config({}, pathlib.Path.cwd() / DEFAULT_CONFIG)
|
||||
resolved_api_key = api_key
|
||||
if not resolved_api_key:
|
||||
for key_name in _API_KEY_ENV_NAMES:
|
||||
val = env_vars.get(key_name)
|
||||
if val:
|
||||
resolved_api_key = val
|
||||
break
|
||||
val = os.environ.get(key_name)
|
||||
if val:
|
||||
resolved_api_key = val
|
||||
break
|
||||
if not resolved_api_key:
|
||||
click.secho(
|
||||
"No LangSmith API key found. Create one at Settings > API Keys in LangSmith.",
|
||||
fg="yellow",
|
||||
)
|
||||
resolved_api_key = click.prompt("Enter LangSmith API key", hide_input=True)
|
||||
return HostBackendClient(host_url, resolved_api_key)
|
||||
|
||||
|
||||
def _call_host_backend_with_optional_tenant(
|
||||
client: HostBackendClient,
|
||||
operation: Callable[[HostBackendClient], object],
|
||||
) -> object:
|
||||
"""Run *operation*, prompting for a workspace ID on org-scoped 403s.
|
||||
|
||||
On success the original *client* is returned as-is. If the user is
|
||||
prompted for a workspace ID, the tenant header is set on *client*
|
||||
in-place so all subsequent calls through the same instance are
|
||||
tenant-aware.
|
||||
"""
|
||||
try:
|
||||
return operation(client)
|
||||
except HostBackendError as err:
|
||||
if err.status_code == 403 and "requires workspace specification" in err.message:
|
||||
click.secho(
|
||||
"Your API key is org-scoped and requires a workspace ID.",
|
||||
fg="yellow",
|
||||
)
|
||||
click.secho(
|
||||
"Find your workspace ID in LangSmith under Settings > Workspaces.",
|
||||
fg="yellow",
|
||||
)
|
||||
tenant_id = click.prompt("Workspace ID")
|
||||
client._client.headers["X-Tenant-ID"] = tenant_id
|
||||
return operation(client)
|
||||
raise
|
||||
|
||||
|
||||
@OPT_HOST_API_KEY
|
||||
@OPT_HOST_URL
|
||||
@click.option(
|
||||
"--name-contains",
|
||||
default="",
|
||||
help="Only show deployments whose names contain this value.",
|
||||
)
|
||||
@deploy.command("list", help="[Beta] List LangSmith Deployments.")
|
||||
def deploy_list(api_key: str | None, host_url: str | None, name_contains: str) -> None:
|
||||
client = _create_host_backend_client(host_url, api_key)
|
||||
response = _call_host_backend_with_optional_tenant(
|
||||
client,
|
||||
lambda c: c.list_deployments(name_contains=name_contains),
|
||||
)
|
||||
resources = response.get("resources", []) if isinstance(response, dict) else []
|
||||
deployments = [item for item in resources if isinstance(item, dict)]
|
||||
if not deployments:
|
||||
click.echo("No deployments found.")
|
||||
return
|
||||
click.echo(format_deployments_table(deployments))
|
||||
|
||||
|
||||
@OPT_HOST_API_KEY
|
||||
@OPT_HOST_URL
|
||||
@click.option(
|
||||
"--force",
|
||||
is_flag=True,
|
||||
default=False,
|
||||
help="Delete without prompting for confirmation.",
|
||||
)
|
||||
@click.argument("deployment_id")
|
||||
@deploy.command(
|
||||
"delete",
|
||||
help=(
|
||||
"[Beta] Delete a LangSmith Deployment.\n\n"
|
||||
"Use the `deploy list` command to list deployment IDs."
|
||||
),
|
||||
)
|
||||
def deploy_delete(
|
||||
api_key: str | None, host_url: str | None, force: bool, deployment_id: str
|
||||
) -> None:
|
||||
if not force:
|
||||
response = click.prompt(
|
||||
click.style(
|
||||
f"Are you sure you want to delete deployment ID {deployment_id}? (Y/n)",
|
||||
fg="yellow",
|
||||
),
|
||||
default="Y",
|
||||
show_default=False,
|
||||
)
|
||||
if response.strip().lower() not in {"y", "yes"}:
|
||||
raise click.Abort()
|
||||
client = _create_host_backend_client(host_url, api_key)
|
||||
_call_host_backend_with_optional_tenant(
|
||||
client,
|
||||
lambda c: c.delete_deployment(deployment_id),
|
||||
)
|
||||
click.secho(f"Deleted deployment {deployment_id}.", fg="green")
|
||||
|
||||
|
||||
def _normalize_image_name(value: str | None) -> str:
|
||||
"""Sanitize a deployment/directory name into a valid Docker repository name.
|
||||
|
||||
@@ -1076,6 +1263,180 @@ def _normalize_image_tag(value: str) -> str:
|
||||
return value
|
||||
|
||||
|
||||
@OPT_HOST_API_KEY
|
||||
@OPT_HOST_DEPLOYMENT_NAME
|
||||
@click.option(
|
||||
"--deployment-id",
|
||||
help="Deployment ID. If omitted, --name is used to find the deployment.",
|
||||
)
|
||||
@click.option(
|
||||
"--type",
|
||||
"log_type",
|
||||
type=click.Choice(["deploy", "build"]),
|
||||
default="deploy",
|
||||
show_default=True,
|
||||
help=(
|
||||
"Log stream to fetch: 'deploy' shows agent server runtime logs; "
|
||||
"'build' shows build logs (for deployments built remotely)."
|
||||
),
|
||||
)
|
||||
@click.option(
|
||||
"--revision-id",
|
||||
help="Specific revision ID. For build logs, defaults to latest revision.",
|
||||
)
|
||||
@click.option(
|
||||
"--level",
|
||||
type=click.Choice(
|
||||
["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"], case_sensitive=False
|
||||
),
|
||||
help="Filter by log level.",
|
||||
)
|
||||
@click.option(
|
||||
"--limit",
|
||||
type=int,
|
||||
default=100,
|
||||
show_default=True,
|
||||
help="Max log entries to fetch.",
|
||||
)
|
||||
@click.option(
|
||||
"--query",
|
||||
"-q",
|
||||
help="Search string filter.",
|
||||
)
|
||||
@click.option(
|
||||
"--start-time",
|
||||
help="ISO8601 start time (e.g. 2026-03-08T00:00:00Z).",
|
||||
)
|
||||
@click.option(
|
||||
"--end-time",
|
||||
help="ISO8601 end time. (e.g. 2026-03-08T00:00:00Z)",
|
||||
)
|
||||
@click.option(
|
||||
"--follow",
|
||||
"-f",
|
||||
is_flag=True,
|
||||
default=False,
|
||||
help="Continuously poll for new logs.",
|
||||
)
|
||||
@OPT_HOST_URL
|
||||
@deploy.command(
|
||||
"logs",
|
||||
help=(
|
||||
"[Beta] Fetch LangSmith Deployment logs. Use 'deploy' for agent runtime "
|
||||
"logs, or 'build' for remote build logs."
|
||||
),
|
||||
)
|
||||
@log_command
|
||||
def deploy_logs(
|
||||
api_key: str | None,
|
||||
name: str | None,
|
||||
deployment_id: str | None,
|
||||
log_type: str,
|
||||
revision_id: str | None,
|
||||
level: str | None,
|
||||
limit: int,
|
||||
query: str | None,
|
||||
start_time: str | None,
|
||||
end_time: str | None,
|
||||
follow: bool,
|
||||
host_url: str,
|
||||
):
|
||||
env_vars = _parse_env_from_config({}, pathlib.Path.cwd() / DEFAULT_CONFIG)
|
||||
client = _create_host_backend_client(host_url, api_key, env_vars=env_vars)
|
||||
if not deployment_id and not name:
|
||||
name = env_vars.get(_DEPLOYMENT_NAME_ENV)
|
||||
dep_id = _call_host_backend_with_optional_tenant(
|
||||
client, lambda c: resolve_deployment_id(c, deployment_id, name)
|
||||
)
|
||||
|
||||
if log_type == "build" and not revision_id:
|
||||
revisions_resp = client.list_revisions(dep_id, limit=1)
|
||||
resources = (
|
||||
revisions_resp.get("resources", [])
|
||||
if isinstance(revisions_resp, dict)
|
||||
else []
|
||||
)
|
||||
if not resources:
|
||||
raise click.ClickException(
|
||||
"No revisions found for this deployment. Cannot fetch build logs."
|
||||
)
|
||||
revision_id = str(resources[0]["id"])
|
||||
click.secho(f"Using latest revision: {revision_id}", fg="cyan")
|
||||
|
||||
payload: dict = {"limit": limit, "order": "desc"}
|
||||
if level:
|
||||
payload["level"] = level.upper()
|
||||
if query:
|
||||
payload["query"] = query
|
||||
if start_time:
|
||||
payload["start_time"] = start_time
|
||||
if end_time:
|
||||
payload["end_time"] = end_time
|
||||
|
||||
def _fetch(request_payload: dict) -> list[dict]:
|
||||
if log_type == "build":
|
||||
resp = client.get_build_logs(dep_id, revision_id, request_payload)
|
||||
else:
|
||||
resp = client.get_deploy_logs(dep_id, request_payload, revision_id)
|
||||
|
||||
if isinstance(resp, dict):
|
||||
return resp.get("logs", [])
|
||||
return []
|
||||
|
||||
def _print_entries(entries: list[dict], *, reverse: bool = False) -> None:
|
||||
iterable = reversed(entries) if reverse else entries
|
||||
for entry in iterable:
|
||||
line = format_log_entry(entry)
|
||||
fg = level_fg(entry.get("level", ""))
|
||||
click.secho(line, fg=fg)
|
||||
|
||||
def _fetch_and_print(request_payload: dict, *, reverse: bool = False) -> list[dict]:
|
||||
entries = _fetch(request_payload)
|
||||
_print_entries(entries, reverse=reverse)
|
||||
return entries
|
||||
|
||||
def _fetch_and_print_new(request_payload: dict, seen_ids: set[str]) -> list[dict]:
|
||||
entries = _fetch(request_payload)
|
||||
new = [e for e in entries if e.get("id", "") not in seen_ids]
|
||||
if new:
|
||||
_print_entries(new)
|
||||
seen_ids.update(e.get("id", "") for e in new)
|
||||
return new
|
||||
|
||||
# initial log fetch will be newest -> oldest, so we need to reverse
|
||||
entries = _fetch_and_print(payload, reverse=True)
|
||||
|
||||
if not follow:
|
||||
if not entries:
|
||||
click.secho("No log entries found.", fg="yellow")
|
||||
return
|
||||
|
||||
payload["order"] = "asc"
|
||||
seen_ids: set[str] = {e.get("id", "") for e in entries if e.get("id")}
|
||||
|
||||
def _update_start_time(ts) -> None:
|
||||
if ts is None:
|
||||
return
|
||||
if isinstance(ts, (int, float)):
|
||||
dt = datetime.fromtimestamp(ts / 1000, tz=timezone.utc)
|
||||
payload["start_time"] = dt.isoformat()
|
||||
else:
|
||||
payload["start_time"] = str(ts)
|
||||
|
||||
if entries:
|
||||
# entries are in descending order here, so index 0 is the newest log
|
||||
_update_start_time(entries[0].get("timestamp"))
|
||||
|
||||
try:
|
||||
while True:
|
||||
time.sleep(2)
|
||||
new_entries = _fetch_and_print_new(payload, seen_ids)
|
||||
if new_entries:
|
||||
_update_start_time(new_entries[-1].get("timestamp"))
|
||||
except KeyboardInterrupt:
|
||||
click.echo("\nStopped.")
|
||||
|
||||
|
||||
def _get_docker_ignore_content() -> str:
|
||||
"""Return the content of a .dockerignore file.
|
||||
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
"""Helpers for the ``langgraph logs`` CLI command."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import click
|
||||
|
||||
from langgraph_cli.host_backend import HostBackendClient
|
||||
|
||||
|
||||
def resolve_deployment_id(
|
||||
client: HostBackendClient,
|
||||
deployment_id: str | None,
|
||||
name: str | None,
|
||||
) -> str:
|
||||
"""Resolve a deployment ID from --deployment-id or --name."""
|
||||
if deployment_id:
|
||||
return deployment_id
|
||||
if not name:
|
||||
raise click.UsageError("Either --deployment-id or --name is required.")
|
||||
existing = client.list_deployments(name_contains=name)
|
||||
if isinstance(existing, dict):
|
||||
for dep in existing.get("resources", []):
|
||||
if isinstance(dep, dict) and dep.get("name") == name:
|
||||
found_id = dep.get("id")
|
||||
if found_id:
|
||||
return str(found_id)
|
||||
raise click.ClickException(f"Deployment '{name}' not found.")
|
||||
|
||||
|
||||
def format_timestamp(ts) -> str:
|
||||
"""Convert a timestamp (epoch ms or string) to a readable string."""
|
||||
if isinstance(ts, (int, float)):
|
||||
dt = datetime.fromtimestamp(ts / 1000, tz=timezone.utc)
|
||||
return dt.strftime("%Y-%m-%d %H:%M:%S")
|
||||
return str(ts) if ts else ""
|
||||
|
||||
|
||||
def format_log_entry(entry: dict) -> str:
|
||||
"""Format a single log entry for display."""
|
||||
ts = format_timestamp(entry.get("timestamp", ""))
|
||||
level = entry.get("level", "")
|
||||
message = entry.get("message", "")
|
||||
if ts and level:
|
||||
return f"[{ts}] [{level}] {message}"
|
||||
elif ts:
|
||||
return f"[{ts}] {message}"
|
||||
return message
|
||||
|
||||
|
||||
def level_fg(level: str) -> str | None:
|
||||
"""Return click color for a log level."""
|
||||
level_upper = level.upper() if level else ""
|
||||
if level_upper in {"ERROR", "CRITICAL"}:
|
||||
return "red"
|
||||
if level_upper == "WARNING":
|
||||
return "yellow"
|
||||
return None
|
||||
@@ -19,7 +19,12 @@ class HostBackendError(click.ClickException):
|
||||
class HostBackendClient:
|
||||
"""Minimal JSON HTTP client for the host backend deployment service."""
|
||||
|
||||
def __init__(self, base_url: str, api_key: str, tenant_id: str | None = None):
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
tenant_id: str | None = None,
|
||||
):
|
||||
if not base_url:
|
||||
raise click.UsageError("Host backend URL is required")
|
||||
transport = httpx.HTTPTransport(retries=3)
|
||||
@@ -30,7 +35,6 @@ class HostBackendClient:
|
||||
if tenant_id:
|
||||
headers["X-Tenant-ID"] = tenant_id
|
||||
self._base_url = base_url.rstrip("/")
|
||||
self._api_key = api_key
|
||||
self._client = httpx.Client(
|
||||
base_url=self._base_url,
|
||||
headers=headers,
|
||||
@@ -39,10 +43,14 @@ class HostBackendClient:
|
||||
)
|
||||
|
||||
def _request(
|
||||
self, method: str, path: str, payload: dict[str, Any] | None = None
|
||||
self,
|
||||
method: str,
|
||||
path: str,
|
||||
payload: dict[str, Any] | None = None,
|
||||
params: dict[str, Any] | None = None,
|
||||
) -> Any:
|
||||
try:
|
||||
resp = self._client.request(method, path, json=payload)
|
||||
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)
|
||||
@@ -65,12 +73,19 @@ class HostBackendClient:
|
||||
def create_deployment(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
return self._request("POST", "/v2/deployments", payload)
|
||||
|
||||
def list_deployments(self, name_contains: str) -> dict[str, Any]:
|
||||
return self._request("GET", f"/v2/deployments?name_contains={name_contains}")
|
||||
def list_deployments(self, name_contains: str = "") -> dict[str, Any]:
|
||||
return self._request(
|
||||
"GET",
|
||||
"/v2/deployments",
|
||||
params={"name_contains": name_contains},
|
||||
)
|
||||
|
||||
def get_deployment(self, deployment_id: str) -> dict[str, Any]:
|
||||
return self._request("GET", f"/v2/deployments/{deployment_id}")
|
||||
|
||||
def delete_deployment(self, deployment_id: str) -> None:
|
||||
return self._request("DELETE", f"/v2/deployments/{deployment_id}")
|
||||
|
||||
def request_push_token(self, deployment_id: str) -> dict[str, Any]:
|
||||
return self._request(
|
||||
"POST",
|
||||
@@ -105,3 +120,24 @@ class HostBackendClient:
|
||||
"GET",
|
||||
f"/v2/deployments/{deployment_id}/revisions/{revision_id}",
|
||||
)
|
||||
|
||||
def get_build_logs(
|
||||
self, project_id: str, revision_id: str, payload: dict[str, Any]
|
||||
) -> Any:
|
||||
return self._request(
|
||||
"POST",
|
||||
f"/v1/projects/{project_id}/revisions/{revision_id}/build_logs",
|
||||
payload,
|
||||
)
|
||||
|
||||
def get_deploy_logs(
|
||||
self,
|
||||
project_id: str,
|
||||
payload: dict[str, Any],
|
||||
revision_id: str | None = None,
|
||||
) -> Any:
|
||||
if revision_id:
|
||||
path = f"/v1/projects/{project_id}/revisions/{revision_id}/deploy_logs"
|
||||
else:
|
||||
path = f"/v1/projects/{project_id}/deploy_logs"
|
||||
return self._request("POST", path, payload)
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
from collections.abc import Sequence
|
||||
|
||||
import click
|
||||
|
||||
|
||||
@@ -23,3 +25,35 @@ def warn_non_wolfi_distro(config_json: dict) -> None:
|
||||
fg="yellow",
|
||||
)
|
||||
click.secho("") # Empty line for better readability
|
||||
|
||||
|
||||
def _extract_deployment_url(deployment: dict[str, object]) -> str:
|
||||
source_config = deployment.get("source_config")
|
||||
if isinstance(source_config, dict):
|
||||
custom_url = source_config.get("custom_url")
|
||||
if isinstance(custom_url, str) and custom_url:
|
||||
return custom_url
|
||||
return "-"
|
||||
|
||||
|
||||
def format_deployments_table(deployments: Sequence[dict[str, object]]) -> str:
|
||||
headers = ("Deployment ID", "Deployment Name", "Deployment URL")
|
||||
rows = [
|
||||
(
|
||||
str(deployment.get("id", "-") or "-"),
|
||||
str(deployment.get("name", "-") or "-"),
|
||||
_extract_deployment_url(deployment),
|
||||
)
|
||||
for deployment in deployments
|
||||
]
|
||||
widths = [
|
||||
max(len(headers[index]), *(len(row[index]) for row in rows))
|
||||
for index in range(len(headers))
|
||||
]
|
||||
|
||||
def format_row(row: Sequence[str]) -> str:
|
||||
return " ".join(value.ljust(widths[index]) for index, value in enumerate(row))
|
||||
|
||||
lines = [format_row(headers), format_row(tuple("-" * width for width in widths))]
|
||||
lines.extend(format_row(row) for row in rows)
|
||||
return "\n".join(lines)
|
||||
|
||||
@@ -9,6 +9,7 @@ from pathlib import Path
|
||||
|
||||
from click.testing import CliRunner
|
||||
|
||||
import langgraph_cli.cli as cli_module
|
||||
from langgraph_cli.cli import cli, prepare_args_and_stdin
|
||||
from langgraph_cli.config import Config, _get_pip_cleanup_lines, validate_config
|
||||
from langgraph_cli.docker import DEFAULT_POSTGRES_URI, DockerCapabilities, Version
|
||||
@@ -287,6 +288,238 @@ def test_version_option() -> None:
|
||||
)
|
||||
|
||||
|
||||
def test_top_level_help_shows_deploy_subcommands() -> None:
|
||||
runner = CliRunner()
|
||||
|
||||
result = runner.invoke(cli, ["--help"])
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert "deploy" in result.output
|
||||
assert "deploy list" in result.output
|
||||
assert "deploy delete" in result.output
|
||||
assert "[Beta] List LangSmith Deployments." in result.output
|
||||
|
||||
|
||||
def test_top_level_help_truncates_command_descriptions_to_single_line() -> None:
|
||||
runner = CliRunner()
|
||||
|
||||
result = runner.invoke(cli, ["--help"])
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
lines = result.output.splitlines()
|
||||
deploy_line = next(line for line in lines if line.strip().startswith("deploy"))
|
||||
deploy_list_line = next(
|
||||
line for line in lines if line.strip().startswith("deploy list")
|
||||
)
|
||||
|
||||
assert not lines[lines.index(deploy_line) + 1].startswith(" ")
|
||||
assert "..." in deploy_line
|
||||
assert "[Beta] List LangSmith Deployments." in deploy_list_line
|
||||
|
||||
|
||||
def test_deploy_list_command(monkeypatch) -> None:
|
||||
runner = CliRunner()
|
||||
captured: dict[str, str] = {}
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, host_url: str, api_key: str, tenant_id: str | None = None):
|
||||
captured["host_url"] = host_url
|
||||
captured["api_key"] = api_key
|
||||
captured["tenant_id"] = tenant_id or ""
|
||||
|
||||
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"},
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
monkeypatch.setattr(cli_module, "HostBackendClient", FakeClient)
|
||||
|
||||
result = runner.invoke(
|
||||
cli,
|
||||
[
|
||||
"deploy",
|
||||
"list",
|
||||
"--api-key",
|
||||
"test-key",
|
||||
"--host-url",
|
||||
"https://api.example.com",
|
||||
"--name-contains",
|
||||
"alp",
|
||||
],
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert captured == {
|
||||
"host_url": "https://api.example.com",
|
||||
"api_key": "test-key",
|
||||
"tenant_id": "",
|
||||
"name_contains": "alp",
|
||||
}
|
||||
assert "Deployment ID" in result.output
|
||||
assert "Deployment Name" in result.output
|
||||
assert "Deployment URL" in result.output
|
||||
assert "dep-123" in result.output
|
||||
assert "https://beta.example.com" in result.output
|
||||
|
||||
|
||||
def test_deploy_list_command_no_results(monkeypatch) -> None:
|
||||
runner = CliRunner()
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, host_url: str, api_key: str, tenant_id: str | None = None):
|
||||
pass
|
||||
|
||||
def list_deployments(self, name_contains: str = ""):
|
||||
return {"resources": []}
|
||||
|
||||
monkeypatch.setattr(cli_module, "HostBackendClient", FakeClient)
|
||||
|
||||
result = runner.invoke(
|
||||
cli,
|
||||
[
|
||||
"deploy",
|
||||
"list",
|
||||
"--api-key",
|
||||
"test-key",
|
||||
"--host-url",
|
||||
"https://api.example.com",
|
||||
],
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert result.output.strip() == "No deployments found."
|
||||
|
||||
|
||||
def test_deploy_delete_command(monkeypatch) -> None:
|
||||
runner = CliRunner()
|
||||
captured: dict[str, str] = {}
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, host_url: str, api_key: str, tenant_id: str | None = None):
|
||||
captured["host_url"] = host_url
|
||||
captured["api_key"] = api_key
|
||||
captured["tenant_id"] = tenant_id or ""
|
||||
|
||||
def delete_deployment(self, deployment_id: str):
|
||||
captured["deployment_id"] = deployment_id
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(cli_module, "HostBackendClient", FakeClient)
|
||||
|
||||
result = runner.invoke(
|
||||
cli,
|
||||
[
|
||||
"deploy",
|
||||
"delete",
|
||||
"--api-key",
|
||||
"test-key",
|
||||
"--host-url",
|
||||
"https://api.example.com",
|
||||
"dep-123",
|
||||
],
|
||||
input="y\n",
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert captured == {
|
||||
"host_url": "https://api.example.com",
|
||||
"api_key": "test-key",
|
||||
"tenant_id": "",
|
||||
"deployment_id": "dep-123",
|
||||
}
|
||||
assert (
|
||||
"Are you sure you want to delete deployment ID dep-123? (Y/n):" in result.output
|
||||
)
|
||||
assert result.output.strip().endswith("Deleted deployment dep-123.")
|
||||
|
||||
|
||||
def test_deploy_delete_command_cancelled(monkeypatch) -> None:
|
||||
runner = CliRunner()
|
||||
deleted = False
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, host_url: str, api_key: str, tenant_id: str | None = None):
|
||||
pass
|
||||
|
||||
def delete_deployment(self, deployment_id: str):
|
||||
nonlocal deleted
|
||||
deleted = True
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(cli_module, "HostBackendClient", FakeClient)
|
||||
|
||||
result = runner.invoke(
|
||||
cli,
|
||||
[
|
||||
"deploy",
|
||||
"delete",
|
||||
"--api-key",
|
||||
"test-key",
|
||||
"--host-url",
|
||||
"https://api.example.com",
|
||||
"dep-123",
|
||||
],
|
||||
input="n\n",
|
||||
)
|
||||
|
||||
assert result.exit_code == 1, result.output
|
||||
assert not deleted
|
||||
assert "Aborted!" in result.output
|
||||
|
||||
|
||||
def test_deploy_delete_command_force(monkeypatch) -> None:
|
||||
runner = CliRunner()
|
||||
captured: dict[str, str] = {}
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, host_url: str, api_key: str, tenant_id: str | None = None):
|
||||
captured["host_url"] = host_url
|
||||
captured["api_key"] = api_key
|
||||
captured["tenant_id"] = tenant_id or ""
|
||||
|
||||
def delete_deployment(self, deployment_id: str):
|
||||
captured["deployment_id"] = deployment_id
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(cli_module, "HostBackendClient", FakeClient)
|
||||
|
||||
result = runner.invoke(
|
||||
cli,
|
||||
[
|
||||
"deploy",
|
||||
"delete",
|
||||
"--force",
|
||||
"--api-key",
|
||||
"test-key",
|
||||
"--host-url",
|
||||
"https://api.example.com",
|
||||
"dep-123",
|
||||
],
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert "Are you sure you want to delete deployment ID dep-123?" not in result.output
|
||||
assert captured == {
|
||||
"host_url": "https://api.example.com",
|
||||
"api_key": "test-key",
|
||||
"tenant_id": "",
|
||||
"deployment_id": "dep-123",
|
||||
}
|
||||
assert result.output.strip() == "Deleted deployment dep-123."
|
||||
|
||||
|
||||
def test_dockerfile_command_basic() -> None:
|
||||
"""Test the 'dockerfile' command with basic configuration."""
|
||||
runner = CliRunner()
|
||||
|
||||
@@ -1784,7 +1784,9 @@ def test_config_to_compose_distributed_mode():
|
||||
# Executor service is present with correct base image
|
||||
assert "langgraph-executor:" in actual_compose_stdin
|
||||
assert "FROM langchain/langgraph-executor:3.11" in actual_compose_stdin
|
||||
assert 'entrypoint: ["sh", "/storage/executor_entrypoint.sh"]' in actual_compose_stdin
|
||||
assert (
|
||||
'entrypoint: ["sh", "/storage/executor_entrypoint.sh"]' in actual_compose_stdin
|
||||
)
|
||||
|
||||
# Executor has required environment variables
|
||||
assert "EXECUTOR_GRPC_PORT:" in actual_compose_stdin
|
||||
|
||||
@@ -135,6 +135,28 @@ def test_list_deployments(client):
|
||||
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}
|
||||
@@ -160,3 +182,41 @@ def test_list_revisions(client):
|
||||
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,
|
||||
)
|
||||
result = c.get_deploy_logs("proj-1", {"limit": 10})
|
||||
assert result == {"logs": [{"message": "running"}]}
|
||||
|
||||
|
||||
def test_get_deploy_logs_specific_revision():
|
||||
def handler(req: httpx.Request) -> httpx.Response:
|
||||
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,
|
||||
)
|
||||
result = c.get_deploy_logs("proj-1", {"limit": 10}, revision_id="rev-2")
|
||||
assert result == {"logs": []}
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
from langgraph_cli.helpers import format_log_entry, format_timestamp, level_fg
|
||||
|
||||
|
||||
class TestFormatTimestamp:
|
||||
def test_epoch_ms(self):
|
||||
assert format_timestamp(1773119644012) == "2026-03-10 05:14:04"
|
||||
|
||||
def test_string_passthrough(self):
|
||||
assert format_timestamp("2026-03-08T00:00:00Z") == "2026-03-08T00:00:00Z"
|
||||
|
||||
def test_empty(self):
|
||||
assert format_timestamp("") == ""
|
||||
|
||||
def test_none(self):
|
||||
assert format_timestamp(None) == ""
|
||||
|
||||
|
||||
class TestFormatLogEntry:
|
||||
def test_full_entry_epoch(self):
|
||||
entry = {"timestamp": 1773119644012, "level": "ERROR", "message": "boom"}
|
||||
result = format_log_entry(entry)
|
||||
assert result == "[2026-03-10 05:14:04] [ERROR] boom"
|
||||
|
||||
def test_full_entry_string(self):
|
||||
entry = {
|
||||
"timestamp": "2026-03-08T12:00:00Z",
|
||||
"level": "ERROR",
|
||||
"message": "boom",
|
||||
}
|
||||
assert format_log_entry(entry) == "[2026-03-08T12:00:00Z] [ERROR] boom"
|
||||
|
||||
def test_no_level(self):
|
||||
entry = {"timestamp": "2026-03-08T12:00:00Z", "message": "hello"}
|
||||
assert format_log_entry(entry) == "[2026-03-08T12:00:00Z] hello"
|
||||
|
||||
def test_no_timestamp(self):
|
||||
entry = {"message": "bare message"}
|
||||
assert format_log_entry(entry) == "bare message"
|
||||
|
||||
def test_empty_entry(self):
|
||||
assert format_log_entry({}) == ""
|
||||
|
||||
|
||||
class TestLevelFg:
|
||||
def test_error(self):
|
||||
assert level_fg("ERROR") == "red"
|
||||
|
||||
def test_error_lowercase(self):
|
||||
assert level_fg("error") == "red"
|
||||
|
||||
def test_warning(self):
|
||||
assert level_fg("WARNING") == "yellow"
|
||||
|
||||
def test_info_returns_none(self):
|
||||
assert level_fg("INFO") is None
|
||||
|
||||
def test_empty_returns_none(self):
|
||||
assert level_fg("") is None
|
||||
@@ -1,6 +1,11 @@
|
||||
from unittest.mock import patch
|
||||
|
||||
from langgraph_cli.util import clean_empty_lines, warn_non_wolfi_distro
|
||||
from langgraph_cli.util import (
|
||||
_extract_deployment_url,
|
||||
clean_empty_lines,
|
||||
format_deployments_table,
|
||||
warn_non_wolfi_distro,
|
||||
)
|
||||
|
||||
|
||||
def test_clean_empty_lines():
|
||||
@@ -186,3 +191,36 @@ def test_warn_non_wolfi_distro_does_not_modify_config():
|
||||
warn_non_wolfi_distro(config_copy)
|
||||
|
||||
assert config_copy == original_config # Config should remain unchanged
|
||||
|
||||
|
||||
def test_extract_deployment_url_uses_custom_url():
|
||||
deployment = {"source_config": {"custom_url": "https://example.com/custom"}}
|
||||
assert _extract_deployment_url(deployment) == "https://example.com/custom"
|
||||
|
||||
|
||||
def test_extract_deployment_url_defaults_to_dash():
|
||||
assert _extract_deployment_url({"id": "dep-123"}) == "-"
|
||||
|
||||
|
||||
def test_format_deployments_table():
|
||||
output = format_deployments_table(
|
||||
[
|
||||
{
|
||||
"id": "dep-123",
|
||||
"name": "alpha",
|
||||
"source_config": {"custom_url": "https://alpha.example.com"},
|
||||
},
|
||||
{
|
||||
"id": "dep-456",
|
||||
"name": "beta",
|
||||
"url": "https://beta.example.com",
|
||||
},
|
||||
]
|
||||
)
|
||||
assert "Deployment ID" in output
|
||||
assert "Deployment Name" in output
|
||||
assert "Deployment URL" in output
|
||||
assert "dep-123" in output
|
||||
assert "alpha" in output
|
||||
assert "https://alpha.example.com" in output
|
||||
assert "dep-456" in output
|
||||
|
||||
@@ -1,90 +0,0 @@
|
||||
# LangGraph 1.1.0 Release Notes
|
||||
|
||||
## Type-Safe Streaming & Invoke
|
||||
|
||||
LangGraph 1.1 introduces `version="v2"` — a new opt-in streaming format that brings full type safety to `stream()`, `astream()`, `invoke()`, and `ainvoke()`.
|
||||
|
||||
### What's changing
|
||||
|
||||
**v1 (default, unchanged):** `stream()` yields bare tuples like `(stream_mode, data)` or just `data`. `invoke()` returns a plain `dict`. Interrupts are mixed into the output dict under `"__interrupt__"`.
|
||||
|
||||
**v2 (opt-in):** `stream()` yields strongly-typed `StreamPart` dicts with `type`, `ns`, `data`, and (for values) `interrupts` fields. `invoke()` returns a `GraphOutput` object with `.value` and `.interrupts` attributes. When your state schema is a Pydantic model or dataclass, outputs are automatically coerced to the correct type.
|
||||
|
||||
### `invoke()` / `ainvoke()` with `version="v2"`
|
||||
|
||||
```python
|
||||
from langgraph.types import GraphOutput
|
||||
|
||||
result = graph.invoke({"input": "hello"}, version="v2")
|
||||
|
||||
# result is a GraphOutput, not a dict
|
||||
assert isinstance(result, GraphOutput)
|
||||
result.value # your output — dict, Pydantic model, or dataclass
|
||||
result.interrupts # tuple[Interrupt, ...], empty if none occurred
|
||||
```
|
||||
|
||||
With a non-`"values"` stream mode, `invoke(..., stream_mode="updates", version="v2")` returns `list[StreamPart]` instead of `list[tuple]`.
|
||||
|
||||
### `stream()` / `astream()` with `version="v2"`
|
||||
|
||||
```python
|
||||
for part in graph.stream({"input": "hello"}, version="v2"):
|
||||
if part["type"] == "values":
|
||||
part["data"] # OutputT — full state
|
||||
part["interrupts"] # tuple[Interrupt, ...]
|
||||
elif part["type"] == "updates":
|
||||
part["data"] # dict[str, Any]
|
||||
elif part["type"] == "messages":
|
||||
part["data"] # tuple[BaseMessage, dict]
|
||||
elif part["type"] == "custom":
|
||||
part["data"] # Any
|
||||
elif part["type"] == "tasks":
|
||||
part["data"] # TaskPayload | TaskResultPayload
|
||||
elif part["type"] == "debug":
|
||||
part["data"] # DebugPayload
|
||||
```
|
||||
|
||||
Each stream mode has its own `TypedDict` — `ValuesStreamPart`, `UpdatesStreamPart`, `MessagesStreamPart`, `CustomStreamPart`, `CheckpointStreamPart`, `TasksStreamPart`, `DebugStreamPart` — all importable from `langgraph.types`. The union type `StreamPart` is a discriminated union on `part["type"]`, enabling full type narrowing in editors and type checkers.
|
||||
|
||||
### Pydantic & dataclass output coercion
|
||||
|
||||
When your graph's state schema is a Pydantic model or dataclass, `version="v2"` automatically coerces outputs to the declared type:
|
||||
|
||||
```python
|
||||
from pydantic import BaseModel
|
||||
|
||||
class MyState(BaseModel):
|
||||
answer: str
|
||||
count: int
|
||||
|
||||
graph = StateGraph(MyState)
|
||||
# ... build graph ...
|
||||
compiled = graph.compile()
|
||||
|
||||
result = compiled.invoke({"answer": "", "count": 0}, version="v2")
|
||||
assert isinstance(result.value, MyState) # not a dict!
|
||||
```
|
||||
|
||||
### Backward compatibility
|
||||
|
||||
- **Default is still `version="v1"`** — existing code works without changes.
|
||||
- To make migration easier, `GraphOutput` supports old-style best-effort access to graph values and interrupts. Dict-style access (`result["key"]`, `"key" in result`, `result["__interrupt__"]`) still works and delegates to `result.value` / `result.interrupts` under the hood. However, this is **deprecated** and emits a `LangGraphDeprecatedSinceV11` warning. It will be removed in v3.0 — migrate to `result.value` and `result.interrupts` at your convenience.
|
||||
|
||||
```python
|
||||
result = graph.invoke({"input": "hello"}, version="v2")
|
||||
|
||||
# Old style — still works, but deprecated
|
||||
result["input"] # delegates to result.value["input"]
|
||||
result["__interrupt__"] # delegates to result.interrupts
|
||||
"input" in result # delegates to "input" in result.value
|
||||
|
||||
# New style — preferred
|
||||
result.value["input"]
|
||||
result.interrupts
|
||||
```
|
||||
|
||||
## Migration Guide
|
||||
|
||||
1. **No action required** — `version="v1"` remains the default. All existing code continues to work.
|
||||
2. **Adopt v2 incrementally** — Add `version="v2"` to individual `invoke()`/`stream()` calls to get typed outputs.
|
||||
3. **Use typed imports** — Import `GraphOutput`, `StreamPart`, and individual part types from `langgraph.types` for type-safe code.
|
||||
@@ -647,13 +647,24 @@ class PregelLoop:
|
||||
# writes so that interrupt() calls re-fire instead of returning
|
||||
# stale values. But if we're actively resuming, keep them —
|
||||
# multi-interrupt scenarios need previously resolved values preserved.
|
||||
# We check two conditions because resume signals arrive differently:
|
||||
# - Command(resume=...): the outer graph receives resume via input
|
||||
# - CONFIG_KEY_RESUMING: child subgraphs receive it via config from
|
||||
# the parent (their input is a Send arg, not a Command)
|
||||
if self.is_replaying and not (
|
||||
(input_is_command and cast(Command, self.input).resume is not None)
|
||||
or configurable.get(CONFIG_KEY_RESUMING, False)
|
||||
if self.is_replaying and (
|
||||
# Time-travel to a subgraph checkpoint: the parent sets
|
||||
# RESUMING=True (it can't distinguish time-travel from resume),
|
||||
# so we check if this subgraph's own ns is in checkpoint_map.
|
||||
# Normally the map only has ancestor entries (_algo.py); the
|
||||
# subgraph's own entry only appears via get_state(subgraphs=True).
|
||||
(
|
||||
self.is_nested
|
||||
and configurable.get(CONFIG_KEY_CHECKPOINT_NS, "")
|
||||
in configurable.get(CONFIG_KEY_CHECKPOINT_MAP, {})
|
||||
)
|
||||
or not (
|
||||
# Outer graph: resume arrives as Command(resume=...)
|
||||
(input_is_command and cast(Command, self.input).resume is not None)
|
||||
# Subgraphs: resume arrives via config flag from parent
|
||||
# (subgraph input is a Send arg, not a Command)
|
||||
or configurable.get(CONFIG_KEY_RESUMING, False)
|
||||
)
|
||||
):
|
||||
self.checkpoint_pending_writes = [
|
||||
w for w in self.checkpoint_pending_writes if w[1] != RESUME
|
||||
@@ -1127,9 +1138,15 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
def __enter__(self) -> Self:
|
||||
if not self.checkpointer:
|
||||
saved = None
|
||||
elif self.is_nested and (
|
||||
replay_state := self.config[CONF].get(CONFIG_KEY_REPLAY_STATE)
|
||||
):
|
||||
elif self.checkpoint_config[CONF].get(CONFIG_KEY_CHECKPOINT_ID):
|
||||
# Explicit checkpoint_id requested — fetch that exact checkpoint.
|
||||
# This covers both normal replay and subgraphs resolved via
|
||||
# checkpoint_map during time-travel.
|
||||
saved = self.checkpointer.get_tuple(self.checkpoint_config)
|
||||
elif replay_state := self.config[CONF].get(CONFIG_KEY_REPLAY_STATE):
|
||||
# Subgraph replay: the parent is replaying and passed us a
|
||||
# replay_state with its checkpoint_id. Look up our checkpoint
|
||||
# from the parent's checkpoint_map instead of fetching latest.
|
||||
saved = replay_state.get_checkpoint(
|
||||
self.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, ""),
|
||||
self.checkpointer,
|
||||
@@ -1141,9 +1158,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
self.config[CONF].pop(CONFIG_KEY_RESUMING, None)
|
||||
else:
|
||||
# Normal case: fetch the most recent checkpoint for this
|
||||
# graph/thread. If a specific checkpoint_id is in the config,
|
||||
# fetch that exact checkpoint; otherwise fetch the latest one.
|
||||
# Returns None on first invocation (no checkpoints exist yet).
|
||||
# graph/thread. Returns None on first invocation.
|
||||
saved = self.checkpointer.get_tuple(self.checkpoint_config)
|
||||
|
||||
if saved is None:
|
||||
@@ -1322,9 +1337,15 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
async def __aenter__(self) -> Self:
|
||||
if not self.checkpointer:
|
||||
saved = None
|
||||
elif self.is_nested and (
|
||||
replay_state := self.config[CONF].get(CONFIG_KEY_REPLAY_STATE)
|
||||
):
|
||||
elif self.checkpoint_config[CONF].get(CONFIG_KEY_CHECKPOINT_ID):
|
||||
# Explicit checkpoint_id requested — fetch that exact checkpoint.
|
||||
# This covers both normal replay and subgraphs resolved via
|
||||
# checkpoint_map during time-travel.
|
||||
saved = await self.checkpointer.aget_tuple(self.checkpoint_config)
|
||||
elif replay_state := self.config[CONF].get(CONFIG_KEY_REPLAY_STATE):
|
||||
# Subgraph replay: the parent is replaying and passed us a
|
||||
# replay_state with its checkpoint_id. Look up our checkpoint
|
||||
# from the parent's checkpoint_map instead of fetching latest.
|
||||
saved = await replay_state.aget_checkpoint(
|
||||
self.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, ""),
|
||||
self.checkpointer,
|
||||
@@ -1336,9 +1357,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
self.config[CONF].pop(CONFIG_KEY_RESUMING, None)
|
||||
else:
|
||||
# Normal case: fetch the most recent checkpoint for this
|
||||
# graph/thread. If a specific checkpoint_id is in the config,
|
||||
# fetch that exact checkpoint; otherwise fetch the latest one.
|
||||
# Returns None on first invocation (no checkpoints exist yet).
|
||||
# graph/thread. Returns None on first invocation.
|
||||
saved = await self.checkpointer.aget_tuple(self.checkpoint_config)
|
||||
|
||||
if saved is None:
|
||||
|
||||
@@ -25,7 +25,7 @@ except ImportError:
|
||||
_StreamingCallbackHandler = object # type: ignore
|
||||
|
||||
T = TypeVar("T")
|
||||
Meta = tuple[tuple[str, ...], dict[str, Any]]
|
||||
Meta = tuple[tuple[str, ...], dict[str, Any] | None]
|
||||
|
||||
|
||||
def _state_values(obj: Any) -> Sequence[Any]:
|
||||
@@ -56,6 +56,7 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
|
||||
subgraphs: bool,
|
||||
*,
|
||||
parent_ns: tuple[str, ...] | None = None,
|
||||
dedupe_metadata: bool = False,
|
||||
) -> None:
|
||||
"""Configure the handler to stream messages from LLMs and nodes.
|
||||
|
||||
@@ -84,8 +85,10 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
|
||||
self.stream = stream
|
||||
self.subgraphs = subgraphs
|
||||
self.metadata: dict[UUID, Meta] = {}
|
||||
self.emitted_metadata: set[UUID] = set()
|
||||
self.seen: set[int | str] = set()
|
||||
self.parent_ns = parent_ns
|
||||
self.dedupe_metadata = dedupe_metadata
|
||||
|
||||
def _emit(self, meta: Meta, message: BaseMessage, *, dedupe: bool = False) -> None:
|
||||
if dedupe and message.id in self.seen:
|
||||
@@ -155,6 +158,10 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
|
||||
if not isinstance(chunk, ChatGenerationChunk):
|
||||
return
|
||||
if meta := self.metadata.get(run_id):
|
||||
if self.dedupe_metadata and run_id in self.emitted_metadata:
|
||||
meta = (meta[0], None)
|
||||
else:
|
||||
self.emitted_metadata.add(run_id)
|
||||
self._emit(meta, chunk.message)
|
||||
|
||||
def on_llm_end(
|
||||
@@ -170,6 +177,7 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
|
||||
gen = response.generations[0][0]
|
||||
if isinstance(gen, ChatGeneration):
|
||||
self._emit(meta, gen.message, dedupe=True)
|
||||
self.emitted_metadata.discard(run_id)
|
||||
self.metadata.pop(run_id, None)
|
||||
|
||||
def on_llm_error(
|
||||
@@ -180,6 +188,7 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
|
||||
parent_run_id: UUID | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
self.emitted_metadata.discard(run_id)
|
||||
self.metadata.pop(run_id, None)
|
||||
|
||||
def on_chain_start(
|
||||
|
||||
@@ -2456,7 +2456,7 @@ class Pregel(
|
||||
debug: bool | None = None,
|
||||
version: Literal["v2"],
|
||||
**kwargs: Unpack[DeprecatedKwargs],
|
||||
) -> Iterator[StreamPart[OutputT, StateT]]: ...
|
||||
) -> Iterator[StreamPart[StateT, OutputT]]: ...
|
||||
|
||||
@overload
|
||||
def stream(
|
||||
@@ -2614,6 +2614,7 @@ class Pregel(
|
||||
stream.put,
|
||||
subgraphs,
|
||||
parent_ns=tuple(ns_.split(NS_SEP)) if ns_ else None,
|
||||
dedupe_metadata="compact" in stream_modes,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -2787,7 +2788,7 @@ class Pregel(
|
||||
debug: bool | None = None,
|
||||
version: Literal["v2"],
|
||||
**kwargs: Unpack[DeprecatedKwargs],
|
||||
) -> AsyncIterator[StreamPart[OutputT, StateT]]: ...
|
||||
) -> AsyncIterator[StreamPart[StateT, OutputT]]: ...
|
||||
|
||||
@overload
|
||||
def astream(
|
||||
@@ -2965,6 +2966,7 @@ class Pregel(
|
||||
stream_put,
|
||||
subgraphs,
|
||||
parent_ns=tuple(ns_.split(NS_SEP)) if ns_ else None,
|
||||
dedupe_metadata="compact" in stream_modes,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -3194,7 +3196,7 @@ class Pregel(
|
||||
durability: Durability | None = None,
|
||||
version: Literal["v2"],
|
||||
**kwargs: Any,
|
||||
) -> list[StreamPart[OutputT, StateT]]: ...
|
||||
) -> list[StreamPart[StateT, OutputT]]: ...
|
||||
|
||||
@overload
|
||||
def invoke(
|
||||
@@ -3364,7 +3366,7 @@ class Pregel(
|
||||
durability: Durability | None = None,
|
||||
version: Literal["v2"],
|
||||
**kwargs: Any,
|
||||
) -> list[StreamPart[OutputT, StateT]]: ...
|
||||
) -> list[StreamPart[StateT, OutputT]]: ...
|
||||
|
||||
@overload
|
||||
async def ainvoke(
|
||||
|
||||
@@ -117,7 +117,7 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
subgraphs: bool = False,
|
||||
version: Literal["v2"],
|
||||
) -> Iterator[StreamPart[OutputT, StateT]]: ...
|
||||
) -> Iterator[StreamPart[StateT, OutputT]]: ...
|
||||
|
||||
@overload
|
||||
@abstractmethod
|
||||
@@ -161,7 +161,7 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
subgraphs: bool = False,
|
||||
version: Literal["v2"],
|
||||
) -> AsyncIterator[StreamPart[OutputT, StateT]]: ...
|
||||
) -> AsyncIterator[StreamPart[StateT, OutputT]]: ...
|
||||
|
||||
@overload
|
||||
@abstractmethod
|
||||
|
||||
@@ -31,6 +31,7 @@ from langgraph_sdk.client import (
|
||||
)
|
||||
from langgraph_sdk.schema import (
|
||||
Checkpoint,
|
||||
Context,
|
||||
QueryParamTypes,
|
||||
ThreadState,
|
||||
)
|
||||
@@ -108,6 +109,45 @@ class RemoteException(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def _restore_message_metadata(
|
||||
data: Any, metadata_by_message_id: dict[str, dict[str, Any]]
|
||||
) -> Any:
|
||||
"""Restore deduplicated message metadata using the message id as cache key."""
|
||||
if not (isinstance(data, list) and len(data) == 2):
|
||||
return data
|
||||
message, metadata = data
|
||||
if not isinstance(message, dict):
|
||||
return data
|
||||
message_id = message.get("id")
|
||||
if isinstance(message_id, str):
|
||||
if isinstance(metadata, dict):
|
||||
metadata_by_message_id[message_id] = metadata
|
||||
else:
|
||||
metadata = metadata_by_message_id.get(message_id)
|
||||
return (message, metadata)
|
||||
|
||||
|
||||
def _merge_values_patch(
|
||||
ns: tuple[str, ...],
|
||||
mode: str,
|
||||
data: Any,
|
||||
values_by_ns: dict[tuple[str, ...], dict[str, Any]],
|
||||
) -> tuple[str, Any]:
|
||||
"""Merge `values-patch` events back into full values snapshots."""
|
||||
if mode != "values-patch" or not isinstance(data, dict):
|
||||
return mode, data
|
||||
values = data.get("values")
|
||||
if not isinstance(values, dict):
|
||||
return "values", values if values is not None else {}
|
||||
merged = dict(values_by_ns.get(ns, {}))
|
||||
merged.update(values)
|
||||
for key in data.get("deleted_keys", ()):
|
||||
if isinstance(key, str):
|
||||
merged.pop(key, None)
|
||||
values_by_ns[ns] = merged
|
||||
return "values", merged
|
||||
|
||||
|
||||
class RemoteGraph(PregelProtocol):
|
||||
"""The `RemoteGraph` class is a client implementation for calling remote
|
||||
APIs that implement the LangGraph Server API specification.
|
||||
@@ -691,6 +731,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
@@ -707,6 +748,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
@@ -722,6 +764,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
@@ -734,7 +777,7 @@ class RemoteGraph(PregelProtocol):
|
||||
"""Create a run and stream the results.
|
||||
|
||||
This method calls `POST /threads/{thread_id}/runs/stream` if a `thread_id`
|
||||
is speciffed in the `configurable` field of the config or
|
||||
is specified in the `configurable` field of the config or
|
||||
`POST /runs/stream` otherwise.
|
||||
|
||||
Args:
|
||||
@@ -762,6 +805,8 @@ class RemoteGraph(PregelProtocol):
|
||||
else:
|
||||
command = None
|
||||
thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None)
|
||||
message_metadata_by_id: dict[str, dict[str, Any]] = {}
|
||||
values_by_ns: dict[tuple[str, ...], dict[str, Any]] = {}
|
||||
|
||||
for chunk in sync_client.runs.stream(
|
||||
thread_id=thread_id,
|
||||
@@ -769,6 +814,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input=input,
|
||||
command=command,
|
||||
config=sanitized_config,
|
||||
context=context,
|
||||
stream_mode=stream_modes,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
@@ -793,6 +839,11 @@ class RemoteGraph(PregelProtocol):
|
||||
if caller_ns := (config or {}).get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_NS):
|
||||
caller_ns = tuple(caller_ns.split(NS_SEP))
|
||||
ns = caller_ns + ns
|
||||
mode, data = _merge_values_patch(ns, mode, chunk.data, values_by_ns)
|
||||
if mode != chunk.event:
|
||||
chunk = chunk._replace(data=data)
|
||||
elif data is not chunk.data:
|
||||
chunk = chunk._replace(data=data)
|
||||
# stream to parent stream
|
||||
if stream is not None and mode in stream.modes:
|
||||
stream((ns, mode, chunk.data))
|
||||
@@ -810,7 +861,9 @@ class RemoteGraph(PregelProtocol):
|
||||
continue
|
||||
|
||||
if chunk.event.startswith("messages"):
|
||||
chunk = chunk._replace(data=tuple(chunk.data))
|
||||
chunk = chunk._replace(
|
||||
data=_restore_message_metadata(chunk.data, message_metadata_by_id)
|
||||
)
|
||||
|
||||
# emit chunk
|
||||
if version == "v2":
|
||||
@@ -822,11 +875,6 @@ class RemoteGraph(PregelProtocol):
|
||||
)
|
||||
yield {"type": mode, "ns": ns, "data": chunk.data, "interrupts": ints}
|
||||
elif subgraphs:
|
||||
if NS_SEP in chunk.event:
|
||||
mode, ns_ = chunk.event.split(NS_SEP, 1)
|
||||
ns = tuple(ns_.split(NS_SEP))
|
||||
else:
|
||||
mode, ns = chunk.event, ()
|
||||
if req_single:
|
||||
yield ns, chunk.data
|
||||
else:
|
||||
@@ -842,6 +890,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
@@ -858,6 +907,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
@@ -873,6 +923,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
@@ -885,7 +936,7 @@ class RemoteGraph(PregelProtocol):
|
||||
"""Create a run and stream the results.
|
||||
|
||||
This method calls `POST /threads/{thread_id}/runs/stream` if a `thread_id`
|
||||
is speciffed in the `configurable` field of the config or
|
||||
is specified in the `configurable` field of the config or
|
||||
`POST /runs/stream` otherwise.
|
||||
|
||||
Args:
|
||||
@@ -913,6 +964,8 @@ class RemoteGraph(PregelProtocol):
|
||||
else:
|
||||
command = None
|
||||
thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None)
|
||||
message_metadata_by_id: dict[str, dict[str, Any]] = {}
|
||||
values_by_ns: dict[tuple[str, ...], dict[str, Any]] = {}
|
||||
|
||||
async for chunk in client.runs.stream(
|
||||
thread_id=thread_id,
|
||||
@@ -920,6 +973,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input=input,
|
||||
command=command,
|
||||
config=sanitized_config,
|
||||
context=context,
|
||||
stream_mode=stream_modes,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
@@ -944,6 +998,11 @@ class RemoteGraph(PregelProtocol):
|
||||
if caller_ns := (config or {}).get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_NS):
|
||||
caller_ns = tuple(caller_ns.split(NS_SEP))
|
||||
ns = caller_ns + ns
|
||||
mode, data = _merge_values_patch(ns, mode, chunk.data, values_by_ns)
|
||||
if mode != chunk.event:
|
||||
chunk = chunk._replace(data=data)
|
||||
elif data is not chunk.data:
|
||||
chunk = chunk._replace(data=data)
|
||||
# stream to parent stream
|
||||
if stream is not None and mode in stream.modes:
|
||||
stream((ns, mode, chunk.data))
|
||||
@@ -961,7 +1020,9 @@ class RemoteGraph(PregelProtocol):
|
||||
continue
|
||||
|
||||
if chunk.event.startswith("messages"):
|
||||
chunk = chunk._replace(data=tuple(chunk.data))
|
||||
chunk = chunk._replace(
|
||||
data=_restore_message_metadata(chunk.data, message_metadata_by_id)
|
||||
)
|
||||
|
||||
# emit chunk
|
||||
if version == "v2":
|
||||
@@ -973,11 +1034,6 @@ class RemoteGraph(PregelProtocol):
|
||||
)
|
||||
yield {"type": mode, "ns": ns, "data": chunk.data, "interrupts": ints}
|
||||
elif subgraphs:
|
||||
if NS_SEP in chunk.event:
|
||||
mode, ns_ = chunk.event.split(NS_SEP, 1)
|
||||
ns = tuple(ns_.split(NS_SEP))
|
||||
else:
|
||||
mode, ns = chunk.event, ()
|
||||
if req_single:
|
||||
yield ns, chunk.data
|
||||
else:
|
||||
@@ -1009,6 +1065,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
@@ -1023,6 +1080,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
@@ -1036,6 +1094,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
@@ -1061,6 +1120,7 @@ class RemoteGraph(PregelProtocol):
|
||||
for chunk in self.stream( # type: ignore[misc, call-overload]
|
||||
input,
|
||||
config=config,
|
||||
context=context,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
headers=headers,
|
||||
@@ -1087,6 +1147,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
@@ -1101,6 +1162,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
@@ -1114,6 +1176,7 @@ class RemoteGraph(PregelProtocol):
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
@@ -1139,6 +1202,7 @@ class RemoteGraph(PregelProtocol):
|
||||
async for chunk in self.astream( # type: ignore[misc, call-overload]
|
||||
input,
|
||||
config=config,
|
||||
context=context,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
headers=headers,
|
||||
|
||||
@@ -116,7 +116,14 @@ def ensure_valid_checkpointer(checkpointer: Checkpointer) -> Checkpointer:
|
||||
|
||||
|
||||
StreamMode = Literal[
|
||||
"values", "updates", "checkpoints", "tasks", "debug", "messages", "custom"
|
||||
"values",
|
||||
"updates",
|
||||
"checkpoints",
|
||||
"tasks",
|
||||
"debug",
|
||||
"messages",
|
||||
"custom",
|
||||
"compact",
|
||||
]
|
||||
"""How the stream method should emit outputs.
|
||||
|
||||
@@ -275,13 +282,14 @@ class MessagesStreamPart(TypedDict):
|
||||
"""Stream part emitted for `stream_mode="messages"`.
|
||||
|
||||
`data` is a 2-tuple of `(message, metadata)` where `message` is a
|
||||
`BaseMessage` (e.g. `AIMessageChunk`) and `metadata` is a dict containing
|
||||
keys like `langgraph_step`, `langgraph_node`, `langgraph_triggers`, etc.
|
||||
`BaseMessage` (e.g. `AIMessageChunk`) and `metadata` is either a dict containing
|
||||
keys like `langgraph_step`, `langgraph_node`, `langgraph_triggers`, etc. or
|
||||
`None` for deduplicated follow-up chunks when `stream_mode` includes `"compact"`.
|
||||
"""
|
||||
|
||||
type: Literal["messages"]
|
||||
ns: tuple[str, ...]
|
||||
data: tuple[AnyMessage, dict[str, Any]]
|
||||
data: tuple[AnyMessage, dict[str, Any] | None]
|
||||
|
||||
|
||||
class CustomStreamPart(TypedDict):
|
||||
@@ -335,7 +343,7 @@ StreamPart = TypeAliasType(
|
||||
| CheckpointStreamPart[StateT]
|
||||
| TasksStreamPart
|
||||
| DebugStreamPart[StateT],
|
||||
type_params=(OutputT, StateT),
|
||||
type_params=(StateT, OutputT),
|
||||
)
|
||||
"""A discriminated union of all v2 stream part types.
|
||||
|
||||
@@ -346,7 +354,7 @@ async for part in graph.astream(input, version="v2"):
|
||||
if part["type"] == "values":
|
||||
part["data"] # OutputT — full state (pydantic/dataclass/dict)
|
||||
elif part["type"] == "messages":
|
||||
part["data"] # tuple[BaseMessage, dict] — (message, metadata)
|
||||
part["data"] # tuple[BaseMessage, dict | None] — (message, metadata)
|
||||
elif part["type"] == "custom":
|
||||
part["data"] # Any — user-defined
|
||||
```
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "langgraph"
|
||||
version = "1.1.0"
|
||||
version = "1.1.2"
|
||||
description = "Building stateful, multi-actor applications with LLMs"
|
||||
authors = []
|
||||
requires-python = ">=3.10"
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import re
|
||||
import sys
|
||||
from dataclasses import dataclass
|
||||
from typing import Annotated
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
@@ -10,6 +11,7 @@ from langchain_core.runnables import RunnableConfig
|
||||
from langchain_core.runnables.graph import Edge as DrawableEdge
|
||||
from langchain_core.runnables.graph import Node as DrawableNode
|
||||
from langgraph_sdk.schema import StreamPart
|
||||
from pydantic import BaseModel
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.errors import GraphInterrupt
|
||||
@@ -880,6 +882,80 @@ def test_stream_sanitizes_thread_id():
|
||||
assert not passed_config["configurable"]
|
||||
|
||||
|
||||
def test_stream_restores_messages_and_merges_values_patch():
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.runs.stream.return_value = [
|
||||
StreamPart(
|
||||
event="messages|tools:call_1",
|
||||
data=[
|
||||
{"id": "msg-1", "type": "AIMessageChunk", "content": "hel"},
|
||||
{
|
||||
"langgraph_checkpoint_ns": "tools:call_1",
|
||||
"langgraph_node": "agent",
|
||||
},
|
||||
],
|
||||
),
|
||||
StreamPart(
|
||||
event="messages|tools:call_1",
|
||||
data=[
|
||||
{"id": "msg-1", "type": "AIMessageChunk", "content": "lo"},
|
||||
None,
|
||||
],
|
||||
),
|
||||
StreamPart(
|
||||
event="values|tools:call_1",
|
||||
data={"messages": [{"type": "human", "content": "hi"}], "count": 1},
|
||||
),
|
||||
StreamPart(
|
||||
event="values-patch|tools:call_1",
|
||||
data={"values": {"count": 2}, "deleted_keys": ["messages"]},
|
||||
),
|
||||
]
|
||||
remote_pregel = RemoteGraph("test_graph_id", sync_client=mock_sync_client)
|
||||
|
||||
parts = list(
|
||||
remote_pregel.stream(
|
||||
{"input": "data"},
|
||||
config={"configurable": {"thread_id": "thread_1"}},
|
||||
stream_mode=["messages", "values", "compact"],
|
||||
subgraphs=True,
|
||||
version="v2",
|
||||
)
|
||||
)
|
||||
|
||||
message_parts = [part for part in parts if part["type"] == "messages"]
|
||||
assert message_parts[0]["data"][1] == {
|
||||
"langgraph_checkpoint_ns": "tools:call_1",
|
||||
"langgraph_node": "agent",
|
||||
}
|
||||
assert message_parts[1]["data"][1] == {
|
||||
"langgraph_checkpoint_ns": "tools:call_1",
|
||||
"langgraph_node": "agent",
|
||||
}
|
||||
value_parts = [part for part in parts if part["type"] == "values"]
|
||||
assert value_parts == [
|
||||
{
|
||||
"type": "values",
|
||||
"ns": ("tools:call_1",),
|
||||
"data": {"messages": [{"type": "human", "content": "hi"}], "count": 1},
|
||||
"interrupts": (),
|
||||
},
|
||||
{
|
||||
"type": "values",
|
||||
"ns": ("tools:call_1",),
|
||||
"data": {"count": 2},
|
||||
"interrupts": (),
|
||||
},
|
||||
]
|
||||
_, kwargs = mock_sync_client.runs.stream.call_args
|
||||
assert set(kwargs["stream_mode"]) == {
|
||||
"messages-tuple",
|
||||
"values",
|
||||
"compact",
|
||||
"updates",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_ainvoke():
|
||||
# set up test
|
||||
@@ -908,6 +984,265 @@ async def test_ainvoke():
|
||||
assert result == {"messages": [{"type": "human", "content": "world"}]}
|
||||
|
||||
|
||||
def test_stream_context():
|
||||
"""Test that context is passed through to the SDK client in stream."""
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.runs.stream.return_value = [
|
||||
StreamPart(event="values", data={"chunk": "data1"}),
|
||||
]
|
||||
|
||||
remote_pregel = RemoteGraph(
|
||||
"test_graph_id",
|
||||
sync_client=mock_sync_client,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "thread_1"}}
|
||||
context = {"model_name": "anthropic", "user_id": "123"}
|
||||
stream_parts = list(
|
||||
remote_pregel.stream(
|
||||
{"input": "data"},
|
||||
config,
|
||||
context=context,
|
||||
stream_mode="values",
|
||||
)
|
||||
)
|
||||
|
||||
assert stream_parts == [{"chunk": "data1"}]
|
||||
_, kwargs = mock_sync_client.runs.stream.call_args
|
||||
assert kwargs["context"] == {"model_name": "anthropic", "user_id": "123"}
|
||||
|
||||
|
||||
def test_stream_context_none():
|
||||
"""Test that context defaults to None when not provided."""
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.runs.stream.return_value = [
|
||||
StreamPart(event="values", data={"chunk": "data1"}),
|
||||
]
|
||||
|
||||
remote_pregel = RemoteGraph(
|
||||
"test_graph_id",
|
||||
sync_client=mock_sync_client,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "thread_1"}}
|
||||
list(remote_pregel.stream({"input": "data"}, config, stream_mode="values"))
|
||||
|
||||
_, kwargs = mock_sync_client.runs.stream.call_args
|
||||
assert kwargs["context"] is None
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_astream_context():
|
||||
"""Test that context is passed through to the SDK client in astream."""
|
||||
mock_async_client = MagicMock()
|
||||
async_iter = MagicMock()
|
||||
async_iter.__aiter__.return_value = [
|
||||
StreamPart(event="values", data={"chunk": "data1"}),
|
||||
]
|
||||
mock_async_client.runs.stream.return_value = async_iter
|
||||
|
||||
remote_pregel = RemoteGraph(
|
||||
"test_graph_id",
|
||||
client=mock_async_client,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "thread_1"}}
|
||||
context = {"model_name": "anthropic"}
|
||||
chunks = []
|
||||
async for chunk in remote_pregel.astream(
|
||||
{"input": "data"},
|
||||
config,
|
||||
context=context,
|
||||
stream_mode="values",
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
assert chunks == [{"chunk": "data1"}]
|
||||
_, kwargs = mock_async_client.runs.stream.call_args
|
||||
assert kwargs["context"] == {"model_name": "anthropic"}
|
||||
|
||||
|
||||
def test_invoke_context():
|
||||
"""Test that context is passed through to the SDK client in invoke."""
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.runs.stream.return_value = [
|
||||
StreamPart(event="values", data={"result": "done"}),
|
||||
]
|
||||
|
||||
remote_pregel = RemoteGraph(
|
||||
"test_graph_id",
|
||||
sync_client=mock_sync_client,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "thread_1"}}
|
||||
context = {"model_name": "openai"}
|
||||
result = remote_pregel.invoke({"input": "data"}, config, context=context)
|
||||
|
||||
assert result == {"result": "done"}
|
||||
_, kwargs = mock_sync_client.runs.stream.call_args
|
||||
assert kwargs["context"] == {"model_name": "openai"}
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_ainvoke_context():
|
||||
"""Test that context is passed through to the SDK client in ainvoke."""
|
||||
mock_async_client = MagicMock()
|
||||
async_iter = MagicMock()
|
||||
async_iter.__aiter__.return_value = [
|
||||
StreamPart(event="values", data={"result": "done"}),
|
||||
]
|
||||
mock_async_client.runs.stream.return_value = async_iter
|
||||
|
||||
remote_pregel = RemoteGraph(
|
||||
"test_graph_id",
|
||||
client=mock_async_client,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "thread_1"}}
|
||||
context = {"user_id": "456"}
|
||||
result = await remote_pregel.ainvoke({"input": "data"}, config, context=context)
|
||||
|
||||
assert result == {"result": "done"}
|
||||
_, kwargs = mock_async_client.runs.stream.call_args
|
||||
assert kwargs["context"] == {"user_id": "456"}
|
||||
|
||||
|
||||
def test_stream_context_dataclass():
|
||||
"""Test that a dataclass context is passed through to the SDK client."""
|
||||
|
||||
@dataclass
|
||||
class MyContext:
|
||||
model_name: str
|
||||
user_id: str
|
||||
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.runs.stream.return_value = [
|
||||
StreamPart(event="values", data={"chunk": "data1"}),
|
||||
]
|
||||
|
||||
remote_pregel = RemoteGraph(
|
||||
"test_graph_id",
|
||||
sync_client=mock_sync_client,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "thread_1"}}
|
||||
ctx = MyContext(model_name="anthropic", user_id="123")
|
||||
list(
|
||||
remote_pregel.stream(
|
||||
{"input": "data"}, config, context=ctx, stream_mode="values"
|
||||
)
|
||||
)
|
||||
|
||||
_, kwargs = mock_sync_client.runs.stream.call_args
|
||||
assert kwargs["context"] == ctx
|
||||
|
||||
|
||||
def test_stream_context_base_model():
|
||||
"""Test that a BaseModel context is passed through to the SDK client."""
|
||||
|
||||
class MyContext(BaseModel):
|
||||
model_name: str
|
||||
user_id: str
|
||||
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.runs.stream.return_value = [
|
||||
StreamPart(event="values", data={"chunk": "data1"}),
|
||||
]
|
||||
|
||||
remote_pregel = RemoteGraph(
|
||||
"test_graph_id",
|
||||
sync_client=mock_sync_client,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "thread_1"}}
|
||||
ctx = MyContext(model_name="anthropic", user_id="123")
|
||||
list(
|
||||
remote_pregel.stream(
|
||||
{"input": "data"}, config, context=ctx, stream_mode="values"
|
||||
)
|
||||
)
|
||||
|
||||
_, kwargs = mock_sync_client.runs.stream.call_args
|
||||
assert kwargs["context"] == ctx
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_astream_restores_messages_and_merges_values_patch():
|
||||
mock_async_client = MagicMock()
|
||||
async_iter = MagicMock()
|
||||
async_iter.__aiter__.return_value = [
|
||||
StreamPart(
|
||||
event="messages|tools:call_1",
|
||||
data=[
|
||||
{"id": "msg-1", "type": "AIMessageChunk", "content": "hel"},
|
||||
{
|
||||
"langgraph_checkpoint_ns": "tools:call_1",
|
||||
"langgraph_node": "agent",
|
||||
},
|
||||
],
|
||||
),
|
||||
StreamPart(
|
||||
event="messages|tools:call_1",
|
||||
data=[
|
||||
{"id": "msg-1", "type": "AIMessageChunk", "content": "lo"},
|
||||
None,
|
||||
],
|
||||
),
|
||||
StreamPart(
|
||||
event="values|tools:call_1",
|
||||
data={"messages": [{"type": "human", "content": "hi"}], "count": 1},
|
||||
),
|
||||
StreamPart(
|
||||
event="values-patch|tools:call_1",
|
||||
data={"values": {"count": 2}, "deleted_keys": ["messages"]},
|
||||
),
|
||||
]
|
||||
mock_async_client.runs.stream.return_value = async_iter
|
||||
remote_pregel = RemoteGraph("test_graph_id", client=mock_async_client)
|
||||
|
||||
parts = []
|
||||
async for part in remote_pregel.astream(
|
||||
{"input": "data"},
|
||||
config={"configurable": {"thread_id": "thread_1"}},
|
||||
stream_mode=["messages", "values", "compact"],
|
||||
subgraphs=True,
|
||||
version="v2",
|
||||
):
|
||||
parts.append(part)
|
||||
|
||||
message_parts = [part for part in parts if part["type"] == "messages"]
|
||||
assert message_parts[0]["data"][1] == {
|
||||
"langgraph_checkpoint_ns": "tools:call_1",
|
||||
"langgraph_node": "agent",
|
||||
}
|
||||
assert message_parts[1]["data"][1] == {
|
||||
"langgraph_checkpoint_ns": "tools:call_1",
|
||||
"langgraph_node": "agent",
|
||||
}
|
||||
value_parts = [part for part in parts if part["type"] == "values"]
|
||||
assert value_parts == [
|
||||
{
|
||||
"type": "values",
|
||||
"ns": ("tools:call_1",),
|
||||
"data": {"messages": [{"type": "human", "content": "hi"}], "count": 1},
|
||||
"interrupts": (),
|
||||
},
|
||||
{
|
||||
"type": "values",
|
||||
"ns": ("tools:call_1",),
|
||||
"data": {"count": 2},
|
||||
"interrupts": (),
|
||||
},
|
||||
]
|
||||
_, kwargs = mock_async_client.runs.stream.call_args
|
||||
assert set(kwargs["stream_mode"]) == {
|
||||
"messages-tuple",
|
||||
"values",
|
||||
"compact",
|
||||
"updates",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.skip(
|
||||
"Unskip this test to manually test the LangSmith Deployment integration"
|
||||
)
|
||||
|
||||
@@ -90,6 +90,25 @@ def _make_messages_graph() -> StateGraph[
|
||||
return builder
|
||||
|
||||
|
||||
def _make_streaming_messages_graph() -> StateGraph[
|
||||
MessagesState, None, MessagesState, MessagesState
|
||||
]:
|
||||
model = FakeChatModel(messages=[AIMessage(content="hello world", id="ai-1")])
|
||||
|
||||
def call_model(state: MessagesState) -> dict[str, Any]:
|
||||
streamed = model.stream(state["messages"])
|
||||
message = next(streamed)
|
||||
for chunk in streamed:
|
||||
message += chunk
|
||||
return {"messages": message}
|
||||
|
||||
builder = StateGraph(MessagesState, input_schema=MessagesState)
|
||||
builder.add_node("call_model", call_model)
|
||||
builder.add_edge(START, "call_model")
|
||||
builder.add_edge("call_model", END)
|
||||
return builder
|
||||
|
||||
|
||||
def _make_custom_graph() -> Any:
|
||||
@entrypoint()
|
||||
def graph(inputs: Any, *, writer: StreamWriter) -> Any:
|
||||
@@ -164,6 +183,26 @@ class TestV1BackwardsCompat:
|
||||
ns, _data = chunk
|
||||
assert isinstance(ns, tuple)
|
||||
|
||||
def test_stream_v1_messages_keep_metadata_on_every_chunk(self) -> None:
|
||||
graph = _make_streaming_messages_graph().compile()
|
||||
chunks = list(graph.stream(_MSG_INPUT, stream_mode="messages"))
|
||||
metadata = [meta for _message, meta in chunks]
|
||||
assert len(metadata) >= 3
|
||||
assert all(isinstance(meta, dict) for meta in metadata)
|
||||
|
||||
def test_stream_v1_messages_compact_dedupes_metadata(self) -> None:
|
||||
graph = _make_streaming_messages_graph().compile()
|
||||
chunks = list(graph.stream(_MSG_INPUT, stream_mode=["messages", "compact"]))
|
||||
metadata = [
|
||||
meta
|
||||
for mode, payload in chunks
|
||||
if mode == "messages"
|
||||
for _message, meta in [payload]
|
||||
]
|
||||
assert len(metadata) >= 3
|
||||
assert isinstance(metadata[0], dict)
|
||||
assert all(meta is None for meta in metadata[1:])
|
||||
|
||||
|
||||
# --- v2 sync stream ---
|
||||
|
||||
@@ -205,6 +244,26 @@ class TestV2Stream:
|
||||
assert isinstance(metadata, dict)
|
||||
assert "langgraph_node" in metadata
|
||||
|
||||
def test_messages_streaming_compact_dedupes_metadata(self) -> None:
|
||||
graph = _make_streaming_messages_graph().compile()
|
||||
chunks = list(
|
||||
graph.stream(
|
||||
_MSG_INPUT,
|
||||
stream_mode=["messages", "compact"],
|
||||
version="v2",
|
||||
)
|
||||
)
|
||||
msg_chunks = [c for c in chunks if c["type"] == "messages"]
|
||||
assert len(msg_chunks) >= 3
|
||||
first_message, first_metadata = msg_chunks[0]["data"]
|
||||
assert isinstance(first_message, BaseMessage)
|
||||
assert isinstance(first_metadata, dict)
|
||||
assert "langgraph_node" in first_metadata
|
||||
for chunk in msg_chunks[1:]:
|
||||
message, metadata = chunk["data"]
|
||||
assert isinstance(message, BaseMessage)
|
||||
assert metadata is None
|
||||
|
||||
def test_custom(self) -> None:
|
||||
graph = _make_custom_graph()
|
||||
chunks = list(graph.stream({"key": "val"}, stream_mode="custom", version="v2"))
|
||||
@@ -544,6 +603,29 @@ class TestV2StreamAsync:
|
||||
assert isinstance(metadata, dict)
|
||||
assert "langgraph_node" in metadata
|
||||
|
||||
@NEEDS_CONTEXTVARS
|
||||
@pytest.mark.anyio
|
||||
async def test_messages_streaming_compact_dedupes_metadata(self) -> None:
|
||||
graph = _make_streaming_messages_graph().compile()
|
||||
chunks = [
|
||||
c
|
||||
async for c in graph.astream(
|
||||
_MSG_INPUT,
|
||||
stream_mode=["messages", "compact"],
|
||||
version="v2",
|
||||
)
|
||||
]
|
||||
msg_chunks = [c for c in chunks if c["type"] == "messages"]
|
||||
assert len(msg_chunks) >= 3
|
||||
first_message, first_metadata = msg_chunks[0]["data"]
|
||||
assert isinstance(first_message, BaseMessage)
|
||||
assert isinstance(first_metadata, dict)
|
||||
assert "langgraph_node" in first_metadata
|
||||
for chunk in msg_chunks[1:]:
|
||||
message, metadata = chunk["data"]
|
||||
assert isinstance(message, BaseMessage)
|
||||
assert metadata is None
|
||||
|
||||
@NEEDS_CONTEXTVARS
|
||||
@pytest.mark.anyio
|
||||
async def test_custom(self) -> None:
|
||||
@@ -1129,7 +1211,7 @@ _OutputT = TypeVar("_OutputT")
|
||||
_StateT = TypeVar("_StateT")
|
||||
|
||||
|
||||
def _check_type_narrowing(part: StreamPart[_OutputT, _StateT]) -> None:
|
||||
def _check_type_narrowing(part: StreamPart[_StateT, _OutputT]) -> None:
|
||||
"""Compile-time type narrowing checks — never called at runtime."""
|
||||
if part["type"] == "values":
|
||||
assert_type(part, ValuesStreamPart[_OutputT])
|
||||
|
||||
@@ -1116,6 +1116,631 @@ def test_subgraph_replay_from_subgraph_checkpoint(
|
||||
assert "post" in final_result["value"]
|
||||
|
||||
|
||||
def test_subgraph_time_travel_to_first_interrupt(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel to a subgraph checkpoint at the FIRST interrupt.
|
||||
|
||||
Architecture:
|
||||
Parent: START --> executor (subgraph, checkpointer=True) --> END
|
||||
Executor: START --> step_a --> ask_1 (interrupt) --> ask_2 (interrupt) --> END
|
||||
|
||||
Flow: run through both interrupts, then time travel back to the subgraph
|
||||
checkpoint captured at the first interrupt. ask_1 should re-fire,
|
||||
step_a should NOT re-run. Then resume through both interrupts with new answers.
|
||||
"""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
def ask_2(state: State) -> State:
|
||||
called.append("ask_2")
|
||||
answer = interrupt("Question 2?")
|
||||
return {"value": [f"ask_2:{answer}"]}
|
||||
|
||||
executor = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_node("ask_2", ask_2)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "ask_2")
|
||||
.add_edge("ask_2", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("executor", executor)
|
||||
.add_edge(START, "executor")
|
||||
.compile(checkpointer=sync_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# Run until first interrupt (ask_1)
|
||||
result = graph.invoke({"value": []}, config)
|
||||
assert "__interrupt__" in result
|
||||
assert result["__interrupt__"][0].value == "Question 1?"
|
||||
|
||||
# Capture subgraph state at the first interrupt
|
||||
parent_state = graph.get_state(config, subgraphs=True)
|
||||
sub_config_at_first = parent_state.tasks[0].state.config
|
||||
|
||||
# Resume first interrupt
|
||||
result = graph.invoke(Command(resume="answer_1"), config)
|
||||
assert result["__interrupt__"][0].value == "Question 2?"
|
||||
|
||||
# Resume second interrupt to complete
|
||||
result = graph.invoke(Command(resume="answer_2"), config)
|
||||
assert "__interrupt__" not in result
|
||||
|
||||
# --- Scenario 1: Replay from subgraph checkpoint at 1st interrupt ---
|
||||
called.clear()
|
||||
replay_result = graph.invoke(None, sub_config_at_first)
|
||||
assert "__interrupt__" in replay_result
|
||||
assert replay_result["__interrupt__"][0].value == "Question 1?"
|
||||
# step_a should NOT re-run — it was before this checkpoint
|
||||
assert "step_a" not in called
|
||||
# ask_1 re-fires because the interrupt replays
|
||||
assert "ask_1" in called
|
||||
|
||||
# --- Scenario 2: Fork from subgraph checkpoint at 1st interrupt ---
|
||||
called.clear()
|
||||
fork_config = graph.update_state(sub_config_at_first, {"value": ["forked"]})
|
||||
fork_result = graph.invoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
assert fork_result["__interrupt__"][0].value == "Question 1?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" in called
|
||||
|
||||
|
||||
def test_subgraph_time_travel_to_second_interrupt(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel to a subgraph checkpoint at the SECOND interrupt.
|
||||
|
||||
Architecture:
|
||||
Parent: START --> executor (subgraph, checkpointer=True) --> END
|
||||
Executor: START --> step_a --> ask_1 (interrupt) --> ask_2 (interrupt) --> END
|
||||
|
||||
Flow: run through both interrupts resuming each, then time travel back to the
|
||||
subgraph checkpoint at the second interrupt. Only ask_2 should re-fire.
|
||||
Then resume with a new answer and verify state.
|
||||
"""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
def ask_2(state: State) -> State:
|
||||
called.append("ask_2")
|
||||
answer = interrupt("Question 2?")
|
||||
return {"value": [f"ask_2:{answer}"]}
|
||||
|
||||
executor = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_node("ask_2", ask_2)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "ask_2")
|
||||
.add_edge("ask_2", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("executor", executor)
|
||||
.add_edge(START, "executor")
|
||||
.compile(checkpointer=sync_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# Run until first interrupt (ask_1)
|
||||
result = graph.invoke({"value": []}, config)
|
||||
assert result["__interrupt__"][0].value == "Question 1?"
|
||||
|
||||
# Resume first interrupt
|
||||
result = graph.invoke(Command(resume="answer_1"), config)
|
||||
assert result["__interrupt__"][0].value == "Question 2?"
|
||||
|
||||
# Capture subgraph state at the second interrupt
|
||||
parent_state = graph.get_state(config, subgraphs=True)
|
||||
sub_config = parent_state.tasks[0].state.config
|
||||
|
||||
# Resume second interrupt to complete the graph
|
||||
result = graph.invoke(Command(resume="answer_2"), config)
|
||||
assert "__interrupt__" not in result
|
||||
|
||||
# --- Scenario 1: Replay from subgraph checkpoint at 2nd interrupt ---
|
||||
called.clear()
|
||||
replay_result = graph.invoke(None, sub_config)
|
||||
assert "__interrupt__" in replay_result
|
||||
assert replay_result["__interrupt__"][0].value == "Question 2?"
|
||||
# step_a and ask_1 should NOT re-run — they were before this checkpoint
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" not in called
|
||||
|
||||
# --- Scenario 2: Fork from subgraph checkpoint at 2nd interrupt ---
|
||||
called.clear()
|
||||
fork_config = graph.update_state(sub_config, {"value": ["forked"]})
|
||||
fork_result = graph.invoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
assert fork_result["__interrupt__"][0].value == "Question 2?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" not in called
|
||||
|
||||
|
||||
def test_subgraph_time_travel_after_completion(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel to a subgraph checkpoint AFTER both interrupts are resolved.
|
||||
|
||||
Architecture:
|
||||
Parent: START --> executor (subgraph, checkpointer=True) --> END
|
||||
Executor: START --> step_a --> ask_1 (interrupt) --> ask_2 (interrupt) --> END
|
||||
|
||||
After completing the full flow, capture the subgraph's final state checkpoint
|
||||
and replay from it — should be a no-op (no nodes re-run).
|
||||
"""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
def ask_2(state: State) -> State:
|
||||
called.append("ask_2")
|
||||
answer = interrupt("Question 2?")
|
||||
return {"value": [f"ask_2:{answer}"]}
|
||||
|
||||
executor = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_node("ask_2", ask_2)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "ask_2")
|
||||
.add_edge("ask_2", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("executor", executor)
|
||||
.add_edge(START, "executor")
|
||||
.compile(checkpointer=sync_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# Run through both interrupts
|
||||
graph.invoke({"value": []}, config)
|
||||
graph.invoke(Command(resume="answer_1"), config)
|
||||
|
||||
# Before resuming 2nd interrupt, get state history to find the
|
||||
# subgraph checkpoint that will exist after ask_2 completes
|
||||
graph.invoke(Command(resume="answer_2"), config)
|
||||
|
||||
# Get the final parent state — no pending tasks
|
||||
final_state = graph.get_state(config)
|
||||
assert len(final_state.tasks) == 0
|
||||
|
||||
# Replay from the final parent checkpoint — should be a no-op
|
||||
called.clear()
|
||||
replay_result = graph.invoke(None, final_state.config)
|
||||
assert "__interrupt__" not in replay_result
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" not in called
|
||||
assert "ask_2" not in called
|
||||
# All values should be present
|
||||
assert "step_a_done" in replay_result["value"]
|
||||
assert "ask_1:answer_1" in replay_result["value"]
|
||||
assert "ask_2:answer_2" in replay_result["value"]
|
||||
|
||||
|
||||
def test_3_levels_deep_time_travel_to_first_interrupt(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel to the innermost subgraph checkpoint at the FIRST interrupt.
|
||||
|
||||
Architecture:
|
||||
Parent: START --> outer (subgraph, checkpointer=True) --> END
|
||||
Outer: START --> inner (subgraph, checkpointer=True) --> END
|
||||
Inner: START --> step_a --> ask_1 (interrupt) --> ask_2 (interrupt) --> END
|
||||
"""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
def ask_2(state: State) -> State:
|
||||
called.append("ask_2")
|
||||
answer = interrupt("Question 2?")
|
||||
return {"value": [f"ask_2:{answer}"]}
|
||||
|
||||
inner = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_node("ask_2", ask_2)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "ask_2")
|
||||
.add_edge("ask_2", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
middle = (
|
||||
StateGraph(State)
|
||||
.add_node("inner", inner)
|
||||
.add_edge(START, "inner")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("outer", middle)
|
||||
.add_edge(START, "outer")
|
||||
.compile(checkpointer=sync_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# Run until first interrupt
|
||||
result = graph.invoke({"value": []}, config)
|
||||
assert result["__interrupt__"][0].value == "Question 1?"
|
||||
|
||||
# Capture innermost subgraph state at the first interrupt
|
||||
parent_state = graph.get_state(config, subgraphs=True)
|
||||
mid_state = parent_state.tasks[0].state
|
||||
inner_config = mid_state.tasks[0].state.config
|
||||
|
||||
# Resume through both interrupts to complete
|
||||
graph.invoke(Command(resume="answer_1"), config)
|
||||
graph.invoke(Command(resume="answer_2"), config)
|
||||
|
||||
# --- Scenario 1: Replay from innermost checkpoint at 1st interrupt ---
|
||||
called.clear()
|
||||
replay_result = graph.invoke(None, inner_config)
|
||||
assert "__interrupt__" in replay_result
|
||||
assert replay_result["__interrupt__"][0].value == "Question 1?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" in called
|
||||
|
||||
# --- Scenario 2: Fork from innermost checkpoint at 1st interrupt ---
|
||||
called.clear()
|
||||
fork_config = graph.update_state(inner_config, {"value": ["forked"]})
|
||||
fork_result = graph.invoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
assert fork_result["__interrupt__"][0].value == "Question 1?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" in called
|
||||
|
||||
|
||||
def test_3_levels_deep_time_travel_to_second_interrupt(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel to the innermost subgraph checkpoint at the SECOND interrupt.
|
||||
|
||||
Architecture:
|
||||
Parent: START --> outer (subgraph, checkpointer=True) --> END
|
||||
Outer: START --> inner (subgraph, checkpointer=True) --> END
|
||||
Inner: START --> step_a --> ask_1 (interrupt) --> ask_2 (interrupt) --> END
|
||||
"""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
def ask_2(state: State) -> State:
|
||||
called.append("ask_2")
|
||||
answer = interrupt("Question 2?")
|
||||
return {"value": [f"ask_2:{answer}"]}
|
||||
|
||||
inner = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_node("ask_2", ask_2)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "ask_2")
|
||||
.add_edge("ask_2", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
middle = (
|
||||
StateGraph(State)
|
||||
.add_node("inner", inner)
|
||||
.add_edge(START, "inner")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("outer", middle)
|
||||
.add_edge(START, "outer")
|
||||
.compile(checkpointer=sync_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# Run until first interrupt
|
||||
graph.invoke({"value": []}, config)
|
||||
|
||||
# Resume first interrupt
|
||||
result = graph.invoke(Command(resume="answer_1"), config)
|
||||
assert result["__interrupt__"][0].value == "Question 2?"
|
||||
|
||||
# Capture innermost subgraph state at the second interrupt
|
||||
parent_state = graph.get_state(config, subgraphs=True)
|
||||
mid_state = parent_state.tasks[0].state
|
||||
inner_config = mid_state.tasks[0].state.config
|
||||
|
||||
# Resume second interrupt to complete
|
||||
graph.invoke(Command(resume="answer_2"), config)
|
||||
|
||||
# --- Scenario 1: Replay from innermost checkpoint at 2nd interrupt ---
|
||||
called.clear()
|
||||
replay_result = graph.invoke(None, inner_config)
|
||||
assert "__interrupt__" in replay_result
|
||||
assert replay_result["__interrupt__"][0].value == "Question 2?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" not in called
|
||||
|
||||
# --- Scenario 2: Fork from innermost checkpoint at 2nd interrupt ---
|
||||
called.clear()
|
||||
fork_config = graph.update_state(inner_config, {"value": ["forked"]})
|
||||
fork_result = graph.invoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
assert fork_result["__interrupt__"][0].value == "Question 2?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" not in called
|
||||
|
||||
|
||||
def test_3_levels_deep_time_travel_to_middle_subgraph(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel to the MIDDLE-level subgraph checkpoint (not innermost).
|
||||
|
||||
Architecture:
|
||||
Parent: START --> outer (subgraph, checkpointer=True) --> END
|
||||
Outer: START --> inner (subgraph, checkpointer=True) --> END
|
||||
Inner: START --> step_a --> ask_1 (interrupt) --> ask_2 (interrupt) --> END
|
||||
|
||||
After completing the full flow, time travel back to the middle subgraph's
|
||||
checkpoint at the second interrupt. The middle subgraph should replay the
|
||||
inner subgraph from the correct point.
|
||||
"""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
def ask_2(state: State) -> State:
|
||||
called.append("ask_2")
|
||||
answer = interrupt("Question 2?")
|
||||
return {"value": [f"ask_2:{answer}"]}
|
||||
|
||||
inner = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_node("ask_2", ask_2)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "ask_2")
|
||||
.add_edge("ask_2", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
middle = (
|
||||
StateGraph(State)
|
||||
.add_node("inner", inner)
|
||||
.add_edge(START, "inner")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("outer", middle)
|
||||
.add_edge(START, "outer")
|
||||
.compile(checkpointer=sync_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# Run until first interrupt
|
||||
graph.invoke({"value": []}, config)
|
||||
|
||||
# Resume first, capture middle config at second interrupt
|
||||
graph.invoke(Command(resume="answer_1"), config)
|
||||
parent_state = graph.get_state(config, subgraphs=True)
|
||||
mid_config = parent_state.tasks[0].state.config
|
||||
|
||||
# Resume second to complete
|
||||
graph.invoke(Command(resume="answer_2"), config)
|
||||
|
||||
# --- Scenario 1: Replay from middle-level subgraph checkpoint ---
|
||||
# The middle subgraph's checkpoint knows about the inner subgraph's state
|
||||
# via checkpoint_map, so the inner replays from the correct point.
|
||||
called.clear()
|
||||
replay_result = graph.invoke(None, mid_config)
|
||||
assert "__interrupt__" in replay_result
|
||||
|
||||
# --- Scenario 2: Fork from middle-level subgraph checkpoint ---
|
||||
called.clear()
|
||||
fork_config = graph.update_state(mid_config, {"value": ["forked"]})
|
||||
fork_result = graph.invoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
|
||||
|
||||
def test_3_levels_deep_middle_has_interrupts(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel when the MIDDLE subgraph itself has interrupts.
|
||||
|
||||
Architecture:
|
||||
Parent: START --> outer (subgraph, checkpointer=True) --> END
|
||||
Outer: START --> pre (interrupt) --> inner (subgraph, checkpointer=True) --> END
|
||||
Inner: START --> step_a --> ask_1 (interrupt) --> END
|
||||
|
||||
Flow: run through both interrupts (pre then ask_1), then time travel back to
|
||||
the middle subgraph checkpoint at each interrupt point.
|
||||
"""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
def pre(state: State) -> State:
|
||||
called.append("pre")
|
||||
answer = interrupt("Pre-question?")
|
||||
return {"value": [f"pre:{answer}"]}
|
||||
|
||||
def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
inner = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
middle = (
|
||||
StateGraph(State)
|
||||
.add_node("pre", pre)
|
||||
.add_node("inner", inner)
|
||||
.add_edge(START, "pre")
|
||||
.add_edge("pre", "inner")
|
||||
.add_edge("inner", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("outer", middle)
|
||||
.add_edge(START, "outer")
|
||||
.compile(checkpointer=sync_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# Run until first interrupt (pre in middle subgraph)
|
||||
result = graph.invoke({"value": []}, config)
|
||||
assert result["__interrupt__"][0].value == "Pre-question?"
|
||||
|
||||
# Capture middle subgraph config at the pre interrupt
|
||||
parent_state = graph.get_state(config, subgraphs=True)
|
||||
mid_config_at_pre = parent_state.tasks[0].state.config
|
||||
|
||||
# Resume pre, hits ask_1 in inner subgraph
|
||||
result = graph.invoke(Command(resume="pre_answer"), config)
|
||||
assert result["__interrupt__"][0].value == "Question 1?"
|
||||
|
||||
# Capture middle subgraph config at the ask_1 interrupt
|
||||
parent_state = graph.get_state(config, subgraphs=True)
|
||||
mid_config_at_ask1 = parent_state.tasks[0].state.config
|
||||
|
||||
# Resume ask_1 to complete
|
||||
result = graph.invoke(Command(resume="answer_1"), config)
|
||||
assert "__interrupt__" not in result
|
||||
|
||||
# --- Time travel to middle checkpoint at pre interrupt ---
|
||||
called.clear()
|
||||
replay_result = graph.invoke(None, mid_config_at_pre)
|
||||
assert "__interrupt__" in replay_result
|
||||
assert replay_result["__interrupt__"][0].value == "Pre-question?"
|
||||
# pre should re-fire (interrupt replays), but nothing else should run
|
||||
assert "pre" in called
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" not in called
|
||||
|
||||
# Fork from middle checkpoint at pre interrupt
|
||||
called.clear()
|
||||
fork_config = graph.update_state(mid_config_at_pre, {"value": ["forked"]})
|
||||
fork_result = graph.invoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
assert fork_result["__interrupt__"][0].value == "Pre-question?"
|
||||
assert "pre" in called
|
||||
assert "step_a" not in called
|
||||
|
||||
# --- Time travel to middle checkpoint at ask_1 interrupt ---
|
||||
called.clear()
|
||||
replay_result = graph.invoke(None, mid_config_at_ask1)
|
||||
assert "__interrupt__" in replay_result
|
||||
assert replay_result["__interrupt__"][0].value == "Question 1?"
|
||||
# pre should NOT re-run (it completed before this checkpoint)
|
||||
assert "pre" not in called
|
||||
# ask_1 re-fires
|
||||
assert "ask_1" in called
|
||||
|
||||
# Fork from middle checkpoint at ask_1 interrupt
|
||||
called.clear()
|
||||
fork_config = graph.update_state(mid_config_at_ask1, {"value": ["forked"]})
|
||||
fork_result = graph.invoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
assert fork_result["__interrupt__"][0].value == "Question 1?"
|
||||
assert "pre" not in called
|
||||
assert "ask_1" in called
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Section 6: __copy__ / update_state(None)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -1051,6 +1051,552 @@ async def test_subgraph_interrupt_full_flow_no_sub_checkpointer(
|
||||
assert "post" in final_result["value"]
|
||||
|
||||
|
||||
@NEEDS_CONTEXTVARS
|
||||
async def test_subgraph_time_travel_to_first_interrupt_async(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel to a subgraph checkpoint at the FIRST interrupt (async)."""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
async def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
async def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
async def ask_2(state: State) -> State:
|
||||
called.append("ask_2")
|
||||
answer = interrupt("Question 2?")
|
||||
return {"value": [f"ask_2:{answer}"]}
|
||||
|
||||
executor = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_node("ask_2", ask_2)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "ask_2")
|
||||
.add_edge("ask_2", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("executor", executor)
|
||||
.add_edge(START, "executor")
|
||||
.compile(checkpointer=async_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# Run until first interrupt (ask_1)
|
||||
result = await graph.ainvoke({"value": []}, config)
|
||||
assert result["__interrupt__"][0].value == "Question 1?"
|
||||
|
||||
# Capture subgraph state at the first interrupt
|
||||
parent_state = await graph.aget_state(config, subgraphs=True)
|
||||
sub_config_at_first = parent_state.tasks[0].state.config
|
||||
|
||||
# Resume through both interrupts to complete
|
||||
await graph.ainvoke(Command(resume="answer_1"), config)
|
||||
await graph.ainvoke(Command(resume="answer_2"), config)
|
||||
|
||||
# --- Scenario 1: Replay from subgraph checkpoint at 1st interrupt ---
|
||||
called.clear()
|
||||
replay_result = await graph.ainvoke(None, sub_config_at_first)
|
||||
assert "__interrupt__" in replay_result
|
||||
assert replay_result["__interrupt__"][0].value == "Question 1?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" in called
|
||||
|
||||
# --- Scenario 2: Fork from subgraph checkpoint at 1st interrupt ---
|
||||
called.clear()
|
||||
fork_config = await graph.aupdate_state(sub_config_at_first, {"value": ["forked"]})
|
||||
fork_result = await graph.ainvoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
assert fork_result["__interrupt__"][0].value == "Question 1?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" in called
|
||||
|
||||
|
||||
@NEEDS_CONTEXTVARS
|
||||
async def test_subgraph_time_travel_to_second_interrupt_async(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel to a subgraph checkpoint at the SECOND interrupt (async)."""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
async def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
async def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
async def ask_2(state: State) -> State:
|
||||
called.append("ask_2")
|
||||
answer = interrupt("Question 2?")
|
||||
return {"value": [f"ask_2:{answer}"]}
|
||||
|
||||
executor = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_node("ask_2", ask_2)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "ask_2")
|
||||
.add_edge("ask_2", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("executor", executor)
|
||||
.add_edge(START, "executor")
|
||||
.compile(checkpointer=async_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# Run until first interrupt
|
||||
await graph.ainvoke({"value": []}, config)
|
||||
|
||||
# Resume first interrupt
|
||||
result = await graph.ainvoke(Command(resume="answer_1"), config)
|
||||
assert result["__interrupt__"][0].value == "Question 2?"
|
||||
|
||||
# Capture subgraph state at the second interrupt
|
||||
parent_state = await graph.aget_state(config, subgraphs=True)
|
||||
sub_config = parent_state.tasks[0].state.config
|
||||
|
||||
# Resume second interrupt to complete
|
||||
await graph.ainvoke(Command(resume="answer_2"), config)
|
||||
|
||||
# --- Scenario 1: Replay from subgraph checkpoint at 2nd interrupt ---
|
||||
called.clear()
|
||||
replay_result = await graph.ainvoke(None, sub_config)
|
||||
assert "__interrupt__" in replay_result
|
||||
assert replay_result["__interrupt__"][0].value == "Question 2?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" not in called
|
||||
|
||||
# --- Scenario 2: Fork from subgraph checkpoint at 2nd interrupt ---
|
||||
called.clear()
|
||||
fork_config = await graph.aupdate_state(sub_config, {"value": ["forked"]})
|
||||
fork_result = await graph.ainvoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
assert fork_result["__interrupt__"][0].value == "Question 2?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" not in called
|
||||
|
||||
|
||||
@NEEDS_CONTEXTVARS
|
||||
async def test_subgraph_time_travel_after_completion_async(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel to a subgraph checkpoint AFTER both interrupts resolved (async)."""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
async def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
async def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
async def ask_2(state: State) -> State:
|
||||
called.append("ask_2")
|
||||
answer = interrupt("Question 2?")
|
||||
return {"value": [f"ask_2:{answer}"]}
|
||||
|
||||
executor = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_node("ask_2", ask_2)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "ask_2")
|
||||
.add_edge("ask_2", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("executor", executor)
|
||||
.add_edge(START, "executor")
|
||||
.compile(checkpointer=async_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
await graph.ainvoke({"value": []}, config)
|
||||
await graph.ainvoke(Command(resume="answer_1"), config)
|
||||
await graph.ainvoke(Command(resume="answer_2"), config)
|
||||
|
||||
final_state = await graph.aget_state(config)
|
||||
assert len(final_state.tasks) == 0
|
||||
|
||||
# Replay from the final parent checkpoint — should be a no-op
|
||||
called.clear()
|
||||
replay_result = await graph.ainvoke(None, final_state.config)
|
||||
assert "__interrupt__" not in replay_result
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" not in called
|
||||
assert "ask_2" not in called
|
||||
assert "step_a_done" in replay_result["value"]
|
||||
assert "ask_1:answer_1" in replay_result["value"]
|
||||
assert "ask_2:answer_2" in replay_result["value"]
|
||||
|
||||
|
||||
@NEEDS_CONTEXTVARS
|
||||
async def test_3_levels_deep_time_travel_to_first_interrupt_async(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel to innermost subgraph checkpoint at FIRST interrupt (async, 3 levels)."""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
async def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
async def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
async def ask_2(state: State) -> State:
|
||||
called.append("ask_2")
|
||||
answer = interrupt("Question 2?")
|
||||
return {"value": [f"ask_2:{answer}"]}
|
||||
|
||||
inner = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_node("ask_2", ask_2)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "ask_2")
|
||||
.add_edge("ask_2", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
middle = (
|
||||
StateGraph(State)
|
||||
.add_node("inner", inner)
|
||||
.add_edge(START, "inner")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("outer", middle)
|
||||
.add_edge(START, "outer")
|
||||
.compile(checkpointer=async_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
result = await graph.ainvoke({"value": []}, config)
|
||||
assert result["__interrupt__"][0].value == "Question 1?"
|
||||
|
||||
parent_state = await graph.aget_state(config, subgraphs=True)
|
||||
mid_state = parent_state.tasks[0].state
|
||||
inner_config = mid_state.tasks[0].state.config
|
||||
|
||||
await graph.ainvoke(Command(resume="answer_1"), config)
|
||||
await graph.ainvoke(Command(resume="answer_2"), config)
|
||||
|
||||
# --- Scenario 1: Replay from innermost checkpoint at 1st interrupt ---
|
||||
called.clear()
|
||||
replay_result = await graph.ainvoke(None, inner_config)
|
||||
assert "__interrupt__" in replay_result
|
||||
assert replay_result["__interrupt__"][0].value == "Question 1?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" in called
|
||||
|
||||
# --- Scenario 2: Fork from innermost checkpoint at 1st interrupt ---
|
||||
called.clear()
|
||||
fork_config = await graph.aupdate_state(inner_config, {"value": ["forked"]})
|
||||
fork_result = await graph.ainvoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
assert fork_result["__interrupt__"][0].value == "Question 1?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" in called
|
||||
|
||||
|
||||
@NEEDS_CONTEXTVARS
|
||||
async def test_3_levels_deep_time_travel_to_second_interrupt_async(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel to innermost subgraph checkpoint at SECOND interrupt (async, 3 levels)."""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
async def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
async def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
async def ask_2(state: State) -> State:
|
||||
called.append("ask_2")
|
||||
answer = interrupt("Question 2?")
|
||||
return {"value": [f"ask_2:{answer}"]}
|
||||
|
||||
inner = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_node("ask_2", ask_2)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "ask_2")
|
||||
.add_edge("ask_2", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
middle = (
|
||||
StateGraph(State)
|
||||
.add_node("inner", inner)
|
||||
.add_edge(START, "inner")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("outer", middle)
|
||||
.add_edge(START, "outer")
|
||||
.compile(checkpointer=async_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
await graph.ainvoke({"value": []}, config)
|
||||
|
||||
result = await graph.ainvoke(Command(resume="answer_1"), config)
|
||||
assert result["__interrupt__"][0].value == "Question 2?"
|
||||
|
||||
parent_state = await graph.aget_state(config, subgraphs=True)
|
||||
mid_state = parent_state.tasks[0].state
|
||||
inner_config = mid_state.tasks[0].state.config
|
||||
|
||||
await graph.ainvoke(Command(resume="answer_2"), config)
|
||||
|
||||
# --- Scenario 1: Replay ---
|
||||
called.clear()
|
||||
replay_result = await graph.ainvoke(None, inner_config)
|
||||
assert "__interrupt__" in replay_result
|
||||
assert replay_result["__interrupt__"][0].value == "Question 2?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" not in called
|
||||
|
||||
# --- Scenario 2: Fork ---
|
||||
called.clear()
|
||||
fork_config = await graph.aupdate_state(inner_config, {"value": ["forked"]})
|
||||
fork_result = await graph.ainvoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
assert fork_result["__interrupt__"][0].value == "Question 2?"
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" not in called
|
||||
|
||||
|
||||
@NEEDS_CONTEXTVARS
|
||||
async def test_3_levels_deep_time_travel_to_middle_subgraph_async(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel to the MIDDLE-level subgraph checkpoint (async, 3 levels)."""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
async def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
async def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
async def ask_2(state: State) -> State:
|
||||
called.append("ask_2")
|
||||
answer = interrupt("Question 2?")
|
||||
return {"value": [f"ask_2:{answer}"]}
|
||||
|
||||
inner = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_node("ask_2", ask_2)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "ask_2")
|
||||
.add_edge("ask_2", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
middle = (
|
||||
StateGraph(State)
|
||||
.add_node("inner", inner)
|
||||
.add_edge(START, "inner")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("outer", middle)
|
||||
.add_edge(START, "outer")
|
||||
.compile(checkpointer=async_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
await graph.ainvoke({"value": []}, config)
|
||||
await graph.ainvoke(Command(resume="answer_1"), config)
|
||||
|
||||
parent_state = await graph.aget_state(config, subgraphs=True)
|
||||
mid_config = parent_state.tasks[0].state.config
|
||||
|
||||
await graph.ainvoke(Command(resume="answer_2"), config)
|
||||
|
||||
# --- Scenario 1: Replay from middle-level subgraph checkpoint ---
|
||||
# The middle subgraph's checkpoint knows about the inner subgraph's state
|
||||
# via checkpoint_map, so the inner replays from the correct point.
|
||||
called.clear()
|
||||
replay_result = await graph.ainvoke(None, mid_config)
|
||||
assert "__interrupt__" in replay_result
|
||||
|
||||
# --- Scenario 2: Fork from middle-level subgraph checkpoint ---
|
||||
called.clear()
|
||||
fork_config = await graph.aupdate_state(mid_config, {"value": ["forked"]})
|
||||
fork_result = await graph.ainvoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
|
||||
|
||||
@NEEDS_CONTEXTVARS
|
||||
async def test_3_levels_deep_middle_has_interrupts_async(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Time travel when the MIDDLE subgraph itself has interrupts (async)."""
|
||||
|
||||
called: list[str] = []
|
||||
|
||||
async def pre(state: State) -> State:
|
||||
called.append("pre")
|
||||
answer = interrupt("Pre-question?")
|
||||
return {"value": [f"pre:{answer}"]}
|
||||
|
||||
async def step_a(state: State) -> State:
|
||||
called.append("step_a")
|
||||
return {"value": ["step_a_done"]}
|
||||
|
||||
async def ask_1(state: State) -> State:
|
||||
called.append("ask_1")
|
||||
answer = interrupt("Question 1?")
|
||||
return {"value": [f"ask_1:{answer}"]}
|
||||
|
||||
inner = (
|
||||
StateGraph(State)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("ask_1", ask_1)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "ask_1")
|
||||
.add_edge("ask_1", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
middle = (
|
||||
StateGraph(State)
|
||||
.add_node("pre", pre)
|
||||
.add_node("inner", inner)
|
||||
.add_edge(START, "pre")
|
||||
.add_edge("pre", "inner")
|
||||
.add_edge("inner", "__end__")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_node("outer", middle)
|
||||
.add_edge(START, "outer")
|
||||
.compile(checkpointer=async_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# Run until first interrupt (pre in middle subgraph)
|
||||
result = await graph.ainvoke({"value": []}, config)
|
||||
assert result["__interrupt__"][0].value == "Pre-question?"
|
||||
|
||||
# Capture middle subgraph config at the pre interrupt
|
||||
parent_state = await graph.aget_state(config, subgraphs=True)
|
||||
mid_config_at_pre = parent_state.tasks[0].state.config
|
||||
|
||||
# Resume pre, hits ask_1 in inner subgraph
|
||||
result = await graph.ainvoke(Command(resume="pre_answer"), config)
|
||||
assert result["__interrupt__"][0].value == "Question 1?"
|
||||
|
||||
# Capture middle subgraph config at the ask_1 interrupt
|
||||
parent_state = await graph.aget_state(config, subgraphs=True)
|
||||
mid_config_at_ask1 = parent_state.tasks[0].state.config
|
||||
|
||||
# Resume ask_1 to complete
|
||||
result = await graph.ainvoke(Command(resume="answer_1"), config)
|
||||
assert "__interrupt__" not in result
|
||||
|
||||
# --- Time travel to middle checkpoint at pre interrupt ---
|
||||
called.clear()
|
||||
replay_result = await graph.ainvoke(None, mid_config_at_pre)
|
||||
assert "__interrupt__" in replay_result
|
||||
assert replay_result["__interrupt__"][0].value == "Pre-question?"
|
||||
assert "pre" in called
|
||||
assert "step_a" not in called
|
||||
assert "ask_1" not in called
|
||||
|
||||
# Fork from middle checkpoint at pre interrupt
|
||||
called.clear()
|
||||
fork_config = await graph.aupdate_state(mid_config_at_pre, {"value": ["forked"]})
|
||||
fork_result = await graph.ainvoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
assert fork_result["__interrupt__"][0].value == "Pre-question?"
|
||||
assert "pre" in called
|
||||
assert "step_a" not in called
|
||||
|
||||
# --- Time travel to middle checkpoint at ask_1 interrupt ---
|
||||
called.clear()
|
||||
replay_result = await graph.ainvoke(None, mid_config_at_ask1)
|
||||
assert "__interrupt__" in replay_result
|
||||
assert replay_result["__interrupt__"][0].value == "Question 1?"
|
||||
assert "pre" not in called
|
||||
assert "ask_1" in called
|
||||
|
||||
# Fork from middle checkpoint at ask_1 interrupt
|
||||
called.clear()
|
||||
fork_config = await graph.aupdate_state(mid_config_at_ask1, {"value": ["forked"]})
|
||||
fork_result = await graph.ainvoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
assert fork_result["__interrupt__"][0].value == "Question 1?"
|
||||
assert "pre" not in called
|
||||
assert "ask_1" in called
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Section 6: __copy__ / update_state(None)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Generated
+12
-14
@@ -1367,7 +1367,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph"
|
||||
version = "1.1.0"
|
||||
version = "1.1.2"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -3615,21 +3615,19 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "tornado"
|
||||
version = "6.5.4"
|
||||
version = "6.5.5"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/37/1d/0a336abf618272d53f62ebe274f712e213f5a03c0b2339575430b8362ef2/tornado-6.5.4.tar.gz", hash = "sha256:a22fa9047405d03260b483980635f0b041989d8bcc9a313f8fe18b411d84b1d7", size = 513632, upload-time = "2025-12-15T19:21:03.836Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/f8/f1/3173dfa4a18db4a9b03e5d55325559dab51ee653763bb8745a75af491286/tornado-6.5.5.tar.gz", hash = "sha256:192b8f3ea91bd7f1f50c06955416ed76c6b72f96779b962f07f911b91e8d30e9", size = 516006, upload-time = "2026-03-10T21:31:02.067Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/ab/a9/e94a9d5224107d7ce3cc1fab8d5dc97f5ea351ccc6322ee4fb661da94e35/tornado-6.5.4-cp39-abi3-macosx_10_9_universal2.whl", hash = "sha256:d6241c1a16b1c9e4cc28148b1cda97dd1c6cb4fb7068ac1bedc610768dff0ba9", size = 443909, upload-time = "2025-12-15T19:20:48.382Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/db/7e/f7b8d8c4453f305a51f80dbb49014257bb7d28ccb4bbb8dd328ea995ecad/tornado-6.5.4-cp39-abi3-macosx_10_9_x86_64.whl", hash = "sha256:2d50f63dda1d2cac3ae1fa23d254e16b5e38153758470e9956cbc3d813d40843", size = 442163, upload-time = "2025-12-15T19:20:49.791Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/ba/b5/206f82d51e1bfa940ba366a8d2f83904b15942c45a78dd978b599870ab44/tornado-6.5.4-cp39-abi3-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:d1cf66105dc6acb5af613c054955b8137e34a03698aa53272dbda4afe252be17", size = 445746, upload-time = "2025-12-15T19:20:51.491Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/8e/9d/1a3338e0bd30ada6ad4356c13a0a6c35fbc859063fa7eddb309183364ac1/tornado-6.5.4-cp39-abi3-manylinux_2_5_i686.manylinux1_i686.manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:50ff0a58b0dc97939d29da29cd624da010e7f804746621c78d14b80238669335", size = 445083, upload-time = "2025-12-15T19:20:52.778Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/50/d4/e51d52047e7eb9a582da59f32125d17c0482d065afd5d3bc435ff2120dc5/tornado-6.5.4-cp39-abi3-manylinux_2_5_x86_64.manylinux1_x86_64.manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:e5fb5e04efa54cf0baabdd10061eb4148e0be137166146fff835745f59ab9f7f", size = 445315, upload-time = "2025-12-15T19:20:53.996Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/27/07/2273972f69ca63dbc139694a3fc4684edec3ea3f9efabf77ed32483b875c/tornado-6.5.4-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:9c86b1643b33a4cd415f8d0fe53045f913bf07b4a3ef646b735a6a86047dda84", size = 446003, upload-time = "2025-12-15T19:20:56.101Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/d1/83/41c52e47502bf7260044413b6770d1a48dda2f0246f95ee1384a3cd9c44a/tornado-6.5.4-cp39-abi3-musllinux_1_2_i686.whl", hash = "sha256:6eb82872335a53dd063a4f10917b3efd28270b56a33db69009606a0312660a6f", size = 445412, upload-time = "2025-12-15T19:20:57.398Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/10/c7/bc96917f06cbee182d44735d4ecde9c432e25b84f4c2086143013e7b9e52/tornado-6.5.4-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:6076d5dda368c9328ff41ab5d9dd3608e695e8225d1cd0fd1e006f05da3635a8", size = 445392, upload-time = "2025-12-15T19:20:58.692Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/0c/1a/d7592328d037d36f2d2462f4bc1fbb383eec9278bc786c1b111cbbd44cfa/tornado-6.5.4-cp39-abi3-win32.whl", hash = "sha256:1768110f2411d5cd281bac0a090f707223ce77fd110424361092859e089b38d1", size = 446481, upload-time = "2025-12-15T19:21:00.008Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/d6/6d/c69be695a0a64fd37a97db12355a035a6d90f79067a3cf936ec2b1dc38cd/tornado-6.5.4-cp39-abi3-win_amd64.whl", hash = "sha256:fa07d31e0cd85c60713f2b995da613588aa03e1303d75705dca6af8babc18ddc", size = 446886, upload-time = "2025-12-15T19:21:01.287Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/50/49/8dc3fd90902f70084bd2cd059d576ddb4f8bb44c2c7c0e33a11422acb17e/tornado-6.5.4-cp39-abi3-win_arm64.whl", hash = "sha256:053e6e16701eb6cbe641f308f4c1a9541f91b6261991160391bfc342e8a551a1", size = 445910, upload-time = "2025-12-15T19:21:02.571Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/59/8c/77f5097695f4dd8255ecbd08b2a1ed8ba8b953d337804dd7080f199e12bf/tornado-6.5.5-cp39-abi3-macosx_10_9_universal2.whl", hash = "sha256:487dc9cc380e29f58c7ab88f9e27cdeef04b2140862e5076a66fb6bb68bb1bfa", size = 445983, upload-time = "2026-03-10T21:30:44.28Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/ab/5e/7625b76cd10f98f1516c36ce0346de62061156352353ef2da44e5c21523c/tornado-6.5.5-cp39-abi3-macosx_10_9_x86_64.whl", hash = "sha256:65a7f1d46d4bb41df1ac99f5fcb685fb25c7e61613742d5108b010975a9a6521", size = 444246, upload-time = "2026-03-10T21:30:46.571Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/b2/04/7b5705d5b3c0fab088f434f9c83edac1573830ca49ccf29fb83bf7178eec/tornado-6.5.5-cp39-abi3-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:e74c92e8e65086b338fd56333fb9a68b9f6f2fe7ad532645a290a464bcf46be5", size = 447229, upload-time = "2026-03-10T21:30:48.273Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/34/01/74e034a30ef59afb4097ef8659515e96a39d910b712a89af76f5e4e1f93c/tornado-6.5.5-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:435319e9e340276428bbdb4e7fa732c2d399386d1de5686cb331ec8eee754f07", size = 448192, upload-time = "2026-03-10T21:30:51.22Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/be/00/fe9e02c5a96429fce1a1d15a517f5d8444f9c412e0bb9eadfbe3b0fc55bf/tornado-6.5.5-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:3f54aa540bdbfee7b9eb268ead60e7d199de5021facd276819c193c0fb28ea4e", size = 448039, upload-time = "2026-03-10T21:30:53.52Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/82/9e/656ee4cec0398b1d18d0f1eb6372c41c6b889722641d84948351ae19556d/tornado-6.5.5-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:36abed1754faeb80fbd6e64db2758091e1320f6bba74a4cf8c09cd18ccce8aca", size = 447445, upload-time = "2026-03-10T21:30:55.541Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/5a/76/4921c00511f88af86a33de770d64141170f1cfd9c00311aea689949e274e/tornado-6.5.5-cp39-abi3-win32.whl", hash = "sha256:dd3eafaaeec1c7f2f8fdcd5f964e8907ad788fe8a5a32c4426fbbdda621223b7", size = 448582, upload-time = "2026-03-10T21:30:57.142Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/2c/23/f6c6112a04d28eed765e374435fb1a9198f73e1ec4b4024184f21faeb1ad/tornado-6.5.5-cp39-abi3-win_amd64.whl", hash = "sha256:6443a794ba961a9f619b1ae926a2e900ac20c34483eea67be4ed8f1e58d3ef7b", size = 448990, upload-time = "2026-03-10T21:30:58.857Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/b7/c8/876602cbc96469911f0939f703453c1157b0c826ecb05bdd32e023397d4e/tornado-6.5.5-cp39-abi3-win_arm64.whl", hash = "sha256:2c9a876e094109333f888539ddb2de4361743e5d21eece20688e3e351e4990a6", size = 448016, upload-time = "2026-03-10T21:31:00.43Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
||||
Generated
+1
-1
@@ -268,7 +268,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph"
|
||||
version = "1.1.0"
|
||||
version = "1.1.2"
|
||||
source = { editable = "../langgraph" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -3,6 +3,6 @@ from langgraph_sdk.client import get_client, get_sync_client
|
||||
from langgraph_sdk.encryption import Encryption
|
||||
from langgraph_sdk.encryption.types import EncryptionContext
|
||||
|
||||
__version__ = "0.3.10"
|
||||
__version__ = "0.3.11"
|
||||
|
||||
__all__ = ["Auth", "Encryption", "EncryptionContext", "get_client", "get_sync_client"]
|
||||
|
||||
@@ -4,10 +4,11 @@ from __future__ import annotations
|
||||
|
||||
import warnings
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from datetime import datetime, tzinfo
|
||||
from typing import Any
|
||||
|
||||
from langgraph_sdk._async.http import HttpClient
|
||||
from langgraph_sdk._shared.utilities import _resolve_timezone
|
||||
from langgraph_sdk.schema import (
|
||||
All,
|
||||
Config,
|
||||
@@ -70,6 +71,7 @@ class CronClient:
|
||||
multitask_strategy: str | None = None,
|
||||
end_time: datetime | None = None,
|
||||
enabled: bool | None = None,
|
||||
timezone: str | tzinfo | None = None,
|
||||
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
|
||||
stream_subgraphs: bool | None = None,
|
||||
stream_resumable: bool | None = None,
|
||||
@@ -84,7 +86,7 @@ class CronClient:
|
||||
assistant_id: The assistant ID or graph name to use for the cron job.
|
||||
If using graph name, will default to first assistant created from that graph.
|
||||
schedule: The cron schedule to execute this job on.
|
||||
Schedules are interpreted in UTC.
|
||||
Schedules are interpreted in UTC unless a timezone is specified.
|
||||
input: The input to the graph.
|
||||
metadata: Metadata to assign to the cron job runs.
|
||||
config: The configuration for the assistant.
|
||||
@@ -100,6 +102,7 @@ class CronClient:
|
||||
Must be one of 'reject', 'interrupt', 'rollback', or 'enqueue'.
|
||||
end_time: The time to stop running the cron job. If not provided, the cron job will run indefinitely.
|
||||
enabled: Whether the cron job is enabled or not.
|
||||
timezone: IANA timezone for the cron schedule. Accepts a string (e.g. 'America/New_York') or a ``datetime.tzinfo`` instance (e.g. ``ZoneInfo("America/New_York")``).
|
||||
stream_mode: The stream mode(s) to use.
|
||||
stream_subgraphs: Whether to stream output from subgraphs.
|
||||
stream_resumable: Whether to persist the stream chunks in order to resume the stream later.
|
||||
@@ -152,6 +155,7 @@ class CronClient:
|
||||
"webhook": webhook,
|
||||
"end_time": end_time.isoformat() if end_time else None,
|
||||
"enabled": enabled,
|
||||
"timezone": _resolve_timezone(timezone),
|
||||
"stream_mode": stream_mode,
|
||||
"stream_subgraphs": stream_subgraphs,
|
||||
"stream_resumable": stream_resumable,
|
||||
@@ -184,6 +188,7 @@ class CronClient:
|
||||
multitask_strategy: str | None = None,
|
||||
end_time: datetime | None = None,
|
||||
enabled: bool | None = None,
|
||||
timezone: str | tzinfo | None = None,
|
||||
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
|
||||
stream_subgraphs: bool | None = None,
|
||||
stream_resumable: bool | None = None,
|
||||
@@ -197,7 +202,7 @@ class CronClient:
|
||||
assistant_id: The assistant ID or graph name to use for the cron job.
|
||||
If using graph name, will default to first assistant created from that graph.
|
||||
schedule: The cron schedule to execute this job on.
|
||||
Schedules are interpreted in UTC.
|
||||
Schedules are interpreted in UTC unless a timezone is specified.
|
||||
input: The input to the graph.
|
||||
metadata: Metadata to assign to the cron job runs.
|
||||
config: The configuration for the assistant.
|
||||
@@ -215,6 +220,7 @@ class CronClient:
|
||||
Must be one of 'reject', 'interrupt', 'rollback', or 'enqueue'.
|
||||
end_time: The time to stop running the cron job. If not provided, the cron job will run indefinitely.
|
||||
enabled: Whether the cron job is enabled or not.
|
||||
timezone: IANA timezone for the cron schedule. Accepts a string (e.g. 'America/New_York') or a ``datetime.tzinfo`` instance (e.g. ``ZoneInfo("America/New_York")``).
|
||||
stream_mode: The stream mode(s) to use.
|
||||
stream_subgraphs: Whether to stream output from subgraphs.
|
||||
stream_resumable: Whether to persist the stream chunks in order to resume the stream later.
|
||||
@@ -268,6 +274,7 @@ class CronClient:
|
||||
"on_run_completed": on_run_completed,
|
||||
"end_time": end_time.isoformat() if end_time else None,
|
||||
"enabled": enabled,
|
||||
"timezone": _resolve_timezone(timezone),
|
||||
"stream_mode": stream_mode,
|
||||
"stream_subgraphs": stream_subgraphs,
|
||||
"stream_resumable": stream_resumable,
|
||||
@@ -324,6 +331,7 @@ class CronClient:
|
||||
interrupt_after: All | list[str] | None = None,
|
||||
on_run_completed: OnCompletionBehavior | None = None,
|
||||
enabled: bool | None = None,
|
||||
timezone: str | tzinfo | None = None,
|
||||
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
|
||||
stream_subgraphs: bool | None = None,
|
||||
stream_resumable: bool | None = None,
|
||||
@@ -336,7 +344,7 @@ class CronClient:
|
||||
Args:
|
||||
cron_id: The cron ID to update.
|
||||
schedule: The cron schedule to execute this job on.
|
||||
Schedules are interpreted in UTC.
|
||||
Schedules are interpreted in UTC unless a timezone is specified.
|
||||
end_time: The end date to stop running the cron.
|
||||
input: The input to the graph.
|
||||
metadata: Metadata to assign to the cron job runs.
|
||||
@@ -350,6 +358,7 @@ class CronClient:
|
||||
after execution. 'keep' creates a new thread for each execution but does not
|
||||
clean them up.
|
||||
enabled: Enable or disable the cron job.
|
||||
timezone: IANA timezone for the cron schedule. Accepts a string (e.g. 'America/New_York') or a ``datetime.tzinfo`` instance (e.g. ``ZoneInfo("America/New_York")``).
|
||||
stream_mode: The stream mode(s) to use.
|
||||
stream_subgraphs: Whether to stream output from subgraphs.
|
||||
stream_resumable: Whether to persist the stream chunks in order to resume the stream later.
|
||||
@@ -384,6 +393,7 @@ class CronClient:
|
||||
"interrupt_after": interrupt_after,
|
||||
"on_run_completed": on_run_completed,
|
||||
"enabled": enabled,
|
||||
"timezone": _resolve_timezone(timezone),
|
||||
"stream_mode": stream_mode,
|
||||
"stream_subgraphs": stream_subgraphs,
|
||||
"stream_resumable": stream_resumable,
|
||||
|
||||
@@ -6,13 +6,17 @@ import functools
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, cast
|
||||
from datetime import tzinfo
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
import httpx
|
||||
|
||||
import langgraph_sdk
|
||||
from langgraph_sdk.schema import RunCreateMetadata
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
RESERVED_HEADERS = ("x-api-key",)
|
||||
|
||||
NOT_PROVIDED = cast(None, object())
|
||||
@@ -125,6 +129,35 @@ def _sse_to_v2_dict(event: str, data: Any) -> dict[str, Any] | None:
|
||||
return result
|
||||
|
||||
|
||||
def _resolve_timezone(tz: str | tzinfo | ZoneInfo | None) -> str | None:
|
||||
"""Convert a timezone argument to an IANA timezone string.
|
||||
|
||||
Accepts:
|
||||
- A string (returned as-is, assumed to be an IANA timezone name)
|
||||
- A ``datetime.tzinfo`` instance (e.g. ``zoneinfo.ZoneInfo("America/New_York")``,
|
||||
``datetime.timezone.utc``). The ``key`` attribute is used if available,
|
||||
otherwise ``tzname(None)`` is used.
|
||||
- ``None`` (returned as ``None``)
|
||||
"""
|
||||
if tz is None or isinstance(tz, str):
|
||||
return tz
|
||||
if isinstance(tz, tzinfo):
|
||||
# ZoneInfo objects have a .key attribute with the IANA name
|
||||
if hasattr(tz, "key"):
|
||||
return tz.key # type: ignore[union-attr]
|
||||
# Fall back to tzname for fixed-offset timezones like datetime.timezone.utc
|
||||
name = tz.tzname(None)
|
||||
if name is not None:
|
||||
return name
|
||||
raise ValueError(
|
||||
f"Cannot determine timezone name from {tz!r}. "
|
||||
"Use a zoneinfo.ZoneInfo instance or pass a string like 'America/New_York'."
|
||||
)
|
||||
raise TypeError(
|
||||
f"Expected str, datetime.tzinfo, or None for timezone, got {type(tz).__name__}"
|
||||
)
|
||||
|
||||
|
||||
def _provided_vals(d: Mapping[str, Any]) -> dict[str, Any]:
|
||||
return {k: v for k, v in d.items() if v is not None}
|
||||
|
||||
|
||||
@@ -4,9 +4,10 @@ from __future__ import annotations
|
||||
|
||||
import warnings
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from datetime import datetime, tzinfo
|
||||
from typing import Any
|
||||
|
||||
from langgraph_sdk._shared.utilities import _resolve_timezone
|
||||
from langgraph_sdk._sync.http import SyncHttpClient
|
||||
from langgraph_sdk.schema import (
|
||||
All,
|
||||
@@ -64,6 +65,7 @@ class SyncCronClient:
|
||||
multitask_strategy: str | None = None,
|
||||
end_time: datetime | None = None,
|
||||
enabled: bool | None = None,
|
||||
timezone: str | tzinfo | None = None,
|
||||
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
|
||||
stream_subgraphs: bool | None = None,
|
||||
stream_resumable: bool | None = None,
|
||||
@@ -78,7 +80,7 @@ class SyncCronClient:
|
||||
assistant_id: The assistant ID or graph name to use for the cron job.
|
||||
If using graph name, will default to first assistant created from that graph.
|
||||
schedule: The cron schedule to execute this job on.
|
||||
Schedules are interpreted in UTC.
|
||||
Schedules are interpreted in UTC unless a timezone is specified.
|
||||
input: The input to the graph.
|
||||
metadata: Metadata to assign to the cron job runs.
|
||||
config: The configuration for the assistant.
|
||||
@@ -92,6 +94,7 @@ class SyncCronClient:
|
||||
Must be one of 'reject', 'interrupt', 'rollback', or 'enqueue'.
|
||||
end_time: The time to stop running the cron job. If not provided, the cron job will run indefinitely.
|
||||
enabled: Whether the cron job is enabled. By default, it is considered enabled.
|
||||
timezone: IANA timezone for the cron schedule. Accepts a string (e.g. 'America/New_York') or a ``datetime.tzinfo`` instance (e.g. ``ZoneInfo("America/New_York")``).
|
||||
stream_mode: The stream mode(s) to use.
|
||||
stream_subgraphs: Whether to stream output from subgraphs.
|
||||
stream_resumable: Whether to persist the stream chunks in order to resume the stream later.
|
||||
@@ -144,6 +147,7 @@ class SyncCronClient:
|
||||
"multitask_strategy": multitask_strategy,
|
||||
"end_time": end_time.isoformat() if end_time else None,
|
||||
"enabled": enabled,
|
||||
"timezone": _resolve_timezone(timezone),
|
||||
"stream_mode": stream_mode,
|
||||
"stream_subgraphs": stream_subgraphs,
|
||||
"stream_resumable": stream_resumable,
|
||||
@@ -174,6 +178,7 @@ class SyncCronClient:
|
||||
multitask_strategy: str | None = None,
|
||||
end_time: datetime | None = None,
|
||||
enabled: bool | None = None,
|
||||
timezone: str | tzinfo | None = None,
|
||||
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
|
||||
stream_subgraphs: bool | None = None,
|
||||
stream_resumable: bool | None = None,
|
||||
@@ -187,7 +192,7 @@ class SyncCronClient:
|
||||
assistant_id: The assistant ID or graph name to use for the cron job.
|
||||
If using graph name, will default to first assistant created from that graph.
|
||||
schedule: The cron schedule to execute this job on.
|
||||
Schedules are interpreted in UTC.
|
||||
Schedules are interpreted in UTC unless a timezone is specified.
|
||||
input: The input to the graph.
|
||||
metadata: Metadata to assign to the cron job runs.
|
||||
config: The configuration for the assistant.
|
||||
@@ -205,6 +210,7 @@ class SyncCronClient:
|
||||
Must be one of 'reject', 'interrupt', 'rollback', or 'enqueue'.
|
||||
end_time: The time to stop running the cron job. If not provided, the cron job will run indefinitely.
|
||||
enabled: Whether the cron job is enabled. By default, it is considered enabled.
|
||||
timezone: IANA timezone for the cron schedule. Accepts a string (e.g. 'America/New_York') or a ``datetime.tzinfo`` instance (e.g. ``ZoneInfo("America/New_York")``).
|
||||
stream_mode: The stream mode(s) to use.
|
||||
stream_subgraphs: Whether to stream output from subgraphs.
|
||||
stream_resumable: Whether to persist the stream chunks in order to resume the stream later.
|
||||
@@ -259,6 +265,7 @@ class SyncCronClient:
|
||||
"multitask_strategy": multitask_strategy,
|
||||
"end_time": end_time.isoformat() if end_time else None,
|
||||
"enabled": enabled,
|
||||
"timezone": _resolve_timezone(timezone),
|
||||
"stream_mode": stream_mode,
|
||||
"stream_subgraphs": stream_subgraphs,
|
||||
"stream_resumable": stream_resumable,
|
||||
@@ -313,6 +320,7 @@ class SyncCronClient:
|
||||
interrupt_after: All | list[str] | None = None,
|
||||
on_run_completed: OnCompletionBehavior | None = None,
|
||||
enabled: bool | None = None,
|
||||
timezone: str | tzinfo | None = None,
|
||||
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
|
||||
stream_subgraphs: bool | None = None,
|
||||
stream_resumable: bool | None = None,
|
||||
@@ -325,7 +333,7 @@ class SyncCronClient:
|
||||
Args:
|
||||
cron_id: The cron ID to update.
|
||||
schedule: The cron schedule to execute this job on.
|
||||
Schedules are interpreted in UTC.
|
||||
Schedules are interpreted in UTC unless a timezone is specified.
|
||||
end_time: The end date to stop running the cron.
|
||||
input: The input to the graph.
|
||||
metadata: Metadata to assign to the cron job runs.
|
||||
@@ -339,6 +347,7 @@ class SyncCronClient:
|
||||
after execution. 'keep' creates a new thread for each execution but does not
|
||||
clean them up.
|
||||
enabled: Enable or disable the cron job.
|
||||
timezone: IANA timezone for the cron schedule. Accepts a string (e.g. 'America/New_York') or a ``datetime.tzinfo`` instance (e.g. ``ZoneInfo("America/New_York")``).
|
||||
stream_mode: The stream mode(s) to use.
|
||||
stream_subgraphs: Whether to stream output from subgraphs.
|
||||
stream_resumable: Whether to persist the stream chunks in order to resume the stream later.
|
||||
@@ -373,6 +382,7 @@ class SyncCronClient:
|
||||
"interrupt_after": interrupt_after,
|
||||
"on_run_completed": on_run_completed,
|
||||
"enabled": enabled,
|
||||
"timezone": _resolve_timezone(timezone),
|
||||
"stream_mode": stream_mode,
|
||||
"stream_subgraphs": stream_subgraphs,
|
||||
"stream_resumable": stream_resumable,
|
||||
|
||||
@@ -58,6 +58,7 @@ StreamMode = Literal[
|
||||
"debug",
|
||||
"custom",
|
||||
"messages-tuple",
|
||||
"compact",
|
||||
]
|
||||
"""
|
||||
Defines the mode of streaming:
|
||||
@@ -69,6 +70,7 @@ Defines the mode of streaming:
|
||||
- "tasks": Stream task start and finish events.
|
||||
- "debug": Stream detailed debug information.
|
||||
- "custom": Stream custom events.
|
||||
- "compact": Enable compact streaming payloads for other selected modes.
|
||||
"""
|
||||
|
||||
DisconnectMode = Literal["cancel", "continue"]
|
||||
@@ -385,6 +387,8 @@ class Cron(TypedDict):
|
||||
"""The end date to stop running the cron."""
|
||||
schedule: str
|
||||
"""The schedule to run, cron format."""
|
||||
timezone: str | None
|
||||
"""IANA timezone for the cron schedule (e.g. 'America/New_York'). Defaults to null, which is treated as UTC."""
|
||||
created_at: datetime
|
||||
"""The time the cron was created."""
|
||||
updated_at: datetime
|
||||
@@ -406,6 +410,8 @@ class CronUpdate(TypedDict, total=False):
|
||||
|
||||
schedule: str
|
||||
"""The cron schedule to execute this job on."""
|
||||
timezone: str
|
||||
"""IANA timezone for the cron schedule (e.g. 'America/New_York')."""
|
||||
end_time: datetime
|
||||
"""The end date to stop running the cron."""
|
||||
input: Input
|
||||
@@ -482,6 +488,7 @@ CronSelectField = Literal[
|
||||
"thread_id",
|
||||
"end_time",
|
||||
"schedule",
|
||||
"timezone",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"user_id",
|
||||
@@ -728,6 +735,26 @@ class ValuesStreamPart(TypedDict):
|
||||
"""List of interrupts that occurred during this step."""
|
||||
|
||||
|
||||
class ValuesPatchPayload(TypedDict):
|
||||
"""Incremental patch payload for subgraph `values` updates."""
|
||||
|
||||
values: dict[str, Any]
|
||||
"""Only the changed fields since the previous `values` event for this namespace."""
|
||||
deleted_keys: NotRequired[list[str]]
|
||||
"""Optional list of keys that were removed from the previous values snapshot."""
|
||||
|
||||
|
||||
class ValuesPatchStreamPart(TypedDict):
|
||||
"""Stream part emitted for incremental subgraph value patches (`values-patch`)."""
|
||||
|
||||
type: Literal["values-patch"]
|
||||
"""Stream part type discriminator."""
|
||||
ns: list[str]
|
||||
"""Namespace path of the emitting node (empty for root graph)."""
|
||||
data: ValuesPatchPayload
|
||||
"""Incremental state patch for the namespace."""
|
||||
|
||||
|
||||
class UpdatesStreamPart(TypedDict):
|
||||
"""Stream part emitted for `stream_mode="updates"`."""
|
||||
|
||||
@@ -840,6 +867,7 @@ class MetadataStreamPart(TypedDict):
|
||||
|
||||
StreamPartV2 = (
|
||||
ValuesStreamPart
|
||||
| ValuesPatchStreamPart
|
||||
| UpdatesStreamPart
|
||||
| MessagesPartialStreamPart
|
||||
| MessagesCompleteStreamPart
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Iterator, Sequence
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
@@ -8,7 +9,9 @@ import httpx
|
||||
import pytest
|
||||
from typing_extensions import assert_type
|
||||
|
||||
from langgraph_sdk._async.runs import RunsClient
|
||||
from langgraph_sdk._shared.utilities import _sse_to_v2_dict
|
||||
from langgraph_sdk._sync.runs import SyncRunsClient
|
||||
from langgraph_sdk.client import HttpClient, SyncHttpClient
|
||||
from langgraph_sdk.schema import (
|
||||
CheckpointPayload,
|
||||
@@ -24,6 +27,7 @@ from langgraph_sdk.schema import (
|
||||
TaskResultPayload,
|
||||
TasksStreamPart,
|
||||
UpdatesStreamPart,
|
||||
ValuesPatchStreamPart,
|
||||
ValuesStreamPart,
|
||||
)
|
||||
from langgraph_sdk.sse import BytesLike, BytesLineDecoder, SSEDecoder
|
||||
@@ -375,6 +379,81 @@ def test_sse_to_v2_dict_values_with_interrupts() -> None:
|
||||
assert "__interrupt__" not in result["data"]
|
||||
|
||||
|
||||
def test_sse_to_v2_dict_values_patch() -> None:
|
||||
payload = {"values": {"count": 2}, "deleted_keys": ["stale"]}
|
||||
result = _sse_to_v2_dict("values-patch|tools:call_1", payload)
|
||||
assert result is not None
|
||||
_assert_v2_shape(result)
|
||||
assert result == {
|
||||
"type": "values-patch",
|
||||
"ns": ["tools:call_1"],
|
||||
"data": {"values": {"count": 2}, "deleted_keys": ["stale"]},
|
||||
"interrupts": [],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_runs_stream_includes_compact_mode():
|
||||
async def handler(request: httpx.Request) -> httpx.Response:
|
||||
assert request.method == "POST"
|
||||
assert request.url.path == "/runs/stream"
|
||||
body = json.loads(request.content)
|
||||
assert body["stream_mode"] == ["values", "compact"]
|
||||
return httpx.Response(
|
||||
200,
|
||||
headers={"Content-Type": "text/event-stream"},
|
||||
content=b"event: end\ndata: null\n\n",
|
||||
)
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
async with httpx.AsyncClient(
|
||||
transport=transport, base_url="https://example.com"
|
||||
) as client:
|
||||
runs_client = RunsClient(HttpClient(client))
|
||||
parts = [
|
||||
part
|
||||
async for part in runs_client.stream(
|
||||
thread_id=None,
|
||||
assistant_id="agent",
|
||||
input={"messages": []},
|
||||
stream_mode=["values", "compact"],
|
||||
)
|
||||
]
|
||||
|
||||
assert len(parts) == 1
|
||||
assert parts[0].event == "end"
|
||||
assert parts[0].data is None
|
||||
|
||||
|
||||
def test_sync_runs_stream_includes_compact_mode():
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
assert request.method == "POST"
|
||||
assert request.url.path == "/runs/stream"
|
||||
body = json.loads(request.content)
|
||||
assert body["stream_mode"] == ["values", "compact"]
|
||||
return httpx.Response(
|
||||
200,
|
||||
headers={"Content-Type": "text/event-stream"},
|
||||
content=b"event: end\ndata: null\n\n",
|
||||
)
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
with httpx.Client(transport=transport, base_url="https://example.com") as client:
|
||||
runs_client = SyncRunsClient(SyncHttpClient(client))
|
||||
parts = list(
|
||||
runs_client.stream(
|
||||
thread_id=None,
|
||||
assistant_id="agent",
|
||||
input={"messages": []},
|
||||
stream_mode=["values", "compact"],
|
||||
)
|
||||
)
|
||||
|
||||
assert len(parts) == 1
|
||||
assert parts[0].event == "end"
|
||||
assert parts[0].data is None
|
||||
|
||||
|
||||
# --- client-side v2 stream wrapping ---
|
||||
|
||||
|
||||
@@ -448,6 +527,8 @@ def _check_v2_type_narrowing(part: StreamPartV2) -> None:
|
||||
if part["type"] == "values":
|
||||
assert_type(part, ValuesStreamPart)
|
||||
assert_type(part["data"], dict[str, Any])
|
||||
elif part["type"] == "values-patch":
|
||||
assert_type(part, ValuesPatchStreamPart)
|
||||
elif part["type"] == "updates":
|
||||
assert_type(part, UpdatesStreamPart)
|
||||
assert_type(part["data"], dict[str, Any])
|
||||
|
||||
Generated
+1
-1
@@ -265,7 +265,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph"
|
||||
version = "1.1.0"
|
||||
version = "1.1.2"
|
||||
source = { editable = "../langgraph" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
Reference in New Issue
Block a user