diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index 4471074cc..943144b7a 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -25,6 +25,7 @@ from xxhash import xxh3_128_hexdigest from langgraph._internal._cache import default_cache_key from langgraph._internal._constants import INTERRUPT as _INTERRUPT_KEY +from langgraph._internal._constants import OVERWRITE from langgraph._internal._fields import get_cached_annotated_keys, get_update_as_tuples from langgraph._internal._retry import default_retry_on from langgraph._internal._typing import MISSING, DeprecatedKwargs @@ -934,8 +935,7 @@ def interrupt(value: Any) -> Any: ) -@dataclass(slots=True) -class Overwrite: +class Overwrite(dict[str, Any]): """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 @@ -974,5 +974,16 @@ class Overwrite: ``` """ - value: Any - """The value to write directly to the channel, bypassing any reducer.""" + __slots__ = () + + def __init__(self, value: Any) -> None: + super().__init__({OVERWRITE: value}) + + @property + def value(self) -> Any: + """The value to write directly to the channel, bypassing any reducer.""" + return self[OVERWRITE] + + @value.setter + def value(self, value: Any) -> None: + self[OVERWRITE] = value diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 0aae1318e..d6d3af39e 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -9310,6 +9310,32 @@ def test_overwrite_sequential( assert result == {"messages": ["b"]} +def test_overwrite_survives_json_roundtrip( + sync_checkpointer: BaseCheckpointSaver, +) -> None: + class State(TypedDict): + messages: Annotated[list, operator.add] + + def node_a(state: State): + return {"messages": ["a"]} + + def node_b(state: State): + overwrite = json.loads(json.dumps(Overwrite(["b"]))) + assert overwrite == {"__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) + result = graph.invoke({"messages": ["START"]}, {"configurable": {"thread_id": "1"}}) + + assert result == {"messages": ["b"]} + + @pytest.mark.parametrize("as_json", [False, True]) def test_overwrite_parallel( sync_checkpointer: BaseCheckpointSaver, as_json: bool