From 661476e88d163e224a079f683a438a19daf8889d Mon Sep 17 00:00:00 2001 From: vbarda Date: Wed, 12 Feb 2025 17:50:25 -0800 Subject: [PATCH 1/4] langgraph: fix None handling for pydantic state updates --- libs/langgraph/langgraph/graph/state.py | 4 ++-- libs/langgraph/tests/test_pregel.py | 13 +++++++++++++ 2 files changed, 15 insertions(+), 2 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 8c93c0c1d..ac79cc497 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -33,7 +33,7 @@ from langgraph.channels.dynamic_barrier_value import DynamicBarrierValue, WaitFo from langgraph.channels.ephemeral_value import EphemeralValue from langgraph.channels.last_value import LastValue from langgraph.channels.named_barrier_value import NamedBarrierValue -from langgraph.constants import EMPTY_SEQ, NS_END, NS_SEP, SELF, TAG_HIDDEN +from langgraph.constants import EMPTY_SEQ, MISSING, NS_END, NS_SEP, SELF, TAG_HIDDEN from langgraph.errors import ( ErrorCode, InvalidUpdateError, @@ -694,7 +694,7 @@ class CompiledStateGraph(CompiledGraph): return [ (k, getattr(input, k)) for k in output_keys - if getattr(input, k, None) is not None + if getattr(input, k, MISSING) is not MISSING ] else: msg = create_error_message( diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index ccd4fae94..6c3603499 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -6482,3 +6482,16 @@ def test_node_destinations() -> None: Edge(source="child", target="node_b", data="foo", conditional=True), Edge(source="child", target="node_c", data="bar", conditional=True), ] == graph.edges + + +def test_pydantic_none_state_update() -> None: + from pydantic import BaseModel + + class State(BaseModel): + foo: str | None + + def node_a(state: State) -> State: + return State(foo=None) + + graph = StateGraph(State).add_node(node_a).add_edge(START, "node_a").compile() + assert graph.invoke({"foo": ""}) == {"foo": None} From c36323cba8e16cf90cfc14c2bb94a228c9cb6d31 Mon Sep 17 00:00:00 2001 From: vbarda Date: Wed, 12 Feb 2025 19:17:05 -0800 Subject: [PATCH 2/4] 3.9 --- libs/langgraph/tests/test_pregel.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 6c3603499..7e15db62f 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -6488,7 +6488,7 @@ def test_pydantic_none_state_update() -> None: from pydantic import BaseModel class State(BaseModel): - foo: str | None + foo: Optional[str] def node_a(state: State) -> State: return State(foo=None) From b0e11ae52490a4ea75abb38ba153f50707582ff4 Mon Sep 17 00:00:00 2001 From: vbarda Date: Wed, 12 Feb 2025 20:02:26 -0800 Subject: [PATCH 3/4] update --- libs/langgraph/langgraph/graph/state.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index ac79cc497..8c95b9581 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -691,9 +691,17 @@ class CompiledStateGraph(CompiledGraph): updates.extend(_get_updates(i) or ()) return updates elif get_type_hints(type(input)): + # if input is a Pydantic model, only update values + # for the keys that have been explicitly set by the users + # (this is needed to avoid sending updates for fields with None defaults) + output_keys_ = ( + [k for k in output_keys if k in input.model_fields_set] + if hasattr(input, "model_fields_set") + else output_keys + ) return [ (k, getattr(input, k)) - for k in output_keys + for k in output_keys_ if getattr(input, k, MISSING) is not MISSING ] else: From 04e8342d972271053c90aa3b96a6e8c63e2c8c3e Mon Sep 17 00:00:00 2001 From: vbarda Date: Wed, 12 Feb 2025 20:46:12 -0800 Subject: [PATCH 4/4] pydantic v1 --- libs/langgraph/langgraph/graph/state.py | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 8c95b9581..ae8941618 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -694,11 +694,16 @@ class CompiledStateGraph(CompiledGraph): # if input is a Pydantic model, only update values # for the keys that have been explicitly set by the users # (this is needed to avoid sending updates for fields with None defaults) - output_keys_ = ( - [k for k in output_keys if k in input.model_fields_set] - if hasattr(input, "model_fields_set") - else output_keys - ) + output_keys_ = output_keys + # Pydantic v2 + if hasattr(input, "model_fields_set"): + output_keys_ = [ + k for k in output_keys if k in input.model_fields_set + ] + # Pydantic v1 + elif hasattr(input, "__fields_set__"): + output_keys_ = [k for k in output_keys if k in input.__fields_set__] + return [ (k, getattr(input, k)) for k in output_keys_