diff --git a/dargs/dargs.py b/dargs/dargs.py index f067b01..c251138 100644 --- a/dargs/dargs.py +++ b/dargs/dargs.py @@ -480,6 +480,9 @@ def check_value( """ if allow_ref: value = deepcopy(value) + # ``traverse_value`` only checks descendants, so validate the root value + # explicitly before descending into any sub-fields or variants. + self._check_data(value, []) self.traverse_value( value, key_hook=Argument._check_exist, diff --git a/tests/test_checker.py b/tests/test_checker.py index 4486c7a..4395e00 100644 --- a/tests/test_checker.py +++ b/tests/test_checker.py @@ -100,6 +100,16 @@ def test_sub_fields(self) -> None: with self.assertRaises(ValueError): Argument("base", dict, [Argument("sub1", int), Argument("sub1", int)]) + def test_check_value_validates_root(self) -> None: + """Root types and extra checks are enforced by check_value().""" + with self.assertRaises(ArgumentTypeError): + Argument("value", int).check_value("not an integer") + + positive = Argument("value", int, extra_check=lambda value: value > 0) + positive.check_value(1) + with self.assertRaises(ArgumentValueError): + positive.check_value(0) + def test_sub_repeat_list(self) -> None: ca = Argument( "base", list, [Argument("sub1", int), Argument("sub2", str)], repeat=True