This commit is contained in:
Quanzheng Long
2026-03-13 15:40:49 -07:00
parent 597b3402e6
commit 783b3d3435
4 changed files with 277 additions and 90 deletions
@@ -1,6 +1,7 @@
from __future__ import annotations
import asyncio
import copy
import inspect
from collections.abc import Callable, Coroutine, Sequence
from dataclasses import dataclass
@@ -252,7 +253,7 @@ class _GraphEngineRun:
def _execute_node_for_rust(
self, node_name: str, node_input: Any, state: Any
) -> dict[str, Any]:
before_state_markers = _state_shallow_markers(state)
before_state_snapshot = copy.deepcopy(state) if isinstance(state, dict) else None
if node_name not in self._nodes:
raise ValueError(f"Unknown node `{node_name}`")
@@ -270,11 +271,9 @@ class _GraphEngineRun:
if update is None and isinstance(state, dict):
# Preserve in-place state mutations for prototype nodes like wait_node.
if _has_shallow_state_change(before_state_markers, state):
update = state
elif isinstance(update, dict):
if not _has_shallow_update_change(before_state_markers, update):
update = None
update = state
if isinstance(update, dict):
update = _reduce_update_to_changed_fields(before_state_snapshot, update)
return {
"update": update,
@@ -390,44 +389,16 @@ def _resolve_target_name(target: Any) -> str:
raise ValueError(f"Unsupported node target type: {type(target)!r}")
def _state_shallow_markers(state: Any) -> dict[str, Any] | None:
if not isinstance(state, dict):
return None
return {k: _value_shallow_marker(v) for k, v in state.items()}
def _value_shallow_marker(value: Any) -> Any:
if value is None or isinstance(value, (bool, int, float, str, bytes)):
return ("primitive", value)
if isinstance(value, (list, tuple, set, dict)):
return ("container", type(value).__name__, id(value), len(value))
return ("object", type(value).__name__, id(value))
def _has_shallow_state_change(before: dict[str, Any] | None, state: Any) -> bool:
if before is None or not isinstance(state, dict):
return True
if len(before) != len(state):
return True
for key, old_marker in before.items():
if key not in state:
return True
if old_marker != _value_shallow_marker(state[key]):
return True
return False
def _has_shallow_update_change(
def _reduce_update_to_changed_fields(
before: dict[str, Any] | None, update: dict[str, Any]
) -> bool:
) -> dict[str, Any] | None:
if before is None:
return True
return update
changed: dict[str, Any] = {}
for key, new_value in update.items():
if key not in before:
return True
if _value_shallow_marker(new_value) != before[key]:
return True
return False
if key not in before or before[key] != new_value:
changed[key] = new_value
return changed or None
def _invoke_node(node: Callable[..., Any], ctx: Context, node_input: Any, state: Any) -> Any:
@@ -1,35 +1,125 @@
import asyncio
from dataclasses import dataclass
from pydantic import BaseModel
from typing_extensions import TypedDict
import pytest
from langgraph.advanced_graph import AdvancedStateGraph
from langgraph.advanced_graph import AdvancedStateGraph, CompiledGraphEngine
from langgraph.types import Command, Send
pytestmark = pytest.mark.anyio
@dataclass(frozen=True)
class DataClassPayload:
value: int
class PydanticPayload(BaseModel):
value: int
class InnerTypedDict(TypedDict):
flag: bool
n: int
class UpdateElisionState(TypedDict):
x: int
dc: DataClassPayload
model: PydanticPayload
td: InnerTypedDict
obj: dict[str, int]
items: list[int]
def _initial_state() -> UpdateElisionState:
return {
"x": 0,
"dc": DataClassPayload(0),
"model": PydanticPayload(value=0),
"td": {"flag": False, "n": 0},
"obj": {"n": 0},
"items": [0],
}
async def test_noop_slow_update_does_not_override_fast_update() -> None:
graph = AdvancedStateGraph(UpdateElisionState)
graph: AdvancedStateGraph[UpdateElisionState] = AdvancedStateGraph(UpdateElisionState)
async def start_node(state: UpdateElisionState) -> Command:
return Command(goto=[Send("fast_node", None), Send("slow_node", None)])
async def fast_node(state: UpdateElisionState) -> dict[str, int]:
return {"x": 1}
async def fast_node(state: UpdateElisionState) -> UpdateElisionState:
state.x = 1
state.dc.value = 1
state.model.value = 1
state.td["flag"] = True
state.td["n"] = 1
state.obj["n"] = 1
state.items.append(1)
return state
async def slow_node(state: UpdateElisionState) -> dict[str, int]:
async def slow_node(state: UpdateElisionState) -> UpdateElisionState:
await asyncio.sleep(0.1)
# Returns the same value as initial snapshot.
return {"x": state["x"]}
# Returns the same values as the initial snapshot.
return state
graph.add_entry_node(start_node)
graph.add_node(fast_node)
graph.add_finish_node(slow_node)
result = await graph.compile().ainvoke({"x": 0})
compiled: CompiledGraphEngine[UpdateElisionState] = graph.compile()
initial_state: UpdateElisionState = _initial_state()
result: UpdateElisionState = await compiled.ainvoke(initial_state)
assert result["x"] == 1
assert result["dc"] == DataClassPayload(1)
assert result["model"].value == 1
assert result["td"] == {"flag": True, "n": 1}
assert result["obj"] == {"n": 1}
assert result["items"] == [1]
async def test_changed_slow_update_overrides_fast_update() -> None:
graph: AdvancedStateGraph[UpdateElisionState] = AdvancedStateGraph(UpdateElisionState)
async def start_node(state: UpdateElisionState) -> Command:
return Command(goto=[Send("fast_node", None), Send("slow_node", None)])
async def fast_node(state: UpdateElisionState) -> UpdateElisionState:
state.x = 1
state.dc.value = 1
state.model.value = 1
state.td["flag"] = True
state.td["n"] = 1
state.obj["n"] = 1
state.items.append(1)
return state
async def slow_node(state: UpdateElisionState) -> UpdateElisionState:
await asyncio.sleep(0.1)
# Slow node makes real changes for all field types.
state.x = 2
state.dc.value = 2
state.model.value = 2
state.td["flag"] = False
state.td["n"] = 2
state.obj["n"] = 2
state.items.append(2)
return state
graph.add_entry_node(start_node)
graph.add_node(fast_node)
graph.add_finish_node(slow_node)
compiled: CompiledGraphEngine[UpdateElisionState] = graph.compile()
initial_state: UpdateElisionState = _initial_state()
result: UpdateElisionState = await compiled.ainvoke(initial_state)
assert result["x"] == 2
assert result["dc"] == DataClassPayload(2)
assert result["model"].value == 2
assert result["td"] == {"flag": False, "n": 2}
assert result["obj"] == {"n": 2}
assert result["items"] == [2]