diff --git a/langgraph/graph/state.py b/langgraph/graph/state.py index dfdac148d..0043b6d2e 100644 --- a/langgraph/graph/state.py +++ b/langgraph/graph/state.py @@ -1,3 +1,4 @@ +from dataclasses import dataclass import logging import typing import warnings @@ -291,10 +292,13 @@ class CompiledStateGraph(CompiledGraph): def _get_state_key(input: dict, config: RunnableConfig, *, key: str) -> Any: if input is None: return SKIP_WRITE - elif not isinstance(input, dict): - raise InvalidUpdateError(f"Expected dict, got {input}") - else: + elif isinstance(input, dict): return input.get(key, SKIP_WRITE) + elif get_type_hints(type(input)).get(key): + value = getattr(input, key, SKIP_WRITE) + return value if value is not None else SKIP_WRITE + else: + raise InvalidUpdateError(f"Expected dict, got {input}") # state updaters state_write_entries = ( diff --git a/tests/test_pregel.py b/tests/test_pregel.py index 9df1d382f..c58ff7e7b 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -6029,11 +6029,16 @@ def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class( answer: Optional[str] = None docs: Annotated[list[str], sorted_add] + class StateUpdate(BaseModel): + query: Optional[str] = None + answer: Optional[str] = None + docs: Optional[list[str]] = None + def rewrite_query(data: State) -> State: return {"query": f"query: {data.query}"} def analyzer_one(data: State) -> State: - return {"query": f"analyzed: {data.query}"} + return StateUpdate(query=f"analyzed: {data.query}") def retriever_one(data: State) -> State: return {"docs": ["doc1", "doc2"]} diff --git a/tests/test_pregel_async.py b/tests/test_pregel_async.py index e410ea8d1..34e4a783c 100644 --- a/tests/test_pregel_async.py +++ b/tests/test_pregel_async.py @@ -4524,11 +4524,16 @@ async def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class( answer: Optional[str] = None docs: Annotated[list[str], sorted_add] + class StateUpdate(BaseModel): + query: Optional[str] = None + answer: Optional[str] = None + docs: Optional[list[str]] = None + async def rewrite_query(data: State) -> State: return {"query": f"query: {data.query}"} async def analyzer_one(data: State) -> State: - return {"query": f"analyzed: {data.query}"} + return StateUpdate(query=f"analyzed: {data.query}") async def retriever_one(data: State) -> State: return {"docs": ["doc1", "doc2"]}