mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 04:25:08 +02:00
fix(langgraph): serialize as canonical dict
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user