Compare commits

...
3 changed files with 84 additions and 6 deletions
+2 -2
View File
@@ -154,7 +154,7 @@ class DeltaChannel(Generic[Value], BaseChannel[Any, Any, Any]):
base = _copy.copy(ow_value) if ow_value is not None else self.typ()
start = i + 1
remaining = values[start:]
self.value = self.reducer(base, remaining) if remaining else base
self.value = self.reducer(base, remaining)
def update(self, values: Sequence[Any]) -> bool:
if not values:
@@ -178,7 +178,7 @@ class DeltaChannel(Generic[Value], BaseChannel[Any, Any, Any]):
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 = self.reducer(base, remaining)
return True
base = self.typ() if self.value is MISSING else self.value
self.value = self.reducer(base, list(values))
+15 -4
View File
@@ -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
+67
View File
@@ -9310,6 +9310,73 @@ 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"]}
def test_delta_channel_overwrite_normalizes_json_messages(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
class State(TypedDict):
messages: Annotated[
list[AnyMessage],
DeltaChannel(_messages_delta_reducer),
]
def node_a(state: State):
return {"messages": [HumanMessage(content="original", id="h0")]}
def node_b(state: State):
update = Overwrite([HumanMessage(content="replacement", id="h1")])
overwrite = json.loads(json.dumps(update, default=lambda obj: obj.model_dump()))
assert overwrite == {
"__overwrite__": [
{
"content": "replacement",
"additional_kwargs": {},
"response_metadata": {},
"type": "human",
"name": None,
"id": "h1",
}
]
}
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": []}, {"configurable": {"thread_id": "1"}})
assert result == {"messages": [HumanMessage(content="replacement", id="h1")]}
@pytest.mark.parametrize("as_json", [False, True])
def test_overwrite_parallel(
sync_checkpointer: BaseCheckpointSaver, as_json: bool