mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 04:25:08 +02:00
Return a default value instead
This commit is contained in:
@@ -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]
|
||||
|
||||
@@ -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[<type>]
|
||||
# (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 ...
|
||||
|
||||
@@ -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) == ...
|
||||
|
||||
Reference in New Issue
Block a user