diff --git a/esphome/automation.py b/esphome/automation.py index 859ec2f1f3..6b6af4ed72 100644 --- a/esphome/automation.py +++ b/esphome/automation.py @@ -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 diff --git a/tests/unit_tests/test_automation.py b/tests/unit_tests/test_automation.py index ddd777fdb0..3cd70c082d 100644 --- a/tests/unit_tests/test_automation.py +++ b/tests/unit_tests/test_automation.py @@ -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