mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-29 21:15:11 +02:00
more
This commit is contained in:
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user