diff --git a/esphome/automation.py b/esphome/automation.py index a8c421f065..4a67cd31ba 100644 --- a/esphome/automation.py +++ b/esphome/automation.py @@ -339,9 +339,9 @@ def _check_key_in_schema( schema = schema.schema[markers[part]] -async def _apply_parent(config: ConfigType) -> str: +async def _apply_parent(config: ConfigType, id_key: str = CONF_ID) -> str: # Global-scope qualified so a trigger arg named like the id cannot shadow it. - return f"::{await cg.get_variable(config[CONF_ID])}" + return f"::{await cg.get_variable(config[id_key])}" def _apply_lambda_args(args: TemplateArgsType) -> TemplateArgsType: @@ -396,12 +396,14 @@ def register_apply_action( schema: cv.Schema, *fields: ApplyField | ApplyCall, call: str | None = None, + id_key: str = CONF_ID, ) -> None: """Register an action that only forwards config values to its parent, with no C++ class. - Generates one stateless function for ``ApplyAction``: parent and constants are baked - in, lambdas are called inline with the trigger args. With ``call`` every statement targets - the call object ``auto apply_call = parent->call()``, and ``apply_call.perform()`` is appended. + Generates one stateless function for ``ApplyAction``: the parent (read from + ``id_key``) and constants are baked in, lambdas are called inline with the trigger args. + With ``call`` every statement targets the call object ``auto apply_call = parent->call()``, + and ``apply_call.perform()`` is appended. """ # An action stores the value, so a std::string constant stays in flash on ESP8266. statements_spec = [ @@ -414,6 +416,7 @@ def register_apply_action( ) for c in (f if isinstance(f, ApplyCall) else f.call() for f in fields) ] + _check_key_in_schema(name, schema, id_key) for _, members in statements_spec: for conf_key, _, _ in members: _check_key_in_schema(name, schema, conf_key) @@ -424,7 +427,7 @@ def register_apply_action( template_arg: cg.TemplateArguments, args: TemplateArgsType, ) -> MockObj: - parent = await _apply_parent(config) + parent = await _apply_parent(config, id_key) lambda_args = _apply_lambda_args(args) receiver = "apply_call." if call else f"{parent}->" statements: list[str] = [] @@ -451,7 +454,7 @@ def register_apply_action( def register_apply_condition( - name: str, schema: cv.Schema, check: str | ApplyCall + name: str, schema: cv.Schema, check: str | ApplyCall, id_key: str = CONF_ID ) -> None: """Register a condition that is one expression on its parent, with no C++ class. @@ -463,6 +466,7 @@ def register_apply_condition( """ call = check if isinstance(check, ApplyCall) else ApplyCall(check) members = call.members + _check_key_in_schema(name, schema, id_key) for conf_key, _, _ in members: _check_key_in_schema(name, schema, conf_key) @@ -472,7 +476,7 @@ def register_apply_condition( template_arg: cg.TemplateArguments, args: TemplateArgsType, ) -> MockObj: - parent = await _apply_parent(config) + parent = await _apply_parent(config, id_key) lambda_args = _apply_lambda_args(args) exprs = await _render_values( name, diff --git a/tests/unit_tests/test_automation.py b/tests/unit_tests/test_automation.py index f6aaff3d92..1e1d2e714a 100644 --- a/tests/unit_tests/test_automation.py +++ b/tests/unit_tests/test_automation.py @@ -610,12 +610,13 @@ async def _run_entry( config: dict[str, object], args: list[tuple[object, str]] | None, platform: str, + id_key: str = CONF_ID, ) -> RegistryEntry: """Run a registered builder with the given config, trigger args and platform.""" CORE.data[KEY_CORE] = {KEY_TARGET_PLATFORM: platform} args = args or [] template_arg = cg.TemplateArguments(*(t for t, _ in args)) - await entry.fun({CONF_ID: PARENT_ID, **config}, ID("obj_1"), template_arg, args) + await entry.fun({id_key: PARENT_ID, **config}, ID("obj_1"), template_arg, args) return entry @@ -626,11 +627,12 @@ async def _run_apply_action( args: list[tuple[object, str]] | None = None, call: str | None = None, platform: str = "esp32", + id_key: str = CONF_ID, ) -> RegistryEntry: """Register an apply action and run its builder with the given config.""" actions, _ = registries - register_apply_action("my.apply", None, *fields, call=call) - return await _run_entry(actions["my.apply"], config, args, platform) + register_apply_action("my.apply", None, *fields, call=call, id_key=id_key) + return await _run_entry(actions["my.apply"], config, args, platform, id_key) async def _run_apply_condition( @@ -639,11 +641,12 @@ async def _run_apply_condition( config: dict[str, object], args: list[tuple[object, str]] | None = None, platform: str = "esp32", + id_key: str = CONF_ID, ) -> RegistryEntry: """Register an apply condition and run its builder with the given config.""" _, conditions = registries - register_apply_condition("my.check", None, check) - return await _run_entry(conditions["my.check"], config, args, platform) + register_apply_condition("my.check", None, check, id_key=id_key) + return await _run_entry(conditions["my.check"], config, args, platform, id_key) def _apply_lambda(mock_cg: MockCodegen) -> str: @@ -663,6 +666,17 @@ async def test_register_apply_action_entry( assert str(template_arg) == "" +@pytest.mark.asyncio +async def test_apply_custom_id_key( + registries: tuple[Registry, Registry], mock_cg: MockCodegen +) -> None: + await _run_apply_action(registries, (), {}, id_key="transmitter_id") + mock_cg.get_variable.assert_awaited_once_with(PARENT_ID) + mock_cg.get_variable.reset_mock() + await _run_apply_condition(registries, "is_on()", {}, id_key="transmitter_id") + mock_cg.get_variable.assert_awaited_once_with(PARENT_ID) + + @pytest.mark.asyncio async def test_apply_constants( registries: tuple[Registry, Registry], mock_cg: MockCodegen @@ -806,6 +820,10 @@ def test_apply_registration_checks(registries: tuple[Registry, Registry]) -> Non register_apply_condition( "my.bad_is", schema, ApplyCall("kd == {}", (("kd", cg.float_),)) ) + with pytest.raises(ValueError, match="'parent_id' is not in the schema"): + register_apply_action("my.bad_id", schema, id_key="parent_id") + with pytest.raises(ValueError, match="'parent_id' is not in the schema"): + register_apply_condition("my.bad_is_id", schema, "is_on()", id_key="parent_id") either = cv.Any(schema, cv.Schema({cv.Optional("kd"): cv.float_})) register_apply_action("my.any", either, ApplyField("kd", "set_kd", cg.float_)) for wrapped in ( @@ -818,7 +836,12 @@ def test_apply_registration_checks(registries: tuple[Registry, Registry]) -> Non register_apply_action( "my.bad", wrapped, ApplyField("kd", "set_kd", cg.float_) ) - nested = cv.Schema({cv.Optional("v"): cv.Schema({cv.Optional("dir"): cv.int_})}) + nested = cv.Schema( + { + cv.Required(CONF_ID): cv.string, + cv.Optional("v"): cv.Schema({cv.Optional("dir"): cv.int_}), + } + ) register_apply_action( "my.nested", nested, ApplyField(("v", "dir"), "set_dir", cg.int_) )