mirror of
https://github.com/esphome/esphome.git
synced 2026-09-28 07:20:24 +00:00
[core] Make the generated apply function a template argument of ApplyAction and ApplyCondition (#19655)
This commit is contained in:
+26
-13
@@ -345,14 +345,28 @@ async def _apply_parent(config: ConfigType, id_key: str = CONF_ID) -> str:
|
||||
|
||||
|
||||
def _apply_lambda_args(args: TemplateArgsType) -> TemplateArgsType:
|
||||
# Must match ApplyAction::ApplyFn and ApplyCondition::CheckFn exactly for the function
|
||||
# pointer conversion.
|
||||
# The generated function's parameters; a std::string arg is never copied.
|
||||
return [
|
||||
(cg.RawExpression(f"const std::remove_cvref_t<{cg.safe_exp(t)}> &"), arg)
|
||||
for t, arg in args
|
||||
]
|
||||
|
||||
|
||||
def _apply_function(
|
||||
id_: ID,
|
||||
return_type: SafeExpType,
|
||||
template_arg: cg.TemplateArguments,
|
||||
lambda_args: TemplateArgsType,
|
||||
statements: list[str],
|
||||
) -> MockObj:
|
||||
"""Emit the generated function and declare ``id_`` as the ``ApplyAction`` or
|
||||
``ApplyCondition`` templated on it, so ``play()`` calls it directly."""
|
||||
fn = cg.static_function(
|
||||
f"esphome__{id_.id}__fn", return_type, lambda_args, statements
|
||||
)
|
||||
return cg.new_Pvariable(id_, cg.TemplateArguments(fn, *template_arg))
|
||||
|
||||
|
||||
async def _render_values(
|
||||
name: str,
|
||||
target: str,
|
||||
@@ -403,8 +417,8 @@ def register_apply_action(
|
||||
) -> None:
|
||||
"""Register an action that only forwards config values to its parent, with no C++ class.
|
||||
|
||||
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.
|
||||
Generates one static function with the parent (read from ``id_key``) and constants baked
|
||||
in, lambdas called inline with the trigger args, and an ``ApplyAction`` templated on it.
|
||||
A constant that is an id (``cv.use_id`` under ``cv.templatable``) is the object it names.
|
||||
With ``call`` every statement targets the call object ``auto apply_call = parent->call()``,
|
||||
and ``apply_call.perform()`` is appended.
|
||||
@@ -449,10 +463,9 @@ def register_apply_action(
|
||||
*statements,
|
||||
"apply_call.perform();",
|
||||
]
|
||||
apply_lambda = LambdaExpression(
|
||||
["\n".join(statements)], lambda_args, capture="", return_type=cg.void
|
||||
return _apply_function(
|
||||
action_id, cg.void, template_arg, lambda_args, statements
|
||||
)
|
||||
return cg.new_Pvariable(action_id, template_arg, apply_lambda)
|
||||
|
||||
register_action(name, ApplyAction, schema, synchronous=True)(builder)
|
||||
|
||||
@@ -466,7 +479,7 @@ def register_apply_condition(
|
||||
``ApplyCall`` such as ``ApplyCall("state == {}", ((CONF_STATE, cg.bool_),))`` compares
|
||||
against config values, all of which must be present. Write ``== false`` to negate.
|
||||
String constants are plain literals, so compare a ``std::string`` or ``StringRef`` member.
|
||||
Generates one stateless function for ``ApplyCondition<Ts...>``.
|
||||
Generates one static predicate and an ``ApplyCondition`` templated on it.
|
||||
"""
|
||||
call = check if isinstance(check, ApplyCall) else ApplyCall(check)
|
||||
members = call.members
|
||||
@@ -492,13 +505,13 @@ def register_apply_condition(
|
||||
lambda_args,
|
||||
compare=True,
|
||||
)
|
||||
check_lambda = LambdaExpression(
|
||||
[f"return {parent}->{call.target.format(*exprs)};"],
|
||||
return _apply_function(
|
||||
condition_id,
|
||||
cg.bool_,
|
||||
template_arg,
|
||||
lambda_args,
|
||||
capture="",
|
||||
return_type=cg.bool_,
|
||||
[f"return {parent}->{call.target.format(*exprs)};"],
|
||||
)
|
||||
return cg.new_Pvariable(condition_id, template_arg, check_lambda)
|
||||
|
||||
register_condition(name, ApplyCondition, schema)(builder)
|
||||
|
||||
|
||||
@@ -42,6 +42,7 @@ from esphome.cpp_generator import ( # noqa: F401
|
||||
set_cpp_standard,
|
||||
statement,
|
||||
static_const_array,
|
||||
static_function,
|
||||
templatable,
|
||||
variable,
|
||||
with_local_variable,
|
||||
|
||||
@@ -253,30 +253,19 @@ template<typename... Ts> class StatelessLambdaAction : public Action<Ts...> {
|
||||
void (*f_)(Ts...);
|
||||
};
|
||||
|
||||
/// Runs one codegen-generated function that has the parent and every field baked in, so the
|
||||
/// action holds one pointer. Args pass by const reference so a std::string arg is never copied;
|
||||
/// Runs one codegen-generated function that has the parent and every field baked in. The
|
||||
/// function is a template argument, so play() calls it directly and the object is just the
|
||||
/// Action base. Args pass by const reference so a std::string arg is never copied;
|
||||
/// StatelessLambdaAction keeps by-value parameters because user `lambda:` code owns them.
|
||||
template<typename... Ts> class ApplyAction final : public Action<Ts...> {
|
||||
template<auto Fn, typename... Ts> class ApplyAction final : public Action<Ts...> {
|
||||
public:
|
||||
using ApplyFn = void (*)(const std::remove_cvref_t<Ts> &...);
|
||||
explicit ApplyAction(ApplyFn apply) : apply_(apply) {}
|
||||
|
||||
void play(const Ts &...x) override { this->apply_(x...); }
|
||||
|
||||
protected:
|
||||
ApplyFn apply_;
|
||||
void play(const Ts &...x) override { Fn(x...); }
|
||||
};
|
||||
|
||||
/// Condition counterpart of ApplyAction: one codegen-generated predicate with the parent baked in.
|
||||
template<typename... Ts> class ApplyCondition final : public Condition<Ts...> {
|
||||
template<auto Fn, typename... Ts> class ApplyCondition final : public Condition<Ts...> {
|
||||
public:
|
||||
using CheckFn = bool (*)(const std::remove_cvref_t<Ts> &...);
|
||||
explicit ApplyCondition(CheckFn check) : check_(check) {}
|
||||
|
||||
bool check(const Ts &...x) override { return this->check_(x...); }
|
||||
|
||||
protected:
|
||||
CheckFn check_;
|
||||
bool check(const Ts &...x) override { return Fn(x...); }
|
||||
};
|
||||
|
||||
/// Simple continuation action that calls play_next_ on a parent action.
|
||||
|
||||
@@ -19,7 +19,7 @@ from esphome.core import (
|
||||
TimePeriodNanoseconds,
|
||||
TimePeriodSeconds,
|
||||
)
|
||||
from esphome.helpers import cpp_string_escape, indent_all_but_first_and_last
|
||||
from esphome.helpers import cpp_string_escape, indent, indent_all_but_first_and_last
|
||||
from esphome.types import Expression, SafeExpType, TemplateArgsType
|
||||
from esphome.util import OrderedDict
|
||||
from esphome.yaml_util import ESPHomeDataBase
|
||||
@@ -670,6 +670,28 @@ def new_Pvariable(id_: ID, *args: SafeExpType) -> "MockObj":
|
||||
return Pvariable(id_, rhs)
|
||||
|
||||
|
||||
def static_function(
|
||||
name: str,
|
||||
return_type: SafeExpType,
|
||||
parameters: TemplateArgsType,
|
||||
body: list[str],
|
||||
) -> RawExpression:
|
||||
"""Emit ``static <return_type> <name>(parameters) { body }`` at global scope and return an
|
||||
expression naming it, for use as a template argument or a function pointer.
|
||||
|
||||
Every id the body names must already be declared, which holds when the statements were
|
||||
rendered through ``get_variable`` or ``process_lambda``.
|
||||
"""
|
||||
params = ParameterListExpression(*parameters)
|
||||
add_global(
|
||||
RawStatement(
|
||||
f"static {safe_exp(return_type)} {name}({params}) {{\n"
|
||||
f"{indent(chr(10).join(body))}\n}}"
|
||||
)
|
||||
)
|
||||
return RawExpression(name)
|
||||
|
||||
|
||||
def add(expression: Expression | Statement, prepend: bool = False):
|
||||
"""Add an expression to the codegen section.
|
||||
|
||||
|
||||
@@ -53,8 +53,8 @@ esphome:
|
||||
# Regression test for esphome/esphome#16224: a light.turn_on action
|
||||
# nested inside on_response with capture_response: true puts
|
||||
# `std::string &` into the trigger's Ts..., which exposed a codegen
|
||||
# bug where the apply lambda's parameter list did not match the
|
||||
# ApplyFn signature.
|
||||
# bug where the generated function's parameter list did not match
|
||||
# the automation's arguments.
|
||||
- light.turn_on:
|
||||
id: test_regression_light
|
||||
brightness: 100%
|
||||
|
||||
@@ -509,6 +509,8 @@ class MockCodegen(NamedTuple):
|
||||
get_variable: AsyncMock
|
||||
new_pvariable: MagicMock
|
||||
register_parented: AsyncMock
|
||||
add_global: MagicMock
|
||||
calls: MagicMock # new_pvariable and add_global attached, to check their order
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -520,10 +522,16 @@ def mock_cg() -> Generator[MockCodegen]:
|
||||
patch(
|
||||
"esphome.codegen.register_parented", new_callable=AsyncMock
|
||||
) as register_parented,
|
||||
patch("esphome.cpp_generator.add_global") as add_global,
|
||||
):
|
||||
get_variable.return_value = PARENT_OBJ
|
||||
new_pvariable.return_value = NEW_OBJ
|
||||
yield MockCodegen(get_variable, new_pvariable, register_parented)
|
||||
calls = MagicMock()
|
||||
calls.attach_mock(new_pvariable, "new_pvariable")
|
||||
calls.attach_mock(add_global, "add_global")
|
||||
yield MockCodegen(
|
||||
get_variable, new_pvariable, register_parented, add_global, calls
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -649,8 +657,8 @@ async def _run_apply_condition(
|
||||
return await _run_entry(conditions["my.check"], config, args, platform, id_key)
|
||||
|
||||
|
||||
def _apply_lambda(mock_cg: MockCodegen) -> str:
|
||||
return str(mock_cg.new_pvariable.call_args.args[2])
|
||||
def _apply_definition(mock_cg: MockCodegen) -> str:
|
||||
return str(mock_cg.add_global.call_args.args[0])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -661,9 +669,14 @@ async def test_register_apply_action_entry(
|
||||
assert entry.type_id is ApplyAction
|
||||
assert entry.synchronous is True
|
||||
mock_cg.get_variable.assert_awaited_once_with(PARENT_ID)
|
||||
action_id, template_arg, _ = mock_cg.new_pvariable.call_args.args
|
||||
action_id, template_arg = mock_cg.new_pvariable.call_args.args
|
||||
assert action_id == ID("obj_1")
|
||||
assert str(template_arg) == "<int32_t>"
|
||||
assert str(template_arg) == "<esphome__obj_1__fn, int32_t>"
|
||||
# The definition must precede the storage line that names the function.
|
||||
assert [c[0] for c in mock_cg.calls.mock_calls] == ["add_global", "new_pvariable"]
|
||||
assert _apply_definition(mock_cg).startswith(
|
||||
"static void esphome__obj_1__fn(const std::remove_cvref_t<int32_t> & x) {"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -692,7 +705,7 @@ async def test_apply_constants(
|
||||
)
|
||||
config = {"kp": 0.0, "on": False, "song": "a:b", "position": 0.5}
|
||||
await _run_apply_action(registries, fields, config)
|
||||
text = _apply_lambda(mock_cg)
|
||||
text = _apply_definition(mock_cg)
|
||||
lines = [
|
||||
f"::{PARENT_OBJ}->set_kp(0.0f);",
|
||||
f"::{PARENT_OBJ}->set_on(false);",
|
||||
@@ -715,7 +728,9 @@ async def test_apply_id_constant_is_the_named_object(
|
||||
fields = (ApplyField("target", "switch_to_output", cg.RawExpression("Speaker *")),)
|
||||
await _run_apply_action(registries, fields, {"target": ID("speaker_b")})
|
||||
mock_cg.get_variable.assert_any_await(ID("speaker_b"))
|
||||
assert f"::{PARENT_OBJ}->switch_to_output(::speaker_b);" in _apply_lambda(mock_cg)
|
||||
assert f"::{PARENT_OBJ}->switch_to_output(::speaker_b);" in str(
|
||||
_apply_definition(mock_cg)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -728,7 +743,9 @@ async def test_apply_condition_id_constant_is_the_named_object(
|
||||
check = ApplyCall("is_output({})", (("target", cg.RawExpression("Speaker *")),))
|
||||
await _run_apply_condition(registries, check, {"target": ID("speaker_b")})
|
||||
mock_cg.get_variable.assert_any_await(ID("speaker_b"))
|
||||
assert f"return ::{PARENT_OBJ}->is_output(::speaker_b);" in _apply_lambda(mock_cg)
|
||||
assert f"return ::{PARENT_OBJ}->is_output(::speaker_b);" in str(
|
||||
_apply_definition(mock_cg)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -745,8 +762,10 @@ async def test_apply_lambdas(
|
||||
"ki": Lambda("if (x) return 1.0f;\nreturn 2.0f;"),
|
||||
}
|
||||
await _run_apply_action(registries, fields, config, args=[(cg.int32, "x")])
|
||||
text = _apply_lambda(mock_cg)
|
||||
assert text.startswith("[](const std::remove_cvref_t<int32_t> & x) -> void {")
|
||||
text = _apply_definition(mock_cg)
|
||||
assert text.startswith(
|
||||
"static void esphome__obj_1__fn(const std::remove_cvref_t<int32_t> & x) {"
|
||||
)
|
||||
# The parent is global-scope qualified, so an arg named like the id cannot shadow it.
|
||||
assert f"::{PARENT_OBJ}->set_kp(" in text
|
||||
assert f"::{PARENT_OBJ}->set_kp(static_cast<float>(x * 2));" in text
|
||||
@@ -768,11 +787,11 @@ async def test_apply_call_keys(
|
||||
ApplyCall("set_range({}, {})", (("low", cg.float_), ("high", cg.float_))),
|
||||
)
|
||||
await _run_apply_action(registries, fields, {"low": 1.0, "high": 2.0})
|
||||
assert f"::{PARENT_OBJ}->set_range(1.0f, 2.0f);" in _apply_lambda(mock_cg)
|
||||
assert f"::{PARENT_OBJ}->set_range(1.0f, 2.0f);" in _apply_definition(mock_cg)
|
||||
|
||||
mock_cg.new_pvariable.reset_mock()
|
||||
await _run_apply_action(registries, fields, {})
|
||||
assert "set_range" not in _apply_lambda(mock_cg)
|
||||
assert "set_range" not in _apply_definition(mock_cg)
|
||||
|
||||
with pytest.raises(EsphomeError, match="needs all of"):
|
||||
await _run_apply_action(registries, fields, {"low": 1.0})
|
||||
@@ -784,7 +803,7 @@ async def test_apply_action_call_shape(
|
||||
) -> None:
|
||||
fields = (ApplyField("brightness", "set_brightness", cg.float_),)
|
||||
await _run_apply_action(registries, fields, {"brightness": 0.5}, call="make_call")
|
||||
text = _apply_lambda(mock_cg)
|
||||
text = _apply_definition(mock_cg)
|
||||
lines = [
|
||||
f"auto apply_call = ::{PARENT_OBJ}->make_call();",
|
||||
"apply_call.set_brightness(0.5f);",
|
||||
@@ -814,7 +833,7 @@ async def test_apply_field_nested_key_const_fn_and_type_string(
|
||||
"value": Lambda("return 42;"),
|
||||
}
|
||||
await _run_apply_action(registries, fields, config)
|
||||
text = _apply_lambda(mock_cg)
|
||||
text = _apply_definition(mock_cg)
|
||||
assert f"::{PARENT_OBJ}->set_direction(3);" in text
|
||||
assert f'::{PARENT_OBJ}->set_name("abc", 3);' in text
|
||||
assert (
|
||||
@@ -824,7 +843,7 @@ async def test_apply_field_nested_key_const_fn_and_type_string(
|
||||
|
||||
mock_cg.new_pvariable.reset_mock()
|
||||
await _run_apply_action(registries, fields[:1], {})
|
||||
assert "set_direction" not in _apply_lambda(mock_cg)
|
||||
assert "set_direction" not in _apply_definition(mock_cg)
|
||||
|
||||
|
||||
def test_apply_registration_checks(registries: tuple[Registry, Registry]) -> None:
|
||||
@@ -883,8 +902,9 @@ async def test_apply_string_constant_stays_in_flash_on_esp8266(
|
||||
) -> None:
|
||||
fields = (ApplyField("song", "play", cg.std_string),)
|
||||
await _run_apply_action(registries, fields, {"song": "a:b"}, platform="esp8266")
|
||||
assert f'::{PARENT_OBJ}->play(progmem_string(ESPHOME_F("a:b")));' in _apply_lambda(
|
||||
mock_cg
|
||||
assert (
|
||||
f'::{PARENT_OBJ}->play(progmem_string(ESPHOME_F("a:b")));'
|
||||
in _apply_definition(mock_cg)
|
||||
)
|
||||
|
||||
|
||||
@@ -900,7 +920,7 @@ async def test_apply_literal_with_length_is_plain_on_every_platform(
|
||||
await _run_apply_action(
|
||||
registries, fields, {"option": "h\u00e9llo"}, platform=platform
|
||||
)
|
||||
text = _apply_lambda(mock_cg)
|
||||
text = _apply_definition(mock_cg)
|
||||
assert f'::{PARENT_OBJ}->set_option("h\\303\\251llo", 6);' in text
|
||||
assert "progmem_string" not in text
|
||||
|
||||
@@ -913,11 +933,13 @@ async def test_register_apply_condition_predicate(
|
||||
registries, "is_playing()", {}, args=[(cg.int32, "x")]
|
||||
)
|
||||
assert entry.type_id is ApplyCondition
|
||||
condition_id, template_arg, check = mock_cg.new_pvariable.call_args.args
|
||||
condition_id, template_arg = mock_cg.new_pvariable.call_args.args
|
||||
assert condition_id == ID("obj_1")
|
||||
assert str(template_arg) == "<int32_t>"
|
||||
text = str(check)
|
||||
assert text.startswith("[](const std::remove_cvref_t<int32_t> & x) -> bool {")
|
||||
assert str(template_arg) == "<esphome__obj_1__fn, int32_t>"
|
||||
text = _apply_definition(mock_cg)
|
||||
assert text.startswith(
|
||||
"static bool esphome__obj_1__fn(const std::remove_cvref_t<int32_t> & x) {"
|
||||
)
|
||||
assert f"return ::{PARENT_OBJ}->is_playing();" in text
|
||||
|
||||
|
||||
@@ -927,7 +949,7 @@ async def test_apply_condition_compares_config_value(
|
||||
) -> None:
|
||||
check = ApplyCall("state == {}", (("state", cg.bool_),))
|
||||
await _run_apply_condition(registries, check, {"state": True})
|
||||
assert f"return ::{PARENT_OBJ}->state == true;" in _apply_lambda(mock_cg)
|
||||
assert f"return ::{PARENT_OBJ}->state == true;" in _apply_definition(mock_cg)
|
||||
|
||||
with pytest.raises(EsphomeError, match="needs all of"):
|
||||
await _run_apply_condition(registries, check, {})
|
||||
@@ -940,7 +962,7 @@ async def test_apply_condition_string_constant_is_a_plain_literal(
|
||||
) -> None:
|
||||
check = ApplyCall("state == {}", (("state", cg.std_string),))
|
||||
await _run_apply_condition(registries, check, {"state": "two"}, platform=platform)
|
||||
assert f'return ::{PARENT_OBJ}->state == "two";' in _apply_lambda(mock_cg)
|
||||
assert f'return ::{PARENT_OBJ}->state == "two";' in _apply_definition(mock_cg)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -964,6 +986,6 @@ async def test_apply_condition_string_lambda_paths(
|
||||
await _run_apply_condition(
|
||||
registries, check, {"state": Lambda(body)}, args=[(cg.std_string, "x")]
|
||||
)
|
||||
text = _apply_lambda(mock_cg)
|
||||
text = _apply_definition(mock_cg)
|
||||
assert expected in text
|
||||
assert ("-> std::string {" in text) is called
|
||||
|
||||
Reference in New Issue
Block a user