mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-02 14:35:18 +02:00
Compare commits
4
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
dc53c4acc0 | ||
|
|
9100f2c682 | ||
|
|
79befe67ba | ||
|
|
9af25217c3 |
@@ -0,0 +1,132 @@
|
||||
# delta-channel-dump
|
||||
|
||||
Recover messages (and other channels) from a Postgres-backed LangGraph thread
|
||||
written by **langgraph >= 1.2** (DeltaChannel format) — including **LangGraph
|
||||
Server / langgraph-api** deployments on the Postgres runtime, **deepagents
|
||||
0.6.x**, or any OSS app using `PostgresSaver` — before rolling back to an older
|
||||
runtime such as **deepagents 0.5.x / langgraph < 1.2**.
|
||||
|
||||
langgraph-api uses the same `checkpoints` / `checkpoint_blobs` / `checkpoint_writes`
|
||||
schema as OSS `checkpoint-postgres`; this script reads those tables directly.
|
||||
|
||||
On the older runtime, `add_messages` does not understand the `EXT_DELTA_SNAPSHOT`
|
||||
msgpack ext code and silently returns an empty list for affected channels. This
|
||||
tool reads the raw checkpoint blobs from Postgres and emits JSON you can inspect
|
||||
and re-apply via `update_state` (LangGraph Server SDK) or `graph.update_state`
|
||||
(OSS).
|
||||
|
||||
## Install
|
||||
|
||||
```bash
|
||||
pip install "psycopg[binary]" ormsgpack
|
||||
```
|
||||
|
||||
## Run
|
||||
|
||||
```bash
|
||||
export DATABASE_URI=postgres://user:pass@host:5432/dbname
|
||||
python3 dump.py \
|
||||
--thread-id <uuid> \
|
||||
--channel messages \
|
||||
[--channel files ...] \
|
||||
[--checkpoint-id <uuid>] \
|
||||
[--checkpoint-ns ""] \
|
||||
--output recovery.json
|
||||
```
|
||||
|
||||
- `--thread-id` (required): thread UUID
|
||||
- `--channel` (required, repeatable): channel names to recover
|
||||
- `--checkpoint-id` (optional): target checkpoint; defaults to latest
|
||||
- `--checkpoint-ns` (optional): namespace; defaults to `""`
|
||||
- `--database-uri` (optional): Postgres URI; defaults to `DATABASE_URI` env var
|
||||
- `--output` (optional): output file; defaults to stdout
|
||||
|
||||
## Output
|
||||
|
||||
```json
|
||||
{
|
||||
"thread_id": "...",
|
||||
"checkpoint_ns": "",
|
||||
"target_checkpoint_id": "...",
|
||||
"parent_checkpoint_id": "...",
|
||||
"channels": {
|
||||
"messages": {
|
||||
"delta_kind": "snapshot",
|
||||
"seed_checkpoint_id": "...",
|
||||
"seed_version": "...",
|
||||
"seed": [{ "type": "ai", "content": "...", "id": "ai-0" }],
|
||||
"writes": [
|
||||
{
|
||||
"checkpoint_id": "...",
|
||||
"task_id": "...",
|
||||
"idx": 0,
|
||||
"value": [{ "type": "ai", "content": "...", "id": "ai-10" }]
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
`delta_kind` is one of:
|
||||
|
||||
- `snapshot` — DeltaChannel snapshot blob (`channel_values[ch] == true`)
|
||||
- `legacy_plain` — pre-DeltaChannel inline or blob value
|
||||
- `no_seed` — walked to root without finding a populated ancestor
|
||||
|
||||
`writes` are ordered oldest-to-newest (the order a reducer would replay them).
|
||||
|
||||
## Reducing back to a single list
|
||||
|
||||
For deepagents-style messages, combine seed and writes, then deduplicate:
|
||||
|
||||
```python
|
||||
import json
|
||||
|
||||
data = json.load(open("recovery.json"))
|
||||
ch = data["channels"]["messages"]
|
||||
messages = list(ch["seed"] or [])
|
||||
for w in ch["writes"]:
|
||||
messages.extend(w["value"] or [])
|
||||
|
||||
# Dedup by id, keep last; drop RemoveMessage tombstones
|
||||
by_id = {}
|
||||
for m in messages:
|
||||
if isinstance(m, dict) and m.get("type") == "remove":
|
||||
by_id.pop(m.get("id"), None)
|
||||
elif isinstance(m, dict) and m.get("id"):
|
||||
by_id[m["id"]] = m
|
||||
else:
|
||||
by_id[id(m)] = m
|
||||
reduced = list(by_id.values())
|
||||
```
|
||||
|
||||
This approximates `_messages_delta_reducer` semantics; adjust for your graph.
|
||||
|
||||
## Re-applying
|
||||
|
||||
```python
|
||||
from langgraph_sdk import get_client
|
||||
|
||||
client = get_client(url="http://localhost:8123")
|
||||
await client.threads.update_state(
|
||||
thread_id,
|
||||
values={"messages": reduced},
|
||||
)
|
||||
```
|
||||
|
||||
Review the recovered JSON before calling `update_state`. This tool is
|
||||
read-only and intentionally does not mutate the database.
|
||||
|
||||
## Scope / non-goals (v1)
|
||||
|
||||
- **Postgres only** — OSS `PostgresSaver` or langgraph-api Postgres runtime; not
|
||||
inmem, gRPC core, Mongo, or Redis checkpointer backends
|
||||
- **No AES decryption** (`LANGGRAPH_AES_KEY`) or custom encryption
|
||||
- **No reducer** — raw seed + writes only
|
||||
- **No automatic `update_state`** — operator applies manually
|
||||
|
||||
## Copying
|
||||
|
||||
`dump.py` is self-contained. Copy it anywhere; only `psycopg[binary]` and
|
||||
`ormsgpack` are required at runtime. No langgraph imports.
|
||||
Executable
+469
@@ -0,0 +1,469 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Recover delta-channel state from a Postgres-backed LangGraph thread.
|
||||
|
||||
Works with OSS ``PostgresSaver`` and LangGraph Server / langgraph-api on the
|
||||
Postgres runtime (same checkpoint schema). Use after rolling back from
|
||||
langgraph >= 1.2 / deepagents 0.6.x to an older runtime that does not
|
||||
understand ``EXT_DELTA_SNAPSHOT`` msgpack blobs. The script walks the checkpoint
|
||||
parent chain, decodes msgpack blobs, and emits a JSON dump of per-channel
|
||||
``seed`` plus oldest-to-newest ``writes``. Apply the recovered values manually
|
||||
via ``client.threads.update_state(...)`` (Server) or ``graph.update_state``
|
||||
(OSS).
|
||||
|
||||
Install::
|
||||
|
||||
pip install "psycopg[binary]" ormsgpack
|
||||
|
||||
Run::
|
||||
|
||||
export DATABASE_URI=postgres://...
|
||||
python3 dump.py --thread-id <uuid> --channel messages --output recovery.json
|
||||
|
||||
Scope (v1): Postgres only; no AES/custom encryption; no reducer application.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import ormsgpack
|
||||
import psycopg
|
||||
|
||||
# LangGraph msgpack EXT type codes (langgraph/checkpoint/serde/jsonplus.py).
|
||||
EXT_CONSTRUCTOR_SINGLE_ARG = 0
|
||||
EXT_CONSTRUCTOR_POS_ARGS = 1
|
||||
EXT_CONSTRUCTOR_KW_ARGS = 2
|
||||
EXT_METHOD_SINGLE_ARG = 3
|
||||
EXT_PYDANTIC_V1 = 4
|
||||
EXT_PYDANTIC_V2 = 5
|
||||
EXT_NUMPY_ARRAY = 6
|
||||
EXT_DELTA_SNAPSHOT = 7
|
||||
|
||||
_MSGPACK_OPTION = ormsgpack.OPT_NON_STR_KEYS
|
||||
|
||||
|
||||
def ext_hook(code: int, data: bytes) -> Any:
|
||||
"""Decode LangGraph msgpack EXT payloads to JSON-friendly Python values."""
|
||||
if code == EXT_DELTA_SNAPSHOT:
|
||||
inner = ormsgpack.unpackb(data, ext_hook=ext_hook, option=_MSGPACK_OPTION)
|
||||
return {"__delta_snapshot__": inner}
|
||||
if code == EXT_CONSTRUCTOR_SINGLE_ARG:
|
||||
try:
|
||||
tup = ormsgpack.unpackb(data, ext_hook=ext_hook, option=_MSGPACK_OPTION)
|
||||
if tup[0] == "uuid" and tup[1] == "UUID":
|
||||
hex_ = tup[2]
|
||||
return (
|
||||
f"{hex_[:8]}-{hex_[8:12]}-{hex_[12:16]}-"
|
||||
f"{hex_[16:20]}-{hex_[20:]}"
|
||||
)
|
||||
return tup[2]
|
||||
except Exception:
|
||||
return None
|
||||
if code == EXT_CONSTRUCTOR_POS_ARGS:
|
||||
try:
|
||||
tup = ormsgpack.unpackb(data, ext_hook=ext_hook, option=_MSGPACK_OPTION)
|
||||
if tup[0] == "langgraph.types" and tup[1] == "Send":
|
||||
args = tup[2]
|
||||
if len(args) == 2:
|
||||
return {"__send__": {"node": args[0], "arg": args[1]}}
|
||||
return {
|
||||
"__send__": {
|
||||
"node": args[0],
|
||||
"arg": args[1],
|
||||
"timeout": args[2],
|
||||
}
|
||||
}
|
||||
return tup[2]
|
||||
except Exception:
|
||||
return None
|
||||
if code in (EXT_CONSTRUCTOR_KW_ARGS, EXT_METHOD_SINGLE_ARG):
|
||||
try:
|
||||
tup = ormsgpack.unpackb(data, ext_hook=ext_hook, option=_MSGPACK_OPTION)
|
||||
return tup[2]
|
||||
except Exception:
|
||||
return None
|
||||
if code in (EXT_PYDANTIC_V1, EXT_PYDANTIC_V2):
|
||||
try:
|
||||
tup = ormsgpack.unpackb(data, ext_hook=ext_hook, option=_MSGPACK_OPTION)
|
||||
return tup[2]
|
||||
except Exception:
|
||||
return None
|
||||
if code == EXT_NUMPY_ARRAY:
|
||||
try:
|
||||
dtype_str, shape, order, buf = ormsgpack.unpackb(
|
||||
data, ext_hook=ext_hook, option=_MSGPACK_OPTION
|
||||
)
|
||||
return {
|
||||
"__numpy_array__": {
|
||||
"dtype": dtype_str,
|
||||
"shape": shape,
|
||||
"order": order,
|
||||
"data_b64": base64.b64encode(buf).decode("ascii"),
|
||||
}
|
||||
}
|
||||
except Exception:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def decode_blob(blob_type: str, blob_bytes: bytes | None) -> Any:
|
||||
if blob_type in ("empty", "null") or blob_bytes is None:
|
||||
return None
|
||||
if blob_type == "msgpack":
|
||||
return ormsgpack.unpackb(
|
||||
blob_bytes, ext_hook=ext_hook, option=_MSGPACK_OPTION
|
||||
)
|
||||
if blob_type in ("bytes", "bytearray"):
|
||||
return base64.b64encode(blob_bytes).decode("ascii")
|
||||
raise RuntimeError(
|
||||
f"Unknown blob type {blob_type!r}. "
|
||||
"AES/custom-encrypted deployments are out of scope for v1."
|
||||
)
|
||||
|
||||
|
||||
def delta_unwrap(value: Any) -> Any:
|
||||
if isinstance(value, dict) and "__delta_snapshot__" in value:
|
||||
return value["__delta_snapshot__"]
|
||||
return value
|
||||
|
||||
|
||||
def json_default(obj: Any) -> Any:
|
||||
if isinstance(obj, (bytes, bytearray)):
|
||||
return base64.b64encode(bytes(obj)).decode("ascii")
|
||||
if isinstance(obj, uuid.UUID):
|
||||
return str(obj)
|
||||
raise TypeError(f"Object of type {type(obj).__name__} is not JSON serializable")
|
||||
|
||||
|
||||
def resolve_target_checkpoint_id(
|
||||
conn: psycopg.Connection[Any],
|
||||
thread_id: str,
|
||||
checkpoint_ns: str,
|
||||
checkpoint_id: str | None,
|
||||
) -> str:
|
||||
if checkpoint_id is not None:
|
||||
return checkpoint_id
|
||||
row = conn.execute(
|
||||
"""
|
||||
SELECT checkpoint_id::text
|
||||
FROM checkpoints
|
||||
WHERE thread_id = %s AND checkpoint_ns = %s
|
||||
ORDER BY checkpoint_id DESC
|
||||
LIMIT 1
|
||||
""",
|
||||
(thread_id, checkpoint_ns),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
raise SystemExit(
|
||||
f"No checkpoints found for thread_id={thread_id!r} "
|
||||
f"checkpoint_ns={checkpoint_ns!r}"
|
||||
)
|
||||
return row[0]
|
||||
|
||||
|
||||
def _load_checkpoint(
|
||||
conn: psycopg.Connection[Any],
|
||||
thread_id: str,
|
||||
checkpoint_ns: str,
|
||||
checkpoint_id: str,
|
||||
) -> tuple[dict[str, Any], str | None] | None:
|
||||
row = conn.execute(
|
||||
"""
|
||||
SELECT checkpoint, parent_checkpoint_id::text
|
||||
FROM checkpoints
|
||||
WHERE thread_id = %s AND checkpoint_ns = %s AND checkpoint_id = %s
|
||||
""",
|
||||
(thread_id, checkpoint_ns, checkpoint_id),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return None
|
||||
return row[0], row[1]
|
||||
|
||||
|
||||
def _load_seed(
|
||||
conn: psycopg.Connection[Any],
|
||||
*,
|
||||
thread_id: str,
|
||||
checkpoint_ns: str,
|
||||
channel: str,
|
||||
checkpoint_id: str,
|
||||
channel_values: dict[str, Any],
|
||||
channel_versions: dict[str, str],
|
||||
) -> dict[str, Any]:
|
||||
cv = channel_values[channel]
|
||||
version = channel_versions.get(channel)
|
||||
if cv is True:
|
||||
blob_row = conn.execute(
|
||||
"""
|
||||
SELECT type, blob
|
||||
FROM checkpoint_blobs
|
||||
WHERE thread_id = %s AND checkpoint_ns = %s
|
||||
AND channel = %s AND version = %s
|
||||
""",
|
||||
(thread_id, checkpoint_ns, channel, version),
|
||||
).fetchone()
|
||||
seed_value = None
|
||||
if blob_row is not None and blob_row[0] != "empty":
|
||||
seed_value = delta_unwrap(decode_blob(blob_row[0], blob_row[1]))
|
||||
return {
|
||||
"delta_kind": "snapshot",
|
||||
"seed_checkpoint_id": checkpoint_id,
|
||||
"seed_version": version,
|
||||
"seed": seed_value,
|
||||
}
|
||||
if isinstance(cv, (int, float, str, bool)) or cv is None:
|
||||
return {
|
||||
"delta_kind": "legacy_plain",
|
||||
"seed_checkpoint_id": checkpoint_id,
|
||||
"seed_version": version,
|
||||
"seed": cv,
|
||||
}
|
||||
blob_row = conn.execute(
|
||||
"""
|
||||
SELECT type, blob
|
||||
FROM checkpoint_blobs
|
||||
WHERE thread_id = %s AND checkpoint_ns = %s
|
||||
AND channel = %s AND version = %s
|
||||
""",
|
||||
(thread_id, checkpoint_ns, channel, version),
|
||||
).fetchone()
|
||||
seed_value = cv if blob_row is None else decode_blob(blob_row[0], blob_row[1])
|
||||
return {
|
||||
"delta_kind": "legacy_plain",
|
||||
"seed_checkpoint_id": checkpoint_id,
|
||||
"seed_version": version,
|
||||
"seed": seed_value,
|
||||
}
|
||||
|
||||
|
||||
def _load_writes_for_checkpoint(
|
||||
conn: psycopg.Connection[Any],
|
||||
*,
|
||||
thread_id: str,
|
||||
checkpoint_ns: str,
|
||||
checkpoint_id: str,
|
||||
channel: str,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Load writes for one checkpoint, newest-first by ``(task_id, idx)``.
|
||||
|
||||
``walk_channel`` reverses the accumulated flat list before returning;
|
||||
DESC here yields oldest-first within each checkpoint in the final output,
|
||||
matching ``PostgresSaver._build_delta_channels_writes_history``.
|
||||
"""
|
||||
rows = conn.execute(
|
||||
"""
|
||||
SELECT task_id::text, idx, type, blob
|
||||
FROM checkpoint_writes
|
||||
WHERE thread_id = %s AND checkpoint_ns = %s
|
||||
AND checkpoint_id = %s AND channel = %s
|
||||
ORDER BY task_id DESC, idx DESC
|
||||
""",
|
||||
(thread_id, checkpoint_ns, checkpoint_id, channel),
|
||||
).fetchall()
|
||||
return [
|
||||
{
|
||||
"checkpoint_id": checkpoint_id,
|
||||
"task_id": task_id,
|
||||
"idx": idx,
|
||||
"value": decode_blob(blob_type, blob_bytes),
|
||||
}
|
||||
for task_id, idx, blob_type, blob_bytes in rows
|
||||
]
|
||||
|
||||
|
||||
def walk_channel(
|
||||
conn: psycopg.Connection[Any],
|
||||
*,
|
||||
thread_id: str,
|
||||
checkpoint_ns: str,
|
||||
start_checkpoint_id: str | None,
|
||||
channel: str,
|
||||
) -> dict[str, Any]:
|
||||
"""Walk one channel's parent chain from target.parent backward to seed."""
|
||||
chain_writes_newest_first: list[dict[str, Any]] = []
|
||||
cur = start_checkpoint_id
|
||||
while cur is not None:
|
||||
loaded = _load_checkpoint(conn, thread_id, checkpoint_ns, cur)
|
||||
if loaded is None:
|
||||
break
|
||||
checkpoint_json, parent_id = loaded
|
||||
channel_values = checkpoint_json.get("channel_values") or {}
|
||||
channel_versions = checkpoint_json.get("channel_versions") or {}
|
||||
|
||||
chain_writes_newest_first.extend(
|
||||
_load_writes_for_checkpoint(
|
||||
conn,
|
||||
thread_id=thread_id,
|
||||
checkpoint_ns=checkpoint_ns,
|
||||
checkpoint_id=cur,
|
||||
channel=channel,
|
||||
)
|
||||
)
|
||||
|
||||
if channel in channel_values:
|
||||
seed = _load_seed(
|
||||
conn,
|
||||
thread_id=thread_id,
|
||||
checkpoint_ns=checkpoint_ns,
|
||||
channel=channel,
|
||||
checkpoint_id=cur,
|
||||
channel_values=channel_values,
|
||||
channel_versions=channel_versions,
|
||||
)
|
||||
seed["writes"] = list(reversed(chain_writes_newest_first))
|
||||
return seed
|
||||
|
||||
cur = parent_id
|
||||
|
||||
return {
|
||||
"delta_kind": "no_seed",
|
||||
"seed_checkpoint_id": None,
|
||||
"seed_version": None,
|
||||
"seed": None,
|
||||
"writes": list(reversed(chain_writes_newest_first)),
|
||||
}
|
||||
|
||||
|
||||
def walk_parent_chain(
|
||||
conn: psycopg.Connection[Any],
|
||||
*,
|
||||
thread_id: str,
|
||||
checkpoint_ns: str,
|
||||
target_checkpoint_id: str,
|
||||
channels: list[str],
|
||||
) -> tuple[str | None, dict[str, dict[str, Any]]]:
|
||||
"""Walk parent chain and return per-channel seed + writes (oldest-first)."""
|
||||
parent_row = conn.execute(
|
||||
"""
|
||||
SELECT parent_checkpoint_id::text
|
||||
FROM checkpoints
|
||||
WHERE thread_id = %s AND checkpoint_ns = %s AND checkpoint_id = %s
|
||||
""",
|
||||
(thread_id, checkpoint_ns, target_checkpoint_id),
|
||||
).fetchone()
|
||||
if parent_row is None:
|
||||
raise SystemExit(
|
||||
f"Checkpoint {target_checkpoint_id!r} not found for thread {thread_id!r}"
|
||||
)
|
||||
parent_checkpoint_id = parent_row[0]
|
||||
|
||||
result = {
|
||||
ch: walk_channel(
|
||||
conn,
|
||||
thread_id=thread_id,
|
||||
checkpoint_ns=checkpoint_ns,
|
||||
start_checkpoint_id=parent_checkpoint_id,
|
||||
channel=ch,
|
||||
)
|
||||
for ch in channels
|
||||
}
|
||||
return parent_checkpoint_id, result
|
||||
|
||||
|
||||
def build_output(
|
||||
*,
|
||||
thread_id: str,
|
||||
checkpoint_ns: str,
|
||||
target_checkpoint_id: str,
|
||||
parent_checkpoint_id: str | None,
|
||||
channels: dict[str, dict[str, Any]],
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"thread_id": thread_id,
|
||||
"checkpoint_ns": checkpoint_ns,
|
||||
"target_checkpoint_id": target_checkpoint_id,
|
||||
"parent_checkpoint_id": parent_checkpoint_id,
|
||||
"channels": channels,
|
||||
}
|
||||
|
||||
|
||||
def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description=(
|
||||
"Recover delta-channel seed + writes from Postgres checkpoint data."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--thread-id", required=True, help="Thread UUID")
|
||||
parser.add_argument(
|
||||
"--channel",
|
||||
action="append",
|
||||
required=True,
|
||||
dest="channels",
|
||||
help="Channel name (repeatable)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpoint-id",
|
||||
default=None,
|
||||
help="Target checkpoint UUID (default: latest for thread)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpoint-ns",
|
||||
default="",
|
||||
help='Checkpoint namespace (default: "")',
|
||||
)
|
||||
parser.add_argument(
|
||||
"--database-uri",
|
||||
default=os.environ.get("DATABASE_URI"),
|
||||
help="Postgres URI (default: DATABASE_URI env var)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default="-",
|
||||
help="Output JSON file path (default: stdout)",
|
||||
)
|
||||
return parser.parse_args(argv)
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
args = parse_args(argv)
|
||||
if not args.database_uri:
|
||||
print(
|
||||
"error: --database-uri or DATABASE_URI is required",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 2
|
||||
|
||||
thread_id = str(uuid.UUID(args.thread_id))
|
||||
channels = list(dict.fromkeys(args.channels))
|
||||
|
||||
with psycopg.connect(args.database_uri) as conn:
|
||||
target_checkpoint_id = resolve_target_checkpoint_id(
|
||||
conn,
|
||||
thread_id,
|
||||
args.checkpoint_ns,
|
||||
args.checkpoint_id,
|
||||
)
|
||||
parent_checkpoint_id, channel_data = walk_parent_chain(
|
||||
conn,
|
||||
thread_id=thread_id,
|
||||
checkpoint_ns=args.checkpoint_ns,
|
||||
target_checkpoint_id=target_checkpoint_id,
|
||||
channels=channels,
|
||||
)
|
||||
|
||||
output = build_output(
|
||||
thread_id=thread_id,
|
||||
checkpoint_ns=args.checkpoint_ns,
|
||||
target_checkpoint_id=target_checkpoint_id,
|
||||
parent_checkpoint_id=parent_checkpoint_id,
|
||||
channels=channel_data,
|
||||
)
|
||||
payload = json.dumps(output, indent=2, default=json_default)
|
||||
|
||||
if args.output == "-":
|
||||
print(payload)
|
||||
else:
|
||||
with open(args.output, "w", encoding="utf-8") as f:
|
||||
f.write(payload)
|
||||
f.write("\n")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -20,6 +20,7 @@ from langchain_core.runnables.config import (
|
||||
from langgraph.checkpoint.base import CheckpointMetadata
|
||||
|
||||
from langgraph._internal._constants import (
|
||||
_CHECKPOINT_COORDINATE_KEYS,
|
||||
CONF,
|
||||
CONFIG_KEY_CHECKPOINT_ID,
|
||||
CONFIG_KEY_CHECKPOINT_MAP,
|
||||
@@ -342,6 +343,28 @@ def ensure_config(*configs: RunnableConfig | None) -> RunnableConfig:
|
||||
if _is_not_empty(v)
|
||||
},
|
||||
)
|
||||
# An explicit config that supplies its own checkpoint coordinate (a
|
||||
# thread_id, or any checkpoint_ns/checkpoint_id/checkpoint_map) is addressing
|
||||
# its own checkpoint lineage, so drop the inherited ambient configurable
|
||||
# rather than merging over it: a child graph invoked inside a parent node
|
||||
# would otherwise write its checkpoints under the parent's namespace and
|
||||
# never find them again. An explicit thread_id resets even when it equals the
|
||||
# ambient one, since a child reusing the parent's thread id still addresses
|
||||
# its own root namespace, not the parent task's. Configs that only refine
|
||||
# other keys keep the ambient and shallow-merge over it below.
|
||||
if empty.get(CONF):
|
||||
for config in configs:
|
||||
if config is None:
|
||||
continue
|
||||
explicit_configurable = config.get(CONF)
|
||||
if not explicit_configurable:
|
||||
continue
|
||||
if any(
|
||||
_is_not_empty(explicit_configurable.get(k))
|
||||
for k in _CHECKPOINT_COORDINATE_KEYS
|
||||
):
|
||||
empty[CONF] = {}
|
||||
break
|
||||
for config in configs:
|
||||
if config is None:
|
||||
continue
|
||||
|
||||
@@ -95,6 +95,15 @@ NULL_TASK_ID = sys.intern("00000000-0000-0000-0000-000000000000")
|
||||
OVERWRITE = sys.intern("__overwrite__")
|
||||
# dict key for the overwrite value, used as `{'__overwrite__': value}`
|
||||
|
||||
# Checkpoint coordinate keys: when any of these appear in an explicit
|
||||
# configurable, the caller is addressing its own checkpoint lineage.
|
||||
_CHECKPOINT_COORDINATE_KEYS = (
|
||||
CONFIG_KEY_THREAD_ID,
|
||||
CONFIG_KEY_CHECKPOINT_NS,
|
||||
CONFIG_KEY_CHECKPOINT_ID,
|
||||
CONFIG_KEY_CHECKPOINT_MAP,
|
||||
)
|
||||
|
||||
# redefined to avoid circular import with langgraph.constants
|
||||
_TAG_HIDDEN = sys.intern("langsmith:hidden")
|
||||
|
||||
|
||||
@@ -132,13 +132,23 @@ class GraphRunStream:
|
||||
def abort(self) -> None:
|
||||
"""Stop the run early.
|
||||
|
||||
Closes the mux and marks the stream exhausted. The graph
|
||||
iterator is dropped; any in-flight nodes see the closure on
|
||||
their next yield point. Idempotent.
|
||||
Closes the underlying graph iterator (propagating `GeneratorExit`
|
||||
so in-flight nodes and subgraphs are cancelled), closes the mux,
|
||||
and marks the stream exhausted. Idempotent.
|
||||
"""
|
||||
if self._exhausted:
|
||||
return
|
||||
self._exhausted = True
|
||||
graph_iter = self._graph_iter
|
||||
self._graph_iter = None
|
||||
if (
|
||||
graph_iter is not None
|
||||
and (close := getattr(graph_iter, "close", None)) is not None
|
||||
):
|
||||
try:
|
||||
close()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
self._mux.close()
|
||||
except Exception:
|
||||
@@ -348,6 +358,8 @@ class AsyncGraphRunStream:
|
||||
self._scope_list: list[str] = list(mux.scope)
|
||||
self._pump_cond = asyncio.Condition()
|
||||
self._pumping = False
|
||||
self._anext_task: asyncio.Future[Any] | None = None
|
||||
self._aborting = False
|
||||
for key in mux.native_keys:
|
||||
setattr(self, key, mux.extensions[key])
|
||||
if wire_pump:
|
||||
@@ -407,7 +419,25 @@ class AsyncGraphRunStream:
|
||||
|
||||
try:
|
||||
try:
|
||||
part = await self._graph_aiter.__anext__()
|
||||
# Run the pull as a child task so `abort()` can cancel it
|
||||
# mid-flight. Cancelling propagates `CancelledError` into the
|
||||
# graph generator frame -> Pregel loop -> nested subgraph
|
||||
# nodes, which a bare `aclose()` cannot do while the generator
|
||||
# is running ("asynchronous generator is already running").
|
||||
self._anext_task = asyncio.ensure_future(self._graph_aiter.__anext__())
|
||||
try:
|
||||
part = await self._anext_task
|
||||
except asyncio.CancelledError:
|
||||
if self._aborting:
|
||||
# Abort-initiated cancel: stop gracefully.
|
||||
self._exhausted = True
|
||||
return False
|
||||
# Genuine external cancel of this task: also stop the
|
||||
# in-flight pull, then propagate.
|
||||
self._anext_task.cancel()
|
||||
raise
|
||||
finally:
|
||||
self._anext_task = None
|
||||
event = convert_to_protocol_event(part)
|
||||
self._observe_event(event)
|
||||
await self._mux.apush(event)
|
||||
@@ -428,15 +458,40 @@ class AsyncGraphRunStream:
|
||||
async def abort(self) -> None:
|
||||
"""Stop the run early.
|
||||
|
||||
Marks the stream exhausted, wakes any pump-waiters, and closes
|
||||
the mux. Any `apush` blocked on backpressure wakes and returns
|
||||
without appending. Idempotent.
|
||||
Marks the stream exhausted and wakes any pump-waiters. Cancels an
|
||||
in-flight pull if one is running, then closes the underlying graph
|
||||
iterator, so running nodes and nested subgraphs are cancelled
|
||||
whether or not a pump is mid-pull. Closes the mux; any `apush`
|
||||
blocked on backpressure wakes and returns without appending.
|
||||
Idempotent.
|
||||
"""
|
||||
async with self._pump_cond:
|
||||
if self._exhausted:
|
||||
return
|
||||
self._exhausted = True
|
||||
self._aborting = True
|
||||
graph_aiter = self._graph_aiter
|
||||
self._graph_aiter = None
|
||||
anext_task = self._anext_task
|
||||
self._pump_cond.notify_all()
|
||||
# If a pump is mid-pull, cancel it so the cancellation propagates
|
||||
# into running nodes and nested subgraphs. Once it settles the
|
||||
# generator is no longer running, so the `aclose()` below is a safe
|
||||
# final cleanup (and handles the no-in-flight-pull case directly).
|
||||
if anext_task is not None and not anext_task.done():
|
||||
anext_task.cancel()
|
||||
try:
|
||||
await anext_task
|
||||
except (asyncio.CancelledError, Exception):
|
||||
pass
|
||||
if (
|
||||
graph_aiter is not None
|
||||
and (aclose := getattr(graph_aiter, "aclose", None)) is not None
|
||||
):
|
||||
try:
|
||||
await aclose()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
await self._mux.aclose()
|
||||
except Exception:
|
||||
|
||||
@@ -8,6 +8,7 @@ from langgraph.checkpoint.memory import InMemorySaver
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
|
||||
from langgraph._internal._constants import OVERWRITE
|
||||
from langgraph._internal._typing import MISSING
|
||||
from langgraph.channels.binop import BinaryOperatorAggregate
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
@@ -397,6 +398,49 @@ def test_delta_channel_inmemory_saver_assembles_writes() -> None:
|
||||
assert len(state.values["messages"]) == 4 # 2 human + 2 AI
|
||||
|
||||
|
||||
def test_delta_channel_api_json_overwrite_sentinel_snapshots_and_replays() -> None:
|
||||
def reducer(state: list[str], writes: Sequence[list[str]]) -> list[str]:
|
||||
result = list(state)
|
||||
for write in writes:
|
||||
result.extend(write)
|
||||
return result
|
||||
|
||||
class State(TypedDict):
|
||||
items: Annotated[
|
||||
list[str], DeltaChannel(reducer, list, snapshot_frequency=1000)
|
||||
]
|
||||
|
||||
calls = 0
|
||||
|
||||
def node(state: State) -> dict:
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
if calls == 1:
|
||||
return {"items": {OVERWRITE: ["reset"]}}
|
||||
return {"items": ["after"]}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("node", node)
|
||||
builder.add_edge(START, "node")
|
||||
|
||||
saver = InMemorySaver()
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
config = {"configurable": {"thread_id": "overwrite-json-sentinel"}}
|
||||
|
||||
updates = list(graph.stream({"items": ["before"]}, config, stream_mode=["updates"]))
|
||||
assert updates == [("updates", {"node": {"items": {OVERWRITE: ["reset"]}}})]
|
||||
assert graph.get_state(config).values == {"items": ["reset"]}
|
||||
|
||||
first_saved = saver.get_tuple(config)
|
||||
assert first_saved is not None
|
||||
snapshot = first_saved.checkpoint["channel_values"].get("items")
|
||||
assert isinstance(snapshot, _DeltaSnapshot)
|
||||
assert snapshot.value == ["reset"]
|
||||
|
||||
assert graph.invoke({"items": []}, config) == {"items": ["reset", "after"]}
|
||||
assert graph.get_state(config).values == {"items": ["reset", "after"]}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DeltaChannel — dict reducer
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -607,6 +607,164 @@ class TestStreamV2Async:
|
||||
_ = await anext(aiter(run.values))
|
||||
assert run._exhausted is True
|
||||
|
||||
async def test_abort_cancels_running_subgraph(self) -> None:
|
||||
class CountState(TypedDict):
|
||||
count: int
|
||||
|
||||
runs: list[int] = []
|
||||
|
||||
async def sub_node(state: CountState) -> dict:
|
||||
runs.append(state["count"] + 1)
|
||||
await asyncio.sleep(0.05)
|
||||
return {"count": state["count"] + 1}
|
||||
|
||||
sub_graph = (
|
||||
StateGraph(CountState)
|
||||
.add_node("sub_node", sub_node)
|
||||
.set_entry_point("sub_node")
|
||||
.add_conditional_edges(
|
||||
"sub_node",
|
||||
lambda s: END if s["count"] >= 10 else "sub_node",
|
||||
)
|
||||
.compile()
|
||||
)
|
||||
|
||||
async def main_node(state: CountState) -> None:
|
||||
await sub_graph.ainvoke({"count": 0})
|
||||
|
||||
main_graph = (
|
||||
StateGraph(CountState)
|
||||
.add_node("main_node", main_node)
|
||||
.set_entry_point("main_node")
|
||||
.compile()
|
||||
)
|
||||
|
||||
run = await main_graph.astream_events({"count": 0}, version="v3")
|
||||
async for e in run:
|
||||
if (
|
||||
e["method"] == "values"
|
||||
and e["params"]["namespace"]
|
||||
and e["params"]["data"]["count"] >= 2
|
||||
):
|
||||
break
|
||||
await run.abort()
|
||||
runs_at_abort = len(runs)
|
||||
# Give the (now-cancelled) subgraph a chance to keep looping.
|
||||
await asyncio.sleep(0.3)
|
||||
assert len(runs) == runs_at_abort
|
||||
assert len(runs) < 10
|
||||
|
||||
async def test_abort_cancels_deeply_nested_subgraph(self) -> None:
|
||||
class CountState(TypedDict):
|
||||
count: int
|
||||
|
||||
runs: list[int] = []
|
||||
|
||||
async def deep_node(state: CountState) -> dict:
|
||||
runs.append(state["count"] + 1)
|
||||
await asyncio.sleep(0.05)
|
||||
return {"count": state["count"] + 1}
|
||||
|
||||
# Deepest graph loops until count >= 10.
|
||||
graph = (
|
||||
StateGraph(CountState)
|
||||
.add_node("deep_node", deep_node)
|
||||
.set_entry_point("deep_node")
|
||||
.add_conditional_edges(
|
||||
"deep_node",
|
||||
lambda s: END if s["count"] >= 10 else "deep_node",
|
||||
)
|
||||
.compile()
|
||||
)
|
||||
|
||||
# Wrap it three times: graph -> subgraph -> subgraph -> subgraph.
|
||||
for _ in range(3):
|
||||
|
||||
async def caller(state: CountState, _child: Any = graph) -> dict:
|
||||
return await _child.ainvoke({"count": 0})
|
||||
|
||||
graph = (
|
||||
StateGraph(CountState)
|
||||
.add_node("caller", caller)
|
||||
.set_entry_point("caller")
|
||||
.compile()
|
||||
)
|
||||
|
||||
run = await graph.astream_events({"count": 0}, version="v3")
|
||||
async for e in run:
|
||||
if (
|
||||
e["method"] == "values"
|
||||
and e["params"]["namespace"]
|
||||
and e["params"]["data"]["count"] >= 2
|
||||
):
|
||||
break
|
||||
await run.abort()
|
||||
runs_at_abort = len(runs)
|
||||
# Give the (now-cancelled) nested subgraph a chance to keep looping.
|
||||
await asyncio.sleep(0.3)
|
||||
assert len(runs) == runs_at_abort
|
||||
assert len(runs) < 10
|
||||
|
||||
async def test_abort_cancels_subgraph_during_inflight_pump(self) -> None:
|
||||
class CountState(TypedDict):
|
||||
count: int
|
||||
|
||||
started = asyncio.Event()
|
||||
cancelled = asyncio.Event()
|
||||
|
||||
async def sub_node(state: CountState) -> dict:
|
||||
started.set()
|
||||
try:
|
||||
# Long-running node: still in flight when abort fires.
|
||||
await asyncio.sleep(5)
|
||||
except asyncio.CancelledError:
|
||||
cancelled.set()
|
||||
raise
|
||||
return {"count": state["count"] + 1}
|
||||
|
||||
sub_graph = (
|
||||
StateGraph(CountState)
|
||||
.add_node("sub_node", sub_node)
|
||||
.set_entry_point("sub_node")
|
||||
.compile()
|
||||
)
|
||||
|
||||
async def main_node(state: CountState) -> None:
|
||||
await sub_graph.ainvoke({"count": 0})
|
||||
|
||||
main_graph = (
|
||||
StateGraph(CountState)
|
||||
.add_node("main_node", main_node)
|
||||
.set_entry_point("main_node")
|
||||
.compile()
|
||||
)
|
||||
|
||||
run = await main_graph.astream_events({"count": 0}, version="v3")
|
||||
|
||||
# A consumer task drives the pump. Once the subgraph node is
|
||||
# running, no further event is produced, so the consumer parks
|
||||
# inside _apump_next awaiting graph_aiter.__anext__() — the
|
||||
# generator is "running" and a plain aclose() would raise.
|
||||
async def consume() -> None:
|
||||
async for _e in run:
|
||||
pass
|
||||
|
||||
consumer = asyncio.create_task(consume())
|
||||
try:
|
||||
await asyncio.wait_for(started.wait(), timeout=2.0)
|
||||
# Let the consumer drain and park in __anext__.
|
||||
await asyncio.sleep(0.05)
|
||||
# Abort from a different task while the consumer is in __anext__.
|
||||
await run.abort()
|
||||
# The in-flight subgraph node must observe cancellation.
|
||||
await asyncio.wait_for(cancelled.wait(), timeout=2.0)
|
||||
finally:
|
||||
consumer.cancel()
|
||||
try:
|
||||
await consumer
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
async def test_extensions_has_native_keys(self) -> None:
|
||||
run = await _build_simple_graph().astream_events(
|
||||
{"value": "x", "items": []}, version="v3"
|
||||
|
||||
@@ -639,3 +639,49 @@ def test_stateful_namespace_isolation(
|
||||
"broccoli round 2",
|
||||
"Veggie: broccoli round 2",
|
||||
]
|
||||
|
||||
|
||||
def test_child_with_own_thread_id_keeps_namespace(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""A child graph invoked from inside a parent node with its own thread_id
|
||||
must store and read its checkpoint under its own namespace, not inherit the
|
||||
parent task's checkpoint_ns.
|
||||
"""
|
||||
|
||||
class ChildState(TypedDict):
|
||||
count: int
|
||||
|
||||
def child_node(state: ChildState) -> dict:
|
||||
return {"count": (state.get("count") or 0) + 1}
|
||||
|
||||
child = (
|
||||
StateGraph(ChildState)
|
||||
.add_node("n", child_node)
|
||||
.add_edge(START, "n")
|
||||
.compile(checkpointer=sync_checkpointer)
|
||||
)
|
||||
|
||||
child_thread = str(uuid4())
|
||||
child_config = {"configurable": {"thread_id": child_thread}}
|
||||
|
||||
def parent_node(state: ParentState) -> dict:
|
||||
child.invoke({}, config=child_config)
|
||||
return {"result": "ok"}
|
||||
|
||||
parent = (
|
||||
StateGraph(ParentState)
|
||||
.add_node("p", parent_node)
|
||||
.add_edge(START, "p")
|
||||
.compile(checkpointer=sync_checkpointer)
|
||||
)
|
||||
parent_config = {"configurable": {"thread_id": str(uuid4())}}
|
||||
|
||||
parent.invoke({"result": ""}, config=parent_config)
|
||||
state1 = child.get_state(child_config)
|
||||
assert state1.values.get("count") == 1
|
||||
assert state1.config["configurable"]["checkpoint_ns"] == ""
|
||||
|
||||
parent.invoke({"result": ""}, config=parent_config)
|
||||
state2 = child.get_state(child_config)
|
||||
assert state2.values.get("count") == 2
|
||||
|
||||
@@ -660,3 +660,50 @@ async def test_stateful_namespace_isolation_async(
|
||||
"broccoli round 2",
|
||||
"Veggie: broccoli round 2",
|
||||
]
|
||||
|
||||
|
||||
@NEEDS_CONTEXTVARS
|
||||
async def test_child_with_own_thread_id_keeps_namespace_async(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""A child graph invoked from inside a parent node with its own thread_id
|
||||
must store and read its checkpoint under its own namespace, not inherit the
|
||||
parent task's checkpoint_ns.
|
||||
"""
|
||||
|
||||
class ChildState(TypedDict):
|
||||
count: int
|
||||
|
||||
def child_node(state: ChildState) -> dict:
|
||||
return {"count": (state.get("count") or 0) + 1}
|
||||
|
||||
child = (
|
||||
StateGraph(ChildState)
|
||||
.add_node("n", child_node)
|
||||
.add_edge(START, "n")
|
||||
.compile(checkpointer=async_checkpointer)
|
||||
)
|
||||
|
||||
child_thread = str(uuid4())
|
||||
child_config = {"configurable": {"thread_id": child_thread}}
|
||||
|
||||
async def parent_node(state: ParentState) -> dict:
|
||||
await child.ainvoke({}, config=child_config)
|
||||
return {"result": "ok"}
|
||||
|
||||
parent = (
|
||||
StateGraph(ParentState)
|
||||
.add_node("p", parent_node)
|
||||
.add_edge(START, "p")
|
||||
.compile(checkpointer=async_checkpointer)
|
||||
)
|
||||
parent_config = {"configurable": {"thread_id": str(uuid4())}}
|
||||
|
||||
await parent.ainvoke({"result": ""}, config=parent_config)
|
||||
state1 = await child.aget_state(child_config)
|
||||
assert state1.values.get("count") == 1
|
||||
assert state1.config["configurable"]["checkpoint_ns"] == ""
|
||||
|
||||
await parent.ainvoke({"result": ""}, config=parent_config)
|
||||
state2 = await child.aget_state(child_config)
|
||||
assert state2.values.get("count") == 2
|
||||
|
||||
@@ -506,6 +506,95 @@ def test_ensure_config_configurable_later_wins_per_key() -> None:
|
||||
assert merged["configurable"]["only_b"] == "B"
|
||||
|
||||
|
||||
def test_ensure_config_explicit_configurable_replaces_ambient() -> None:
|
||||
# An explicit checkpoint coordinate (here a new thread_id) starts a fresh
|
||||
# lineage and drops the ambient run context (e.g. a parent task's
|
||||
# checkpoint_ns), so a child graph does not inherit it.
|
||||
from langchain_core.runnables.config import var_child_runnable_config
|
||||
|
||||
token = var_child_runnable_config.set(
|
||||
{"configurable": {"checkpoint_ns": "p:parent-task", "checkpoint_id": "cid"}}
|
||||
)
|
||||
try:
|
||||
merged = ensure_config({"configurable": {"thread_id": "child"}})
|
||||
finally:
|
||||
var_child_runnable_config.reset(token)
|
||||
assert merged["configurable"]["thread_id"] == "child"
|
||||
assert "checkpoint_ns" not in merged["configurable"]
|
||||
assert "checkpoint_id" not in merged["configurable"]
|
||||
|
||||
|
||||
def test_ensure_config_ambient_inherited_when_no_explicit_configurable() -> None:
|
||||
# With no explicit configurable, the ambient run context is inherited
|
||||
# unchanged (stateless subgraph / interrupt-resume pattern).
|
||||
from langchain_core.runnables.config import var_child_runnable_config
|
||||
|
||||
token = var_child_runnable_config.set(
|
||||
{"configurable": {"checkpoint_ns": "p:parent-task"}}
|
||||
)
|
||||
try:
|
||||
merged = ensure_config({"tags": ["t"]})
|
||||
finally:
|
||||
var_child_runnable_config.reset(token)
|
||||
assert merged["configurable"]["checkpoint_ns"] == "p:parent-task"
|
||||
|
||||
|
||||
def test_ensure_config_explicit_configurables_still_merge_over_ambient() -> None:
|
||||
# A new thread_id drops the ambient, but explicit configs still shallow-merge
|
||||
# among themselves, so a with_config(...) value (ls_agent_type) survives
|
||||
# alongside an invoke-time thread_id.
|
||||
from langchain_core.runnables.config import var_child_runnable_config
|
||||
|
||||
token = var_child_runnable_config.set(
|
||||
{"configurable": {"checkpoint_ns": "p:parent-task"}}
|
||||
)
|
||||
try:
|
||||
merged = ensure_config(
|
||||
{"configurable": {"ls_agent_type": "root"}},
|
||||
{"configurable": {"thread_id": "child"}},
|
||||
)
|
||||
finally:
|
||||
var_child_runnable_config.reset(token)
|
||||
assert merged["configurable"]["ls_agent_type"] == "root"
|
||||
assert merged["configurable"]["thread_id"] == "child"
|
||||
assert "checkpoint_ns" not in merged["configurable"]
|
||||
|
||||
|
||||
def test_ensure_config_non_coordinate_config_keeps_ambient_checkpoint_ns() -> None:
|
||||
# A nested subagent is invoked with a non-coordinate configurable key
|
||||
# (ls_agent_type) and no thread_id; it must keep the inherited checkpoint_ns
|
||||
# so it stays a discoverable child of the parent run (deepagents `task` tool).
|
||||
from langchain_core.runnables.config import var_child_runnable_config
|
||||
|
||||
token = var_child_runnable_config.set(
|
||||
{"configurable": {"thread_id": "parent", "checkpoint_ns": "p:parent-task"}}
|
||||
)
|
||||
try:
|
||||
merged = ensure_config({"configurable": {"ls_agent_type": "subagent"}})
|
||||
finally:
|
||||
var_child_runnable_config.reset(token)
|
||||
assert merged["configurable"]["ls_agent_type"] == "subagent"
|
||||
assert merged["configurable"]["checkpoint_ns"] == "p:parent-task"
|
||||
assert merged["configurable"]["thread_id"] == "parent"
|
||||
|
||||
|
||||
def test_ensure_config_same_thread_id_still_clears_ambient() -> None:
|
||||
# A child that reuses the parent's thread_id is still addressing its own root
|
||||
# namespace on that thread, so the parent task's checkpoint_ns must not leak
|
||||
# in; otherwise the child writes state that get_state cannot read back.
|
||||
from langchain_core.runnables.config import var_child_runnable_config
|
||||
|
||||
token = var_child_runnable_config.set(
|
||||
{"configurable": {"thread_id": "shared", "checkpoint_ns": "p:parent-task"}}
|
||||
)
|
||||
try:
|
||||
merged = ensure_config({"configurable": {"thread_id": "shared"}})
|
||||
finally:
|
||||
var_child_runnable_config.reset(token)
|
||||
assert merged["configurable"]["thread_id"] == "shared"
|
||||
assert "checkpoint_ns" not in merged["configurable"]
|
||||
|
||||
|
||||
def test_ensure_config_merges_metadata_across_configs() -> None:
|
||||
a = {"metadata": {"user_id": "U1"}}
|
||||
b = {"metadata": {"correlation_id": "C1"}}
|
||||
|
||||
Reference in New Issue
Block a user