[core] Let register_apply_action read the parent id from another schema key (#19548)

This commit is contained in:
J. Nick Koston
2026-09-24 09:12:04 -04:00
committed by GitHub
parent fed5dd6a9d
commit e92b1e27ca
2 changed files with 41 additions and 14 deletions
+12 -8
View File
@@ -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<Ts...>``: 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<Ts...>``: 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,
+29 -6
View File
@@ -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) == "<int32_t>"
@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_)
)