Compare commits

...
Author SHA1 Message Date
Christian Bromann 741c6f8d50 cr 2026-03-18 15:46:35 -07:00
Christian Bromann 37a5504433 move to stream mode 2026-03-12 15:08:26 -07:00
Christian Bromann f6286dce38 feat(langgraph): add protocol improvements for better streaming 2026-03-12 15:08:25 -07:00
hari-dhanushkodiandGitHub 93a0dfec08 fix(cli): cleanup client creation code (#7140)
make following fixes for deploy cli:
- allow logs command to read in `_DEPLOYMENT_NAME_ENV`
- change prompt for api key to be more user friendly
- refactor prompt for org scoped keys
- make sure --config option is documented in deploy help
2026-03-12 11:40:49 -07:00
Sydney RunkleandGitHub 56834787eb release(langgraph): 1.1.2 (#7135) 2026-03-12 13:07:16 -04:00
96199c4fee feat(cli): add langgraph deploy logs subcommand (#7100)
Add a `langgraph deploy logs` subcommand to fetch build and
deploy/server logs from LangSmith deployments.

## Changes

- **`host_backend.py`**: Add `get_build_logs()` and `get_deploy_logs()`
methods, plus `langsmith_url` support.
- **`helpers.py`** (new): Log formatting (`format_log_entry`,
`format_timestamp`, `level_fg`) and resolution helpers
(`resolve_deployment_id`).

## Usage

```bash
# Deploy logs (latest 10)
langgraph deploy logs --deployment-id <id> --limit 10

# Error logs by name
langgraph deploy logs --name gtm-agent --level ERROR

# Build logs
langgraph deploy logs --deployment-id <id> --type build

# Tail logs (Ctrl+C to stop)
langgraph deploy logs --name gtm-agent --follow
```

## Tests

56 unit tests passing covering host backend methods, log formatting, and
helper functions.

---------

Co-authored-by: hari-dhanushkodi <hari@langchain.dev>
2026-03-12 09:01:44 -07:00
Sydney RunkleandGitHub 682814e944 fix: stream part generic order (#7134)
this is technically breaking but is a fix given dependency of output t
on state t
2026-03-12 11:58:04 -04:00
Sydney RunkleandGitHub 210c4b3877 feat: add context for remote graph api (#7132)
underlying sdk already supports, just exposing for remote graph
2026-03-12 11:57:06 -04:00
Mason DaughertyandGitHub 81489ab858 chore: fix typo (#7131) 2026-03-12 14:44:24 +00:00
18 changed files with 1029 additions and 102 deletions
+211 -68
View File
@@ -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,6 +27,7 @@ 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
@@ -296,6 +298,17 @@ OPT_HOST_API_KEY = click.option(
),
)
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",
@@ -684,15 +697,7 @@ def _deploy_base_options(
def _apply(target: Callable) -> Callable:
decorators = [
OPT_HOST_API_KEY,
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."
),
),
OPT_HOST_DEPLOYMENT_NAME,
click.option(
"--deployment-id",
help=(
@@ -758,7 +763,7 @@ def _deploy_base_options(
"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, allow_extra_args=True),
invoke_without_command=True, # allow `deploy` click group to execute without command
@@ -807,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:
@@ -852,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", []):
@@ -1178,7 +1140,11 @@ def _create_host_backend_client(
resolved_api_key = val
break
if not resolved_api_key:
resolved_api_key = click.prompt("Host API key", hide_input=True)
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)
@@ -1186,6 +1152,13 @@ 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:
@@ -1199,9 +1172,7 @@ def _call_host_backend_with_optional_tenant(
fg="yellow",
)
tenant_id = click.prompt("Workspace ID")
client = HostBackendClient(
client._base_url, client._api_key, tenant_id=tenant_id
)
client._client.headers["X-Tenant-ID"] = tenant_id
return operation(client)
raise
@@ -1218,9 +1189,7 @@ def deploy_list(api_key: str | None, host_url: str | None, name_contains: str) -
client = _create_host_backend_client(host_url, api_key)
response = _call_host_backend_with_optional_tenant(
client,
lambda current_client: current_client.list_deployments(
name_contains=name_contains
),
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)]
@@ -1263,7 +1232,7 @@ def deploy_delete(
client = _create_host_backend_client(host_url, api_key)
_call_host_backend_with_optional_tenant(
client,
lambda current_client: current_client.delete_deployment(deployment_id),
lambda c: c.delete_deployment(deployment_id),
)
click.secho(f"Deleted deployment {deployment_id}.", fg="green")
@@ -1294,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.
+59
View 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
+27 -2
View File
@@ -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,
@@ -116,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)
@@ -182,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
+10 -1
View File
@@ -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(
+6 -4
View File
@@ -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(
+2 -2
View File
@@ -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
+78 -14
View File
@@ -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,
+14 -6
View File
@@ -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
```
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph"
version = "1.1.1"
version = "1.1.2"
description = "Building stateful, multi-actor applications with LLMs"
authors = []
requires-python = ">=3.10"
+335
View File
@@ -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"
)
+83 -1
View File
@@ -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])
+1 -1
View File
@@ -1367,7 +1367,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "1.1.1"
version = "1.1.2"
source = { editable = "." }
dependencies = [
{ name = "langchain-core" },
+1 -1
View File
@@ -268,7 +268,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "1.1.1"
version = "1.1.2"
source = { editable = "../langgraph" }
dependencies = [
{ name = "langchain-core" },
+23
View File
@@ -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"]
@@ -733,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"`."""
@@ -845,6 +867,7 @@ class MetadataStreamPart(TypedDict):
StreamPartV2 = (
ValuesStreamPart
| ValuesPatchStreamPart
| UpdatesStreamPart
| MessagesPartialStreamPart
| MessagesCompleteStreamPart
+81
View File
@@ -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])
+1 -1
View File
@@ -265,7 +265,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "1.1.1"
version = "1.1.2"
source = { editable = "../langgraph" }
dependencies = [
{ name = "langchain-core" },