[core] Let a trigger filter name its parent with {parent} (#20448)

This commit is contained in:
J. Nick Koston
2026-10-09 15:43:22 -10:00
committed by GitHub
parent 02bfa07819
commit 073394b1eb
2 changed files with 102 additions and 20 deletions
+25 -14
View File
@@ -264,13 +264,17 @@ def string_ref_literal(config: ConfigType, value: str) -> str:
return f"StringRef({literal_with_length(config, value)})"
def _names_parent(text: str) -> bool:
return any(f == "parent" for _, f, _, _ in string.Formatter().parse(text))
@dataclass(frozen=True)
class ApplyCall:
"""One statement from config keys, e.g. ``"set_range({}, {})"`` with ``((CONF_LOW, cg.float_), ...)``.
Each arg is ``(conf_key, type_)`` or ``(conf_key, type_, const_fn)``. A ``conf_key`` may be a
path into nested sections. A plain ``str`` ``type_`` is raw C++ type text and may use
``{parent}``. ``const_fn(config, value)`` renders a constant's argument text; a lambda or an
path into nested sections. ``target`` and a plain ``str`` ``type_`` (raw C++ type text) may
name the parent object as ``{parent}``. ``const_fn(config, value)`` renders a constant's argument text; a lambda or an
id bypasses it. The statement is skipped when none of its keys is set, always emitted when it
has no keys, and a partial set is a config error.
"""
@@ -282,13 +286,13 @@ class ApplyCall:
fields = [
f for _, f, _, _ in string.Formatter().parse(self.target) if f is not None
]
if any(fields):
if any(f not in ("", "parent") for f in fields):
raise ValueError(
f"apply target {self.target!r}: only bare {{}} placeholders"
f"apply target {self.target!r}: only {{}} and {{parent}} placeholders"
)
if len(fields) != len(self.args):
if (count := fields.count("")) != len(self.args):
raise ValueError(
f"apply target {self.target!r} has {len(fields)} "
f"apply target {self.target!r} has {count} "
f"placeholder(s) for {len(self.args)} config key(s)"
)
if any(len(arg) not in (2, 3) for arg in self.args):
@@ -296,6 +300,12 @@ class ApplyCall:
f"apply target {self.target!r}: each arg is (conf_key, type_[, const_fn])"
)
@property
def names_parent(self) -> bool:
return _names_parent(self.target) or any(
isinstance(arg[1], str) and _names_parent(arg[1]) for arg in self.args
)
@property
def members(self) -> list[tuple[Any, Any, Any]]:
"""Each arg as ``(conf_key, type_, const_fn or None)``."""
@@ -494,7 +504,7 @@ def register_apply_action(
exprs = await _render_values(
name, target, members, values, config, parent, lambda_args
)
statements.append(f"{receiver}{target.format(*exprs)};")
statements.append(f"{receiver}{target.format(*exprs, parent=parent)};")
if call:
statements = [
f"auto apply_call = {parent}->{call}();",
@@ -526,7 +536,7 @@ async def _render_check(
exprs = await _render_values(
name, target, members, values, config, parent, lambda_args, compare=True
)
return target.format(*exprs)
return target.format(*exprs, parent=parent)
def register_apply_condition(
@@ -1061,21 +1071,22 @@ async def build_trigger_callback(
params: TemplateArgsType,
forward: Sequence[str | Expression] | None = None,
when: str | ApplyCall | None = None,
parent: MockObj | None = None,
) -> LambdaExpression:
"""Build the Automation for ``config`` and return a stateless callback that triggers it.
``params`` are the parent callback's parameters, ``forward`` the expressions passed to
``trigger()`` (default: the parameter names; write the parent as ``parent_ref(var)``),
``when`` a filter the callback returns early on, skipped like any ``ApplyCall`` when none
of its keys is set.
of its keys is set. ``when`` may name ``parent`` as ``{parent}``, e.g.
``"{parent}->is_fully_open()"``.
"""
members: list[tuple[Any, Any, Any]] = []
if when is not None:
call = when if isinstance(when, ApplyCall) else ApplyCall(when)
members = call.members
# A trigger callback has no parent for a str type to name.
if any(isinstance(t, str) and "{parent}" in t for _, t, _ in members):
raise ValueError(f"trigger filter {call.target!r}: a type names {{parent}}")
if parent is None and call.names_parent:
raise ValueError(f"trigger filter {call.target!r} names {{parent}}")
obj = await _new_automation(args, config)
lambda_args = _apply_lambda_args(params)
statements: list[str] = []
@@ -1088,7 +1099,7 @@ async def build_trigger_callback(
members,
values,
config,
None,
None if parent is None else str(parent_ref(parent)),
lambda_args,
)
statements.append(f"if (!({check}))\n return;")
@@ -1137,7 +1148,7 @@ async def build_callback_automation(
"forwarder cannot be combined with params, forward or when"
)
callback = await build_trigger_callback(
args, config, args if params is None else params, forward, when
args, config, args if params is None else params, forward, when, parent
)
cg.add(getattr(parent, callback_method)(callback))
return
+77 -6
View File
@@ -943,6 +943,30 @@ async def test_apply_action_call_shape(
assert positions == sorted(positions)
@pytest.mark.asyncio
async def test_apply_target_names_parent(
registries: tuple[Registry, Registry], mock_cg: MockCodegen
) -> None:
"""{parent} in a target is the global-scoped parent, also beside a call object."""
fields = (ApplyCall("set_peer({parent}->id(), {})", (("level", cg.float_),)),)
await _run_apply_action(registries, fields, {"level": 0.5})
assert (
f"::{PARENT_OBJ}->set_peer(::{PARENT_OBJ}->id(), 0.5f);"
in _apply_definition(mock_cg)
)
mock_cg.new_pvariable.reset_mock()
await _run_apply_action(registries, fields, {"level": 0.5}, call="make_call")
assert f"apply_call.set_peer(::{PARENT_OBJ}->id(), 0.5f);" in _apply_definition(
mock_cg
)
mock_cg.new_pvariable.reset_mock()
await _run_apply_condition(registries, "is_peer({parent}->id())", {})
assert (
f"return ::{PARENT_OBJ}->is_peer(::{PARENT_OBJ}->id());"
in _apply_definition(mock_cg)
)
@pytest.mark.asyncio
async def test_apply_field_nested_key_const_fn_and_type_string(
registries: tuple[Registry, Registry], mock_cg: MockCodegen
@@ -979,8 +1003,10 @@ async def test_apply_field_nested_key_const_fn_and_type_string(
def test_apply_registration_checks(registries: tuple[Registry, Registry]) -> None:
with pytest.raises(ValueError, match="2 placeholder"):
ApplyCall("set_range({}, {})", (("low", cg.float_),))
with pytest.raises(ValueError, match="only bare"):
ApplyCall("if ({}) {parent}->reset()", (("reset", cg.bool_),))
with pytest.raises(ValueError, match="only {} and {parent}"):
ApplyCall("if ({}) {other}->reset()", (("reset", cg.bool_),))
assert ApplyCall("if ({}) {parent}->reset()", (("reset", cg.bool_),)).names_parent
assert not ApplyCall("set_flags({{parent}})").names_parent
ApplyCall("set_flags({{{}}})", (("flags", cg.int_),))
with pytest.raises(ValueError, match="each arg is"):
ApplyCall("set_kp({})", (("kp", cg.float_, None, "extra"),))
@@ -1242,12 +1268,57 @@ async def test_trigger_callback_reshapes_and_filters(mock_cg: MockCodegen) -> No
@pytest.mark.asyncio
async def test_trigger_callback_filter_has_no_parent(mock_cg: MockCodegen) -> None:
"""A filter type that names {parent} is rejected up front, a trigger callback has none."""
when = ApplyCall("mode == {}", (("mode", "{parent}::Mode"),))
with pytest.raises(ValueError, match="names {parent}"):
"""A filter that names {parent} is rejected up front when no parent is given."""
for when in (
ApplyCall("mode == {}", (("mode", "{parent}::Mode"),)),
"{parent}->is_failed() == false",
):
with pytest.raises(ValueError, match="names {parent}"):
await build_trigger_callback(
[],
{**TRIGGER_CONF, "mode": 1},
[(cg.int_, "mode")],
forward=[],
when=when,
)
@pytest.mark.asyncio
async def test_trigger_callback_filter_names_parent(mock_cg: MockCodegen) -> None:
"""With the parent given, {parent} in a filter is the global-scoped parent."""
text = _squash(
await build_trigger_callback(
[], {**TRIGGER_CONF, "mode": 1}, [(cg.int_, "mode")], forward=[], when=when
[],
TRIGGER_CONF,
[(cg.int_, "state")],
forward=[],
when="state == 1 && {parent}->is_failed() == false",
parent=MockObj("improv", "->"),
)
)
assert (
"if (!(state == 1 && ::improv->is_failed() == false)) return; "
f"::{NEW_OBJ}->trigger();"
) in text
@pytest.mark.asyncio
async def test_build_callback_automations_when_names_parent(
mock_cg: MockCodegen,
) -> None:
"""A table entry names the parent through {parent}, so it can live at module level."""
entries = (
CallbackAutomation(
"on_open", "add_on_state_callback", when="{parent}->is_fully_open()"
),
)
await build_callback_automations(
MockObj("valve", "->"), {"on_open": [TRIGGER_CONF]}, entries
)
assert _squash(mock_cg.add.call_args.args[0]) == (
"valve->add_on_state_callback([]() -> void { "
f"if (!(::valve->is_fully_open())) return; ::{NEW_OBJ}->trigger(); }})"
)
@pytest.mark.asyncio