mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-11 18:55:17 +02:00
Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
741c6f8d50 | ||
|
|
37a5504433 | ||
|
|
f6286dce38 | ||
|
|
93a0dfec08 | ||
|
|
56834787eb | ||
|
|
96199c4fee | ||
|
|
682814e944 | ||
|
|
210c4b3877 | ||
|
|
81489ab858 |
+211
-68
@@ -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.
|
||||
|
||||
|
||||
@@ -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,
|
||||
@@ -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
|
||||
@@ -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.1"
|
||||
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])
|
||||
|
||||
Generated
+1
-1
@@ -1367,7 +1367,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph"
|
||||
version = "1.1.1"
|
||||
version = "1.1.2"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
Generated
+1
-1
@@ -268,7 +268,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph"
|
||||
version = "1.1.1"
|
||||
version = "1.1.2"
|
||||
source = { editable = "../langgraph" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.1"
|
||||
version = "1.1.2"
|
||||
source = { editable = "../langgraph" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
Reference in New Issue
Block a user