diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 8f4d8fa3d..f576a98cd 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -44,7 +44,7 @@ from langgraph.pregel.read import ChannelRead, PregelNode from langgraph.pregel.types import All, RetryPolicy from langgraph.pregel.write import SKIP_WRITE, ChannelWrite, ChannelWriteEntry from langgraph.store.base import BaseStore -from langgraph.utils import RunnableCallable, coerce_to_runnable, field_is_optional +from langgraph.utils import RunnableCallable, coerce_to_runnable, get_field_default logger = logging.getLogger(__name__) @@ -492,13 +492,11 @@ class CompiledStateGraph(CompiledGraph): k: ( self.channels[k].UpdateType, ( - None - if field_is_optional( + get_field_default( k, self.channels[k].UpdateType, self.builder.input, ) - else ... ), ) for k in self.builder.schemas[self.builder.input] diff --git a/libs/langgraph/langgraph/utils.py b/libs/langgraph/langgraph/utils.py index 81ab9b543..1fbcc94e3 100644 --- a/libs/langgraph/langgraph/utils.py +++ b/libs/langgraph/langgraph/utils.py @@ -246,8 +246,8 @@ def _is_readonly_type(type_: Any) -> bool: _DEFAULT_KEYS = frozenset() -def field_is_optional(name: str, type_: Any, schema: Type[Any]) -> bool: - """Determine if the field is optional for a graph's input. +def get_field_default(name: str, type_: Any, schema: Type[Any]) -> Any: + """Determine the default value for a field in a state schema. This is based on: If TypedDict: @@ -259,13 +259,15 @@ def field_is_optional(name: str, type_: Any, schema: Type[Any]) -> bool: if name in optional_keys: # Either total=False or explicit NotRequired. # No type annotation trumps this. - return True + return None if _is_required_type(type_): # Handle Required[] # (we already handled NotRequired and total=False) - return False + return ... # Note, we ignore ReadOnly attributes, # as they don't make much sense. (we don't care if you mutate the state in your node) # and mutating state in your node has no effect on our graph state. # Base case is the annotation - return _is_optional_type(type_) + if _is_optional_type(type_): + return None + return ... diff --git a/libs/langgraph/tests/test_utils.py b/libs/langgraph/tests/test_utils.py index 2e5d2b3bb..936fd2847 100644 --- a/libs/langgraph/tests/test_utils.py +++ b/libs/langgraph/tests/test_utils.py @@ -23,7 +23,7 @@ from langgraph.graph import END, StateGraph from langgraph.graph.graph import CompiledGraph from langgraph.utils import ( _is_optional_type, - field_is_optional, + get_field_default, is_async_callable, is_async_generator, ) @@ -200,17 +200,17 @@ def test_is_required(): val_10: str annos = MyBaseTypedDict.__annotations__ - assert not field_is_optional("val_1", annos["val_1"], MyBaseTypedDict) - assert not field_is_optional("val_2", annos["val_2"], MyBaseTypedDict) - assert field_is_optional("val_3", annos["val_3"], MyBaseTypedDict) - assert field_is_optional("val_4", annos["val_4"], MyBaseTypedDict) + assert get_field_default("val_1", annos["val_1"], MyBaseTypedDict) == ... + assert get_field_default("val_2", annos["val_2"], MyBaseTypedDict) == ... + assert get_field_default("val_3", annos["val_3"], MyBaseTypedDict) is None + assert get_field_default("val_4", annos["val_4"], MyBaseTypedDict) is None # See https://peps.python.org/pep-0655/#interaction-with-annotated - assert field_is_optional("val_5", annos["val_5"], MyBaseTypedDict) - assert field_is_optional("val_6", annos["val_6"], MyBaseTypedDict) - assert not field_is_optional("val_7", annos["val_7"], MyBaseTypedDict) - assert not field_is_optional("val_8", annos["val_8"], MyBaseTypedDict) - assert field_is_optional("val_9", annos["val_9"], MyBaseTypedDict) - assert not field_is_optional("val_10", annos["val_10"], MyBaseTypedDict) + assert get_field_default("val_5", annos["val_5"], MyBaseTypedDict) is None + assert get_field_default("val_6", annos["val_6"], MyBaseTypedDict) is None + assert get_field_default("val_7", annos["val_7"], MyBaseTypedDict) == ... + assert get_field_default("val_8", annos["val_8"], MyBaseTypedDict) == ... + assert get_field_default("val_9", annos["val_9"], MyBaseTypedDict) is None + assert get_field_default("val_10", annos["val_10"], MyBaseTypedDict) == ... class MyChildDict(MyBaseTypedDict): val_11: int @@ -223,9 +223,9 @@ def test_is_required(): cannos = MyChildDict.__annotations__ gcannos = MyGrandChildDict.__annotations__ - assert not field_is_optional("val_11", cannos["val_11"], MyChildDict) - assert field_is_optional("val_11b", cannos["val_11b"], MyChildDict) - assert field_is_optional("val_11c", cannos["val_11c"], MyChildDict) - assert field_is_optional("val_12", gcannos["val_12"], MyGrandChildDict) - assert field_is_optional("val_9", gcannos["val_9"], MyGrandChildDict) - assert not field_is_optional("val_13", gcannos["val_13"], MyGrandChildDict) + assert get_field_default("val_11", cannos["val_11"], MyChildDict) == ... + assert get_field_default("val_11b", cannos["val_11b"], MyChildDict) is None + assert get_field_default("val_11c", cannos["val_11c"], MyChildDict) is None + assert get_field_default("val_12", gcannos["val_12"], MyGrandChildDict) is None + assert get_field_default("val_9", gcannos["val_9"], MyGrandChildDict) is None + assert get_field_default("val_13", gcannos["val_13"], MyGrandChildDict) == ...