mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-03 23:15:10 +02:00
Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
af62aff8ac | ||
|
|
4210feccd9 | ||
|
|
941c170c58 | ||
|
|
79befe67ba | ||
|
|
9af25217c3 |
@@ -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")
|
||||
|
||||
|
||||
@@ -177,8 +177,7 @@ class DeltaChannel(Generic[Value], BaseChannel[Any, Any, Any]):
|
||||
if overwrite_value is not None
|
||||
else self.typ()
|
||||
)
|
||||
remaining = [v for i, v in enumerate(values) if i != overwrite_idx]
|
||||
self.value = self.reducer(base, remaining) if remaining else base
|
||||
self.value = base
|
||||
return True
|
||||
base = self.typ() if self.value is MISSING else self.value
|
||||
self.value = self.reducer(base, list(values))
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -186,6 +186,18 @@ def test_delta_channel_overwrite() -> None:
|
||||
assert ch.get()[0].content == "new"
|
||||
|
||||
|
||||
def test_delta_channel_overwrite_bypasses_same_step_reducer_writes() -> None:
|
||||
def list_reducer(state: list, writes: list) -> list:
|
||||
out = list(state)
|
||||
for w in writes:
|
||||
out.extend(w)
|
||||
return out
|
||||
|
||||
ch = DeltaChannel(list_reducer, list).from_checkpoint(MISSING)
|
||||
ch.update([[1], Overwrite([50]), [2]])
|
||||
assert ch.get() == [50]
|
||||
|
||||
|
||||
def test_delta_channel_remove_message_and_replay() -> None:
|
||||
"""RemoveMessage must round-trip correctly when writes are replayed."""
|
||||
spec = DeltaChannel(_messages_delta_reducer, list)
|
||||
|
||||
@@ -9281,6 +9281,13 @@ def test_send_with_untracked_value_overlapping_keys(
|
||||
assert state.values.get("dictionary") == {"session_resource": "legal_value"}
|
||||
|
||||
|
||||
def _delta_list_reducer(state: list, writes: Sequence[list]) -> list:
|
||||
out = list(state)
|
||||
for write in writes:
|
||||
out.extend(write)
|
||||
return out
|
||||
|
||||
|
||||
@pytest.mark.parametrize("as_json", [False, True])
|
||||
def test_overwrite_sequential(
|
||||
sync_checkpointer: BaseCheckpointSaver, as_json: bool
|
||||
@@ -9388,6 +9395,105 @@ def test_overwrite_parallel_error(
|
||||
graph.invoke({"messages": ["START"]}, config)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("as_json", [False, True])
|
||||
def test_delta_channel_overwrite_sequential(
|
||||
sync_checkpointer: BaseCheckpointSaver, as_json: bool
|
||||
) -> None:
|
||||
class State(TypedDict):
|
||||
messages: Annotated[list, DeltaChannel(_delta_list_reducer)]
|
||||
|
||||
def node_a(state: State):
|
||||
return {"messages": ["a"]}
|
||||
|
||||
def node_b(state: State):
|
||||
overwrite = {"__overwrite__": ["b"]} if as_json else Overwrite(["b"])
|
||||
return {"messages": overwrite}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("node_a", node_a)
|
||||
builder.add_node("node_b", node_b)
|
||||
builder.add_edge(START, "node_a")
|
||||
builder.add_edge("node_a", "node_b")
|
||||
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "delta-overwrite-sequential"}}
|
||||
result = graph.invoke({"messages": ["START"]}, config)
|
||||
assert result == {"messages": ["b"]}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("as_json", [False, True])
|
||||
def test_delta_channel_overwrite_parallel(
|
||||
sync_checkpointer: BaseCheckpointSaver, as_json: bool
|
||||
) -> None:
|
||||
class State(TypedDict):
|
||||
messages: Annotated[list, DeltaChannel(_delta_list_reducer)]
|
||||
|
||||
def node_a(state: State):
|
||||
return {"messages": ["a"]}
|
||||
|
||||
def node_b(state: State):
|
||||
overwrite = {"__overwrite__": ["b"]} if as_json else Overwrite(["b"])
|
||||
return {"messages": overwrite}
|
||||
|
||||
def node_c(state: State):
|
||||
return {"messages": ["c"]}
|
||||
|
||||
def node_d(state: State):
|
||||
return {"messages": ["d"]}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("node_a", node_a)
|
||||
builder.add_node("node_b", node_b)
|
||||
builder.add_node("node_c", node_c)
|
||||
builder.add_node("node_d", node_d)
|
||||
builder.add_edge(START, "node_a")
|
||||
builder.add_edge("node_a", "node_b")
|
||||
builder.add_edge("node_a", "node_c")
|
||||
builder.add_edge("node_b", "node_d")
|
||||
builder.add_edge("node_c", "node_d")
|
||||
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "delta-overwrite-parallel"}}
|
||||
result = graph.invoke({"messages": ["START"]}, config)
|
||||
assert result == {"messages": ["b", "d"]}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("as_json", [False, True])
|
||||
def test_delta_channel_overwrite_parallel_error(
|
||||
sync_checkpointer: BaseCheckpointSaver, as_json: bool
|
||||
) -> None:
|
||||
class State(TypedDict):
|
||||
messages: Annotated[list, DeltaChannel(_delta_list_reducer)]
|
||||
|
||||
def node_a(state: State):
|
||||
return {"messages": ["a"]}
|
||||
|
||||
def node_b(state: State):
|
||||
overwrite = {"__overwrite__": ["b"]} if as_json else Overwrite(["b"])
|
||||
return {"messages": overwrite}
|
||||
|
||||
def node_c(state: State):
|
||||
overwrite = {"__overwrite__": ["c"]} if as_json else Overwrite(["c"])
|
||||
return {"messages": overwrite}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("node_a", node_a)
|
||||
builder.add_node("node_b", node_b)
|
||||
builder.add_node("node_c", node_c)
|
||||
builder.add_edge(START, "node_a")
|
||||
builder.add_edge("node_a", "node_b")
|
||||
builder.add_edge("node_a", "node_c")
|
||||
builder.add_edge("node_b", END)
|
||||
builder.add_edge("node_c", END)
|
||||
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "delta-overwrite-parallel-error"}}
|
||||
with pytest.raises(
|
||||
InvalidUpdateError, match="Can receive only one Overwrite value per super-step."
|
||||
):
|
||||
graph.invoke({"messages": ["START"]}, config)
|
||||
|
||||
|
||||
def test_fork_does_not_apply_pending_writes(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
|
||||
@@ -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