Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 29 additions & 0 deletions dagapp/tests/test_not_set_defaults.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
"""``i2``'s ``NotSet`` sentinel in a signature means "required / no default".

See i2mint/i2#48: once ``i2.FuncFactory`` shows ``NotSet`` defaults, a DAG root with
such a default must get the same initial value as one with no default at all.
These tests use ``i2.deco.NotSet`` directly, so they pass with any i2 version.
"""

from i2.deco import NotSet
from meshed import DAG

from dagapp.utils import get_root_values


def f(a: int, b, c: float = 2.0):
return a + b + c


# ``f`` with ``NotSet`` defaults, as a re-landed i2#88 ``FuncFactory`` would show.
def f_with_not_set(a: int = NotSet, b=NotSet, c: float = 2.0):
return a + b + c


f_with_not_set.__name__ = f.__name__


def test_root_values_ignore_not_set():
got = get_root_values(DAG([f_with_not_set]))
assert NotSet not in got.values()
assert got == get_root_values(DAG([f]))
29 changes: 26 additions & 3 deletions dagapp/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,27 @@
import inspect
from inspect import Parameter

try:
from i2 import is_not_set
except ImportError: # older i2: same sentinel, not exported from the root yet
from i2.deco import NotSet as _NotSet

def is_not_set(x) -> bool:
"""Return True iff ``x`` is ``i2``'s ``NotSet`` sentinel."""
return x is _NotSet


def _has_default(param: Parameter) -> bool:
"""True iff ``param`` has a real default (``i2``'s ``NotSet`` doesn't count).

>>> from i2.deco import NotSet
>>> [_has_default(Parameter('x', Parameter.KEYWORD_ONLY, default=d))
... for d in (None, 0, NotSet, Parameter.empty)]
[True, True, False, False]
"""
return param.default is not Parameter.empty and not is_not_set(param.default)


DFLT_VALS = {
int: 0,
float: 0.0,
Expand Down Expand Up @@ -281,7 +302,7 @@ def _compute_node_value(node, funcs):
# If the parameter isn't present in session state but has a default,
# omit it so the function can use its default value.
if name not in st.session_state:
if param.default is not inspect._empty:
if _has_default(param):
continue
# preserve previous behaviour: accessing missing keys will raise
# the same KeyError as before
Expand Down Expand Up @@ -309,9 +330,11 @@ def get_root_values(dag):
Returns the default values for all the root nodes found in dag
"""
root_defaults = dict()
# i2's NotSet sentinel means "no default": fall back on the annotation or 0.0
defaults = {k: v for k, v in dag.sig.defaults.items() if not is_not_set(v)}
for name in dag.sig.names:
if name in dag.sig.defaults:
dflt = dag.sig.defaults[name]
if name in defaults:
dflt = defaults[name]
if dflt is not None:
root_defaults[name] = dflt
elif name in dag.sig.annotations:
Expand Down
Loading