Return a default value instead

This commit is contained in:
William Fu-Hinthorn
2024-08-30 16:28:01 -07:00
parent 51f58279ba
commit 55263c640f
3 changed files with 26 additions and 26 deletions
+2 -4
View File
@@ -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]
+7 -5
View File
@@ -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 ...
+17 -17
View File
@@ -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) == ...