From dc6fa9ed30daa526e299a11a9e164d8ab4840500 Mon Sep 17 00:00:00 2001 From: vbarda Date: Fri, 11 Apr 2025 17:51:56 -0400 Subject: [PATCH 1/5] langgraph: handle pydantic updates consistently in Command --- libs/langgraph/langgraph/graph/state.py | 4 ++-- libs/langgraph/langgraph/types.py | 30 ++++++++++++++++++++++++- 2 files changed, 31 insertions(+), 3 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index ce2f898d0..849ff6b47 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -764,14 +764,14 @@ class CompiledStateGraph(CompiledGraph): updates.extend(_get_updates(i) or ()) return updates elif (t := type(input)) and get_type_hints(t): - # Pydantic v2 + # Pydantic v1 if isinstance(input, BaseModelV1): keep: Optional[set[str]] = input.__fields_set__ defaults = {k: v.default for k, v in t.__fields__.items()} + # Pydantic v2 elif isinstance(input, BaseModel): keep = input.model_fields_set defaults = {k: v.default for k, v in input.model_fields.items()} - # Pydantic v1 else: keep = None defaults = {} diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index 9ce1c2c66..a2a83b55d 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -20,6 +20,8 @@ from typing import ( ) from langchain_core.runnables import Runnable, RunnableConfig +from pydantic import BaseModel +from pydantic.v1 import BaseModel as BaseModelV1 from typing_extensions import Self from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata @@ -69,6 +71,11 @@ else: _DC_KWARGS = {"frozen": True} +# NOTE: this is redefined here separately from langgraph.constants +# to avoid a circular import +MISSING = object() + + def default_retry_on(exc: Exception) -> bool: import httpx import requests @@ -318,7 +325,28 @@ class Command(Generic[N], ToolOutputMixin): ): return self.update elif hints := get_type_hints(type(self.update)): - return [(k, getattr(self.update, k)) for k in hints] + # Pydantic v1 + if isinstance(self.update, BaseModelV1): + keep: Optional[set[str]] = self.update.__fields_set__ + defaults = {k: v.default for k, v in self.update.__fields__.items()} + # Pydantic v2 + elif isinstance(self.update, BaseModel): + keep = self.update.model_fields_set + defaults = {k: v.default for k, v in self.update.model_fields.items()} + else: + keep = None + defaults = {} + + return [ + (k, value) + for k in hints + if (value := getattr(self.update, k, MISSING)) is not MISSING + and ( + value is not None + or defaults.get(k, MISSING) is not None + or (keep is not None and k in keep) + ) + ] elif self.update is not None: return [("__root__", self.update)] else: From 2ed453debe4dee81aaabcbf31b9f6778cadb338a Mon Sep 17 00:00:00 2001 From: vbarda Date: Sat, 12 Apr 2025 10:34:02 -0400 Subject: [PATCH 2/5] factor out util --- libs/langgraph/langgraph/graph/state.py | 29 ++--------------- libs/langgraph/langgraph/types.py | 31 ++----------------- libs/langgraph/langgraph/utils/pydantic.py | 36 +++++++++++++++++++++- 3 files changed, 39 insertions(+), 57 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 849ff6b47..85acf73b3 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -77,7 +77,7 @@ from langgraph.pregel.write import ( from langgraph.store.base import BaseStore from langgraph.types import All, Checkpointer, Command, RetryPolicy from langgraph.utils.fields import get_field_default -from langgraph.utils.pydantic import create_model +from langgraph.utils.pydantic import create_model, get_update_as_tuples from langgraph.utils.runnable import RunnableCallable, RunnableLike, coerce_to_runnable logger = logging.getLogger(__name__) @@ -764,32 +764,7 @@ class CompiledStateGraph(CompiledGraph): updates.extend(_get_updates(i) or ()) return updates elif (t := type(input)) and get_type_hints(t): - # Pydantic v1 - if isinstance(input, BaseModelV1): - keep: Optional[set[str]] = input.__fields_set__ - defaults = {k: v.default for k, v in t.__fields__.items()} - # Pydantic v2 - elif isinstance(input, BaseModel): - keep = input.model_fields_set - defaults = {k: v.default for k, v in input.model_fields.items()} - else: - keep = None - defaults = {} - - # NOTE: This behavior for Pydantic is somewhat inelegant, - # but we keep around for backwards compatibility - # if input is a Pydantic model, only update values - # that are different from the default values or in the keep set - return [ - (k, value) - for k in output_keys - if (value := getattr(input, k, MISSING)) is not MISSING - and ( - value is not None - or defaults.get(k, MISSING) is not None - or (keep is not None and k in keep) - ) - ] + return get_update_as_tuples(input, output_keys) else: msg = create_error_message( message=f"Expected dict, got {input}", diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index a2a83b55d..6acdb3a6b 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -20,11 +20,10 @@ from typing import ( ) from langchain_core.runnables import Runnable, RunnableConfig -from pydantic import BaseModel -from pydantic.v1 import BaseModel as BaseModelV1 from typing_extensions import Self from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata +from langgraph.utils.pydantic import get_update_as_tuples if TYPE_CHECKING: from langgraph.pregel.protocol import PregelProtocol @@ -71,11 +70,6 @@ else: _DC_KWARGS = {"frozen": True} -# NOTE: this is redefined here separately from langgraph.constants -# to avoid a circular import -MISSING = object() - - def default_retry_on(exc: Exception) -> bool: import httpx import requests @@ -325,28 +319,7 @@ class Command(Generic[N], ToolOutputMixin): ): return self.update elif hints := get_type_hints(type(self.update)): - # Pydantic v1 - if isinstance(self.update, BaseModelV1): - keep: Optional[set[str]] = self.update.__fields_set__ - defaults = {k: v.default for k, v in self.update.__fields__.items()} - # Pydantic v2 - elif isinstance(self.update, BaseModel): - keep = self.update.model_fields_set - defaults = {k: v.default for k, v in self.update.model_fields.items()} - else: - keep = None - defaults = {} - - return [ - (k, value) - for k in hints - if (value := getattr(self.update, k, MISSING)) is not MISSING - and ( - value is not None - or defaults.get(k, MISSING) is not None - or (keep is not None and k in keep) - ) - ] + return get_update_as_tuples(self.update, tuple(hints.keys())) elif self.update is not None: return [("__root__", self.update)] else: diff --git a/libs/langgraph/langgraph/utils/pydantic.py b/libs/langgraph/langgraph/utils/pydantic.py index 56cef30e6..35aa53bed 100644 --- a/libs/langgraph/langgraph/utils/pydantic.py +++ b/libs/langgraph/langgraph/utils/pydantic.py @@ -1,12 +1,16 @@ import sys import typing from dataclasses import is_dataclass -from typing import Any, Dict, Optional, Union +from typing import Any, Dict, Optional, Sequence, Union import typing_extensions from pydantic import BaseModel from pydantic.v1 import BaseModel as BaseModelV1 +# NOTE: this is redefined here separately from langgraph.constants +# to avoid a circular import +MISSING = object() + def create_model( model_name: str, @@ -41,6 +45,36 @@ def create_model( return create_model(model_name, **v1_kwargs, **(field_definitions or {})) +def get_update_as_tuples(input: Any, keys: Sequence[str]) -> list[tuple[str, Any]]: + """Get Pydantic state update as a list of (key, value) tuples.""" + # Pydantic v1 + if isinstance(input, BaseModelV1): + keep: Optional[set[str]] = input.__fields_set__ + defaults = {k: v.default for k, v in input.__fields__.items()} + # Pydantic v2 + elif isinstance(input, BaseModel): + keep = input.model_fields_set + defaults = {k: v.default for k, v in input.model_fields.items()} + else: + keep = None + defaults = {} + + # NOTE: This behavior for Pydantic is somewhat inelegant, + # but we keep around for backwards compatibility + # if input is a Pydantic model, only update values + # that are different from the default values or in the keep set + return [ + (k, value) + for k in keys + if (value := getattr(input, k, MISSING)) is not MISSING + and ( + value is not None + or defaults.get(k, MISSING) is not None + or (keep is not None and k in keep) + ) + ] + + def is_supported_by_pydantic(type_: Any) -> bool: """Check if a given "complex" type is supported by pydantic. From 704b78b8fe3cb7df5baba2cd3e5a11da1bfdc39c Mon Sep 17 00:00:00 2001 From: vbarda Date: Sat, 12 Apr 2025 10:45:10 -0400 Subject: [PATCH 3/5] tests --- libs/langgraph/tests/test_pregel.py | 67 +++++++++++++++++++++++++++++ 1 file changed, 67 insertions(+) diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index be3f94ec7..1f85a8357 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -7246,6 +7246,39 @@ def test_pydantic_none_state_update() -> None: assert graph.invoke({"foo": ""}) == {"foo": None} +def test_pydantic_state_update_command() -> None: + from pydantic import BaseModel + + class State(BaseModel): + foo: Optional[str] + + def node_a(state: State) -> State: + return Command(update=State(foo=None)) + + graph = StateGraph(State).add_node(node_a).add_edge(START, "node_a").compile() + assert graph.invoke({"foo": ""}) == {"foo": None} + + class State(BaseModel): + foo: str | None = None + bar: str | None = None + + def node_a(state: State): + return State(foo="foo") + + def node_b(state: State): + return Command(update=State(bar="bar")) + + builder = StateGraph(State) + builder.add_node(node_a) + builder.add_node(node_b) + builder.add_edge(START, "node_a") + builder.add_edge("node_a", "node_b") + builder.add_edge("node_b", END) + graph = builder.compile() + + assert graph.invoke(State()) == {"foo": "foo", "bar": "bar"} + + def test_pydantic_state_mutation() -> None: from pydantic import BaseModel, Field @@ -7280,6 +7313,40 @@ def test_pydantic_state_mutation() -> None: assert graph.invoke({"outer": 1}) == {"outer": 10, "inner": Inner(a=5)} +def test_pydantic_state_mutation_command() -> None: + from pydantic import BaseModel, Field + + class Inner(BaseModel): + a: int = 0 + + class State(BaseModel): + inner: Inner = Inner() + outer: int = 0 + + def my_node(state: State) -> State: + state.inner.a = 5 + state.outer = 10 + return Command(update=state) + + graph = StateGraph(State).add_node(my_node).add_edge(START, "my_node").compile() + + assert graph.invoke({"outer": 1}) == {"outer": 10, "inner": Inner(a=5)} + + # test w/ default_factory + class State(BaseModel): + inner: Inner = Field(default_factory=Inner) + outer: int = 0 + + def my_node(state: State) -> State: + state.inner.a = 5 + state.outer = 10 + return Command(update=state) + + graph = StateGraph(State).add_node(my_node).add_edge(START, "my_node").compile() + + assert graph.invoke({"outer": 1}) == {"outer": 10, "inner": Inner(a=5)} + + def test_get_stream_writer() -> None: class State(TypedDict): foo: str From 2c557e9e46743357b041bc13959ad5ff4fab3932 Mon Sep 17 00:00:00 2001 From: vbarda Date: Mon, 14 Apr 2025 13:00:03 -0400 Subject: [PATCH 4/5] move to fields --- libs/langgraph/langgraph/graph/state.py | 4 +-- libs/langgraph/langgraph/types.py | 2 +- libs/langgraph/langgraph/utils/fields.py | 38 +++++++++++++++++++++- libs/langgraph/langgraph/utils/pydantic.py | 36 +------------------- 4 files changed, 41 insertions(+), 39 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 1a58e023f..fed12fd8b 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -77,8 +77,8 @@ from langgraph.pregel.write import ( ) from langgraph.store.base import BaseStore from langgraph.types import All, Checkpointer, Command, RetryPolicy -from langgraph.utils.fields import get_field_default -from langgraph.utils.pydantic import create_model, get_update_as_tuples +from langgraph.utils.fields import get_field_default, get_update_as_tuples +from langgraph.utils.pydantic import create_model from langgraph.utils.runnable import RunnableLike, coerce_to_runnable logger = logging.getLogger(__name__) diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index 6acdb3a6b..195acd4ff 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -23,7 +23,7 @@ from langchain_core.runnables import Runnable, RunnableConfig from typing_extensions import Self from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata -from langgraph.utils.pydantic import get_update_as_tuples +from langgraph.utils.fields import get_update_as_tuples if TYPE_CHECKING: from langgraph.pregel.protocol import PregelProtocol diff --git a/libs/langgraph/langgraph/utils/fields.py b/libs/langgraph/langgraph/utils/fields.py index 009e5aee1..e39f171bc 100644 --- a/libs/langgraph/langgraph/utils/fields.py +++ b/libs/langgraph/langgraph/utils/fields.py @@ -1,8 +1,14 @@ import dataclasses -from typing import Any, Generator, Optional, Type, Union, get_type_hints +from typing import Any, Generator, Optional, Sequence, Type, Union, get_type_hints +from pydantic import BaseModel +from pydantic.v1 import BaseModel as BaseModelV1 from typing_extensions import Annotated, NotRequired, ReadOnly, Required, get_origin +# NOTE: this is redefined here separately from langgraph.constants +# to avoid a circular import +MISSING = object() + def _is_optional_type(type_: Any) -> bool: """Check if a type is Optional.""" @@ -147,3 +153,33 @@ def get_enhanced_type_hints( pass yield name, typ, default, description + + +def get_update_as_tuples(input: Any, keys: Sequence[str]) -> list[tuple[str, Any]]: + """Get Pydantic state update as a list of (key, value) tuples.""" + # Pydantic v1 + if isinstance(input, BaseModelV1): + keep: Optional[set[str]] = input.__fields_set__ + defaults = {k: v.default for k, v in input.__fields__.items()} + # Pydantic v2 + elif isinstance(input, BaseModel): + keep = input.model_fields_set + defaults = {k: v.default for k, v in input.model_fields.items()} + else: + keep = None + defaults = {} + + # NOTE: This behavior for Pydantic is somewhat inelegant, + # but we keep around for backwards compatibility + # if input is a Pydantic model, only update values + # that are different from the default values or in the keep set + return [ + (k, value) + for k in keys + if (value := getattr(input, k, MISSING)) is not MISSING + and ( + value is not None + or defaults.get(k, MISSING) is not None + or (keep is not None and k in keep) + ) + ] diff --git a/libs/langgraph/langgraph/utils/pydantic.py b/libs/langgraph/langgraph/utils/pydantic.py index 35aa53bed..56cef30e6 100644 --- a/libs/langgraph/langgraph/utils/pydantic.py +++ b/libs/langgraph/langgraph/utils/pydantic.py @@ -1,16 +1,12 @@ import sys import typing from dataclasses import is_dataclass -from typing import Any, Dict, Optional, Sequence, Union +from typing import Any, Dict, Optional, Union import typing_extensions from pydantic import BaseModel from pydantic.v1 import BaseModel as BaseModelV1 -# NOTE: this is redefined here separately from langgraph.constants -# to avoid a circular import -MISSING = object() - def create_model( model_name: str, @@ -45,36 +41,6 @@ def create_model( return create_model(model_name, **v1_kwargs, **(field_definitions or {})) -def get_update_as_tuples(input: Any, keys: Sequence[str]) -> list[tuple[str, Any]]: - """Get Pydantic state update as a list of (key, value) tuples.""" - # Pydantic v1 - if isinstance(input, BaseModelV1): - keep: Optional[set[str]] = input.__fields_set__ - defaults = {k: v.default for k, v in input.__fields__.items()} - # Pydantic v2 - elif isinstance(input, BaseModel): - keep = input.model_fields_set - defaults = {k: v.default for k, v in input.model_fields.items()} - else: - keep = None - defaults = {} - - # NOTE: This behavior for Pydantic is somewhat inelegant, - # but we keep around for backwards compatibility - # if input is a Pydantic model, only update values - # that are different from the default values or in the keep set - return [ - (k, value) - for k in keys - if (value := getattr(input, k, MISSING)) is not MISSING - and ( - value is not None - or defaults.get(k, MISSING) is not None - or (keep is not None and k in keep) - ) - ] - - def is_supported_by_pydantic(type_: Any) -> bool: """Check if a given "complex" type is supported by pydantic. From 0e111b2f44b2ac22378b5f1e3e50d5a74a31ec05 Mon Sep 17 00:00:00 2001 From: vbarda Date: Mon, 14 Apr 2025 13:23:26 -0400 Subject: [PATCH 5/5] 3.9 --- libs/langgraph/tests/test_pregel.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 1f85a8357..82ba6ff96 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -7259,8 +7259,8 @@ def test_pydantic_state_update_command() -> None: assert graph.invoke({"foo": ""}) == {"foo": None} class State(BaseModel): - foo: str | None = None - bar: str | None = None + foo: Optional[str] = None + bar: Optional[str] = None def node_a(state: State): return State(foo="foo")