feat(langgraph): add Overwrite to bypass reducer (#6286)

See https://github.com/langchain-ai/langgraph/pull/6277

Adds langgraph.types.Overwrite, a deterministic way to bypass a reducer.
When encountering a value wrapped with Overwrite,
BinaryOperatorAggregate overwrites the channel value.

<img width="227" height="329" alt="image"
src="https://github.com/user-attachments/assets/f2136117-9aa3-4246-863d-d5df0e7d1df1"
/>

If either node_b or node_c overwrite (but not both), then at END the
channel is equal to the value node_b or node_c wrote. Order of execution
doesn't matter because once an Overwrite value is encountered, regular
values are ignored (self.operator is not called for the rest of the
update)

If multiple nodes overwrite in the same superstep then
InvalidUpdateError is thrown

Usage
```python
from langgraph.types import Overwrite

def node_b(state:State):
    return {"messages": Overwrite(["b"])}
```
or
``` python
def node_b(state:State):
    return {"messages": {"__overwrite__": ["b"]}}
```
This commit is contained in:
Caspar Broekhuizen
2025-10-27 11:31:35 -07:00
committed by GitHub
parent 10abf2deb1
commit 504e91ad5a
4 changed files with 182 additions and 3 deletions
@@ -77,6 +77,8 @@ CONF = cast(Literal["configurable"], sys.intern("configurable"))
# key for the configurable dict in RunnableConfig
NULL_TASK_ID = sys.intern("00000000-0000-0000-0000-000000000000")
# the task_id to use for writes that are not associated with a task
OVERWRITE = sys.intern("__overwrite__")
# dict key for the overwrite value, used as `{'__overwrite__': value}`
# redefined to avoid circular import with langgraph.constants
_TAG_HIDDEN = sys.intern("langsmith:hidden")
+32 -3
View File
@@ -1,12 +1,19 @@
import collections.abc
from collections.abc import Callable, Sequence
from typing import Generic
from typing import Any, Generic
from typing_extensions import NotRequired, Required, Self
from langgraph._internal._constants import OVERWRITE
from langgraph._internal._typing import MISSING
from langgraph.channels.base import BaseChannel, Value
from langgraph.errors import EmptyChannelError
from langgraph.errors import (
EmptyChannelError,
ErrorCode,
InvalidUpdateError,
create_error_message,
)
from langgraph.types import Overwrite
__all__ = ("BinaryOperatorAggregate",)
@@ -22,6 +29,15 @@ def _strip_extras(t): # type: ignore[no-untyped-def]
return t
def _get_overwrite(value: Any) -> tuple[bool, Any]:
"""Inspects the given value and returns (is_overwrite, overwrite_value)."""
if isinstance(value, Overwrite):
return True, value.value
if isinstance(value, dict) and set(value.keys()) == {OVERWRITE}:
return True, value[OVERWRITE]
return False, None
class BinaryOperatorAggregate(Generic[Value], BaseChannel[Value, Value, Value]):
"""Stores the result of applying a binary operator to the current value and each new value.
@@ -89,8 +105,21 @@ class BinaryOperatorAggregate(Generic[Value], BaseChannel[Value, Value, Value]):
if self.value is MISSING:
self.value = values[0]
values = values[1:]
seen_overwrite: bool = False
for value in values:
self.value = self.operator(self.value, value)
is_overwrite, overwrite_value = _get_overwrite(value)
if is_overwrite:
if seen_overwrite:
msg = create_error_message(
message="Can receive only one Overwrite value per super-step.",
error_code=ErrorCode.INVALID_CONCURRENT_GRAPH_UPDATE,
)
raise InvalidUpdateError(msg)
self.value = overwrite_value
seen_overwrite = True
continue
if not seen_overwrite:
self.value = self.operator(self.value, value)
return True
def get(self) -> Value:
+40
View File
@@ -516,3 +516,43 @@ def interrupt(value: Any) -> Any:
),
)
)
@dataclass(slots=True)
class Overwrite:
"""Bypass a reducer and write the wrapped value directly to a BinaryOperatorAggregate channel.
Receiving multiple Overwrite values for the same channel in a single super-step will raise an InvalidUpdateError.
Example:
>>> from typing import Annotated
>>> import operator
>>> from langgraph.graph import StateGraph
>>> from langgraph.types import Overwrite
>>>
>>> class State(TypedDict):
... messages: Annotated[list, operator.add]
>>>
>>> def node_a(state: TypedDict):
... # Normal update: uses the reducer (operator.add)
... return {"messages": ["a"]}
>>>
>>> def node_b(state: State):
... # Overwrite: bypasses the reducer and replaces the entire value
... return {"messages": Overwrite(value=["b"])}
>>>
>>> builder = StateGraph(State)
>>> builder.add_node("node_a", node_a)
>>> builder.add_node("node_b", node_b)
>>> builder.set_entry_point("node_a")
>>> builder.add_edge("node_a", "node_b")
>>> graph = builder.compile()
>>>
>>> # Without Overwrite in node_b, messages would be ["START", "a", "b"]
>>> # With Overwrite, messages is just ["b"]
>>> result = graph.invoke({"messages": ["START"]})
>>> assert result == {"messages": ["b"]}
"""
value: Any
"""The value to write directly to the channel, bypassing any reducer."""
+108
View File
@@ -60,6 +60,7 @@ from langgraph.types import (
Command,
Durability,
Interrupt,
Overwrite,
PregelTask,
RetryPolicy,
Send,
@@ -8597,3 +8598,110 @@ def test_multiple_writes_same_channel_from_same_node(
"values": {"foo": ""},
},
]
@pytest.mark.parametrize("as_json", [False, True])
def test_overwrite_sequential(
sync_checkpointer: BaseCheckpointSaver, as_json: bool
) -> None:
"""Test a sequential chain of nodes where the last node uses Overwrite to bypass a reducer and write a value directly to the channel."""
class State(TypedDict):
messages: Annotated[list, operator.add]
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": "1"}}
result = graph.invoke({"messages": ["START"]}, config)
# a is overwritten by b
assert result == {"messages": ["b"]}
@pytest.mark.parametrize("as_json", [False, True])
def test_overwrite_parallel(
sync_checkpointer: BaseCheckpointSaver, as_json: bool
) -> None:
"""Test parallel nodes where max one node uses Overwrite to bypass a reducer and write a value directly to the channel."""
class State(TypedDict):
messages: Annotated[list, operator.add]
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": "1"}}
result = graph.invoke({"messages": ["START"]}, config)
# a, c are overwritten by b, then d is written
assert result == {"messages": ["b", "d"]}
@pytest.mark.parametrize("as_json", [False, True])
def test_overwrite_parallel_error(
sync_checkpointer: BaseCheckpointSaver, as_json: bool
) -> None:
"""Test parallel nodes where more than one node uses Overwrite to bypass a reducer and write a value directly to the channel. In this case, InvalidUpdateError should be raised."""
class State(TypedDict):
messages: Annotated[list, operator.add]
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": "1"}}
with pytest.raises(
InvalidUpdateError, match="Can receive only one Overwrite value per super-step."
):
graph.invoke({"messages": ["START"]}, config)