mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-26 03:25:06 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e3390a7c48 | ||
|
|
5275579c36 |
@@ -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()
|
base = _copy.copy(ow_value) if ow_value is not None else self.typ()
|
||||||
start = i + 1
|
start = i + 1
|
||||||
remaining = values[start:]
|
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:
|
def update(self, values: Sequence[Any]) -> bool:
|
||||||
if not values:
|
if not values:
|
||||||
@@ -178,7 +178,7 @@ class DeltaChannel(Generic[Value], BaseChannel[Any, Any, Any]):
|
|||||||
else self.typ()
|
else self.typ()
|
||||||
)
|
)
|
||||||
remaining = [v for i, v in enumerate(values) if i != overwrite_idx]
|
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
|
return True
|
||||||
base = self.typ() if self.value is MISSING else self.value
|
base = self.typ() if self.value is MISSING else self.value
|
||||||
self.value = self.reducer(base, list(values))
|
self.value = self.reducer(base, list(values))
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ from xxhash import xxh3_128_hexdigest
|
|||||||
|
|
||||||
from langgraph._internal._cache import default_cache_key
|
from langgraph._internal._cache import default_cache_key
|
||||||
from langgraph._internal._constants import INTERRUPT as _INTERRUPT_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._fields import get_cached_annotated_keys, get_update_as_tuples
|
||||||
from langgraph._internal._retry import default_retry_on
|
from langgraph._internal._retry import default_retry_on
|
||||||
from langgraph._internal._typing import MISSING, DeprecatedKwargs
|
from langgraph._internal._typing import MISSING, DeprecatedKwargs
|
||||||
@@ -934,8 +935,7 @@ def interrupt(value: Any) -> Any:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
class Overwrite(dict[str, Any]):
|
||||||
class Overwrite:
|
|
||||||
"""Bypass a reducer and write the wrapped value directly to a `BinaryOperatorAggregate` channel.
|
"""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
|
Receiving multiple `Overwrite` values for the same channel in a single super-step
|
||||||
@@ -974,5 +974,16 @@ class Overwrite:
|
|||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
|
|
||||||
value: Any
|
__slots__ = ()
|
||||||
"""The value to write directly to the channel, bypassing any reducer."""
|
|
||||||
|
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
|
||||||
|
|||||||
@@ -9310,6 +9310,73 @@ def test_overwrite_sequential(
|
|||||||
assert result == {"messages": ["b"]}
|
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])
|
@pytest.mark.parametrize("as_json", [False, True])
|
||||||
def test_overwrite_parallel(
|
def test_overwrite_parallel(
|
||||||
sync_checkpointer: BaseCheckpointSaver, as_json: bool
|
sync_checkpointer: BaseCheckpointSaver, as_json: bool
|
||||||
|
|||||||
Reference in New Issue
Block a user