move to fields

This commit is contained in:
vbarda
2025-04-14 13:00:03 -04:00
parent d4224a7abb
commit 2c557e9e46
4 changed files with 41 additions and 39 deletions
+2 -2
View File
@@ -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__)
+1 -1
View File
@@ -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
+37 -1
View File
@@ -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)
)
]
+1 -35
View File
@@ -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.