mirror of
https://github.com/esphome/esphome.git
synced 2026-10-09 21:13:12 +00:00
[core] Build trigger callbacks as stateless lambdas (#19606)
This commit is contained in:
+154
-30
@@ -1,4 +1,4 @@
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Sequence
|
||||
from dataclasses import dataclass, field
|
||||
import logging
|
||||
import string
|
||||
@@ -23,7 +23,7 @@ from esphome.const import (
|
||||
)
|
||||
from esphome.core import CORE, ID, EsphomeError, HexInt, Lambda
|
||||
from esphome.cpp_generator import (
|
||||
FlashStringLiteral,
|
||||
Expression,
|
||||
LambdaExpression,
|
||||
MockObj,
|
||||
MockObjClass,
|
||||
@@ -245,9 +245,7 @@ ApplyCondition = cg.esphome_ns.class_("ApplyCondition", Condition)
|
||||
|
||||
def flash_string(config: ConfigType, value: str) -> str:
|
||||
"""Default renderer for ``std::string`` constants; copies the literal out of flash on ESP8266."""
|
||||
if CORE.is_esp8266:
|
||||
return f"progmem_string({FlashStringLiteral(value)})"
|
||||
return str(cg.safe_exp(value))
|
||||
return str(cg.progmem_string(value))
|
||||
|
||||
|
||||
def literal_with_length(config: ConfigType, value: str) -> str:
|
||||
@@ -259,6 +257,13 @@ def literal_with_length(config: ConfigType, value: str) -> str:
|
||||
return f"{cg.safe_exp(value)}, {len(value.encode('utf-8'))}"
|
||||
|
||||
|
||||
def string_ref_literal(config: ConfigType, value: str) -> str:
|
||||
"""Renderer for a ``StringRef`` comparison: a flash literal on ESP8266, else ``StringRef(literal, length)``."""
|
||||
if CORE.is_esp8266:
|
||||
return str(cg.FlashStringLiteral(value))
|
||||
return f"StringRef({literal_with_length(config, value)})"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ApplyCall:
|
||||
"""One statement from config keys, e.g. ``"set_range({}, {})"`` with ``((CONF_LOW, cg.float_), ...)``.
|
||||
@@ -365,9 +370,16 @@ def _check_key_in_schema(
|
||||
schema = schema.schema[markers[part]]
|
||||
|
||||
|
||||
def parent_ref(var: MockObj) -> MockObj:
|
||||
"""``var`` named from global scope, so a trigger argument cannot shadow it.
|
||||
|
||||
Also how a generated callback names its Automation.
|
||||
"""
|
||||
return MockObj(f"::{var}", "->")
|
||||
|
||||
|
||||
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[id_key])}"
|
||||
return str(parent_ref(await cg.get_variable(config[id_key])))
|
||||
|
||||
|
||||
def _apply_lambda_args(args: TemplateArgsType) -> TemplateArgsType:
|
||||
@@ -399,7 +411,7 @@ async def _render_values(
|
||||
members: list[tuple[Any, Any, Any]],
|
||||
values: list[Any],
|
||||
config: ConfigType,
|
||||
parent: str,
|
||||
parent: str | None,
|
||||
lambda_args: TemplateArgsType,
|
||||
compare: bool = False,
|
||||
) -> list[str]:
|
||||
@@ -477,7 +489,7 @@ def register_apply_action(
|
||||
statements: list[str] = []
|
||||
for target, members in statements_spec:
|
||||
values = _apply_values(config, members)
|
||||
if members and all(value is None for value in values):
|
||||
if not _apply_call_active(members, values):
|
||||
continue
|
||||
exprs = await _render_values(
|
||||
name, target, members, values, config, parent, lambda_args
|
||||
@@ -496,6 +508,27 @@ def register_apply_action(
|
||||
register_action(name, ApplyAction, schema, synchronous=True)(builder)
|
||||
|
||||
|
||||
def _apply_call_active(members: list[tuple[Any, Any, Any]], values: list[Any]) -> bool:
|
||||
"""An ``ApplyCall`` is emitted unless it has keys and none of them is set."""
|
||||
return not members or any(value is not None for value in values)
|
||||
|
||||
|
||||
async def _render_check(
|
||||
name: str,
|
||||
target: str,
|
||||
members: list[tuple[Any, Any, Any]],
|
||||
values: list[Any],
|
||||
config: ConfigType,
|
||||
parent: str | None,
|
||||
lambda_args: TemplateArgsType,
|
||||
) -> str:
|
||||
"""Render a boolean check with ``values`` compared against config."""
|
||||
exprs = await _render_values(
|
||||
name, target, members, values, config, parent, lambda_args, compare=True
|
||||
)
|
||||
return target.format(*exprs)
|
||||
|
||||
|
||||
def register_apply_condition(
|
||||
name: str, schema: cv.Schema, check: str | ApplyCall, id_key: str = CONF_ID
|
||||
) -> None:
|
||||
@@ -521,22 +554,16 @@ def register_apply_condition(
|
||||
) -> MockObj:
|
||||
parent = await _apply_parent(config, id_key)
|
||||
lambda_args = _apply_lambda_args(args)
|
||||
exprs = await _render_values(
|
||||
name,
|
||||
call.target,
|
||||
members,
|
||||
_apply_values(config, members),
|
||||
config,
|
||||
parent,
|
||||
lambda_args,
|
||||
compare=True,
|
||||
values = _apply_values(config, members)
|
||||
check = await _render_check(
|
||||
name, call.target, members, values, config, parent, lambda_args
|
||||
)
|
||||
return _apply_function(
|
||||
condition_id,
|
||||
cg.bool_,
|
||||
template_arg,
|
||||
lambda_args,
|
||||
[f"return {parent}->{call.target.format(*exprs)};"],
|
||||
[f"return {parent}->{check};"],
|
||||
)
|
||||
|
||||
register_condition(name, ApplyCondition, schema)(builder)
|
||||
@@ -1056,23 +1083,77 @@ def has_non_synchronous_actions(actions: ConfigType) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
async def build_automation(
|
||||
trigger: MockObj, args: TemplateArgsType, config: ConfigType
|
||||
async def _new_automation(
|
||||
args: TemplateArgsType, config: ConfigType, *ctor_args: MockObj
|
||||
) -> MockObj:
|
||||
arg_types = [arg[0] for arg in args]
|
||||
templ = cg.TemplateArguments(*arg_types)
|
||||
obj = cg.new_Pvariable(config[CONF_AUTOMATION_ID], templ, trigger)
|
||||
"""Create the Automation for ``config`` with its actions."""
|
||||
templ = cg.TemplateArguments(*(arg[0] for arg in args))
|
||||
obj = cg.new_Pvariable(config[CONF_AUTOMATION_ID], templ, *ctor_args)
|
||||
actions = await build_action_list(config[CONF_THEN], templ, args)
|
||||
cg.add(obj.add_actions(actions))
|
||||
return obj
|
||||
|
||||
|
||||
async def build_automation(
|
||||
trigger: MockObj, args: TemplateArgsType, config: ConfigType
|
||||
) -> MockObj:
|
||||
return await _new_automation(args, config, trigger)
|
||||
|
||||
|
||||
async def build_trigger_callback(
|
||||
args: TemplateArgsType,
|
||||
config: ConfigType,
|
||||
params: TemplateArgsType,
|
||||
forward: Sequence[str | Expression] | None = None,
|
||||
when: str | ApplyCall | 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.
|
||||
"""
|
||||
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}}")
|
||||
obj = await _new_automation(args, config)
|
||||
lambda_args = _apply_lambda_args(params)
|
||||
statements: list[str] = []
|
||||
if when is not None:
|
||||
values = _apply_values(config, members)
|
||||
if _apply_call_active(members, values):
|
||||
check = await _render_check(
|
||||
"trigger filter",
|
||||
call.target,
|
||||
members,
|
||||
values,
|
||||
config,
|
||||
None,
|
||||
lambda_args,
|
||||
)
|
||||
statements.append(f"if (!({check}))\n return;")
|
||||
if forward is None:
|
||||
forward = [name for _, name in params]
|
||||
statements.append(f"{parent_ref(obj)}->trigger({', '.join(map(str, forward))});")
|
||||
return LambdaExpression(
|
||||
["\n".join(statements)], lambda_args, capture="", return_type=cg.void
|
||||
)
|
||||
|
||||
|
||||
async def build_callback_automation(
|
||||
parent: MockObj,
|
||||
callback_method: str,
|
||||
args: TemplateArgsType,
|
||||
config: ConfigType,
|
||||
forwarder: MockObj | MockObjClass | None = None,
|
||||
params: TemplateArgsType | None = None,
|
||||
forward: Sequence[str | Expression] | None = None,
|
||||
when: str | ApplyCall | None = None,
|
||||
) -> None:
|
||||
"""Build an Automation and register it as a callback on the parent.
|
||||
|
||||
@@ -1084,6 +1165,9 @@ async def build_callback_automation(
|
||||
pointer-sized (single Automation* field) to fit inline in Callback::ctx_
|
||||
and avoid heap allocation.
|
||||
|
||||
With ``params``, ``forward`` or ``when`` the callback is instead the stateless
|
||||
lambda of ``build_trigger_callback``; ``forwarder`` cannot be combined with them.
|
||||
|
||||
:param parent: The component object (e.g., button, sensor).
|
||||
:param callback_method: Name of the callback method (e.g., "add_on_press_callback").
|
||||
:param args: Automation template args as list of (type, name) tuples.
|
||||
@@ -1092,22 +1176,56 @@ async def build_callback_automation(
|
||||
TriggerForwarder<Ts...>. Pass any struct type whose aggregate init takes
|
||||
a single Automation pointer (e.g., TriggerOnTrueForwarder).
|
||||
"""
|
||||
arg_types = [arg[0] for arg in args]
|
||||
templ = cg.TemplateArguments(*arg_types)
|
||||
obj = cg.new_Pvariable(config[CONF_AUTOMATION_ID], templ)
|
||||
actions = await build_action_list(config[CONF_THEN], templ, args)
|
||||
cg.add(obj.add_actions(actions))
|
||||
if params is not None or forward is not None or when is not None:
|
||||
if forwarder is not None:
|
||||
raise ValueError(
|
||||
"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
|
||||
)
|
||||
cg.add(getattr(parent, callback_method)(callback))
|
||||
return
|
||||
obj = await _new_automation(args, config)
|
||||
# Use template forwarder structs for deduplication. The compiler generates
|
||||
# one operator() per forwarder type; different automation pointers are just
|
||||
# data in the struct.
|
||||
if forwarder is None:
|
||||
forwarder = TriggerForwarder.template(templ)
|
||||
forwarder = TriggerForwarder.template(*(arg[0] for arg in args))
|
||||
# RawExpression for aggregate init — both forwarder and obj are codegen
|
||||
# MockObjs (not user input), and there's no Expression type for positional
|
||||
# aggregate initialization (StructInitializer uses named fields).
|
||||
cg.add(getattr(parent, callback_method)(cg.RawExpression(f"{forwarder}{{{obj}}}")))
|
||||
|
||||
|
||||
async def build_parent_callback_automation(
|
||||
parent: MockObj, callback_method: str, arg: tuple[Any, str], config: ConfigType
|
||||
) -> None:
|
||||
"""Register an Automation that receives ``parent`` from a callback that carries nothing.
|
||||
|
||||
``arg`` is the automation's ``(type, name)``, e.g. ``(Fan.operator("ptr"), "x")``.
|
||||
"""
|
||||
await build_callback_automation(
|
||||
parent, callback_method, [arg], config, params=[], forward=[parent_ref(parent)]
|
||||
)
|
||||
|
||||
|
||||
async def build_trigger_automations(
|
||||
parent: MockObj | None,
|
||||
config: ConfigType,
|
||||
entries: tuple[tuple[str, TemplateArgsType], ...],
|
||||
) -> None:
|
||||
"""Instantiate each entry's Trigger class, with ``parent`` when given, and build its automations.
|
||||
|
||||
``entries`` are ``(conf_key, args)`` pairs; the class comes from the entry's ``CONF_TRIGGER_ID``.
|
||||
"""
|
||||
ctor_args = () if parent is None else (parent,)
|
||||
for conf_key, args in entries:
|
||||
for conf in config.get(conf_key, []):
|
||||
trigger = cg.new_Pvariable(conf[CONF_TRIGGER_ID], *ctor_args)
|
||||
await build_automation(trigger, args, conf)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CallbackAutomation:
|
||||
"""A single callback automation entry for build_callback_automations."""
|
||||
@@ -1116,6 +1234,9 @@ class CallbackAutomation:
|
||||
callback_method: str
|
||||
args: TemplateArgsType = field(default_factory=list)
|
||||
forwarder: MockObj | MockObjClass | None = None
|
||||
params: TemplateArgsType | None = None
|
||||
forward: Sequence[str | Expression] | None = None
|
||||
when: str | ApplyCall | None = None
|
||||
|
||||
|
||||
async def build_callback_automations(
|
||||
@@ -1137,4 +1258,7 @@ async def build_callback_automations(
|
||||
entry.args,
|
||||
conf,
|
||||
forwarder=entry.forwarder,
|
||||
params=entry.params,
|
||||
forward=entry.forward,
|
||||
when=entry.when,
|
||||
)
|
||||
|
||||
@@ -39,6 +39,7 @@ from esphome.cpp_generator import ( # noqa: F401
|
||||
new_variable,
|
||||
process_lambda,
|
||||
progmem_array,
|
||||
progmem_string,
|
||||
safe_exp,
|
||||
set_cpp_standard,
|
||||
shared_progmem_array,
|
||||
|
||||
@@ -269,6 +269,13 @@ class FlashStringLiteral(Literal):
|
||||
return f"ESPHOME_F({cpp_string_escape(self.string)})"
|
||||
|
||||
|
||||
def progmem_string(value: str) -> Expression:
|
||||
"""A ``std::string`` argument from a literal that stays in flash on ESP8266."""
|
||||
if CORE.is_esp8266:
|
||||
return RawExpression(f"progmem_string({FlashStringLiteral(value)})")
|
||||
return safe_exp(value)
|
||||
|
||||
|
||||
class IntLiteral(Literal):
|
||||
__slots__ = ("i",)
|
||||
|
||||
|
||||
@@ -16,10 +16,15 @@ from esphome.automation import (
|
||||
TriggerForwarder,
|
||||
TriggerOnFalseForwarder,
|
||||
TriggerOnTrueForwarder,
|
||||
build_callback_automation,
|
||||
build_callback_automations,
|
||||
build_parent_callback_automation,
|
||||
build_trigger_automations,
|
||||
build_trigger_callback,
|
||||
has_non_synchronous_actions,
|
||||
literal_with_length,
|
||||
maybe_simple_id,
|
||||
parent_ref,
|
||||
progmem_bytes,
|
||||
register_apply_action,
|
||||
register_apply_condition,
|
||||
@@ -29,11 +34,12 @@ from esphome.automation import (
|
||||
register_parented_condition,
|
||||
register_simple_action,
|
||||
register_simple_condition,
|
||||
string_ref_literal,
|
||||
templatable_bytes,
|
||||
)
|
||||
import esphome.codegen as cg
|
||||
import esphome.config_validation as cv
|
||||
from esphome.const import CONF_ID
|
||||
from esphome.const import CONF_AUTOMATION_ID, CONF_ID, CONF_THEN
|
||||
from esphome.core import CORE, ID, KEY_CORE, KEY_TARGET_PLATFORM, EsphomeError, Lambda
|
||||
from esphome.cpp_generator import MockObj, RawExpression
|
||||
from esphome.util import Registry, RegistryEntry
|
||||
@@ -327,7 +333,14 @@ async def test_build_callback_automations_single_entry(
|
||||
(CallbackAutomation("on_state", "add_on_state_callback", [(bool, "x")]),),
|
||||
)
|
||||
mock_build_callback.assert_called_once_with(
|
||||
parent, "add_on_state_callback", [(bool, "x")], conf, forwarder=None
|
||||
parent,
|
||||
"add_on_state_callback",
|
||||
[(bool, "x")],
|
||||
conf,
|
||||
forwarder=None,
|
||||
params=None,
|
||||
forward=None,
|
||||
when=None,
|
||||
)
|
||||
|
||||
|
||||
@@ -347,10 +360,24 @@ async def test_build_callback_automations_multiple_configs(
|
||||
)
|
||||
assert mock_build_callback.call_count == 2
|
||||
mock_build_callback.assert_any_call(
|
||||
parent, "add_on_state_callback", [(bool, "x")], conf1, forwarder=None
|
||||
parent,
|
||||
"add_on_state_callback",
|
||||
[(bool, "x")],
|
||||
conf1,
|
||||
forwarder=None,
|
||||
params=None,
|
||||
forward=None,
|
||||
when=None,
|
||||
)
|
||||
mock_build_callback.assert_any_call(
|
||||
parent, "add_on_state_callback", [(bool, "x")], conf2, forwarder=None
|
||||
parent,
|
||||
"add_on_state_callback",
|
||||
[(bool, "x")],
|
||||
conf2,
|
||||
forwarder=None,
|
||||
params=None,
|
||||
forward=None,
|
||||
when=None,
|
||||
)
|
||||
|
||||
|
||||
@@ -378,9 +405,25 @@ async def test_build_callback_automations_multiple_entries(
|
||||
)
|
||||
assert mock_build_callback.call_count == 2
|
||||
assert mock_build_callback.call_args_list == [
|
||||
call(parent, "add_on_value_callback", [(float, "x")], conf_a, forwarder=None),
|
||||
call(
|
||||
parent, "add_on_raw_value_callback", [(float, "x")], conf_b, forwarder=None
|
||||
parent,
|
||||
"add_on_value_callback",
|
||||
[(float, "x")],
|
||||
conf_a,
|
||||
forwarder=None,
|
||||
params=None,
|
||||
forward=None,
|
||||
when=None,
|
||||
),
|
||||
call(
|
||||
parent,
|
||||
"add_on_raw_value_callback",
|
||||
[(float, "x")],
|
||||
conf_b,
|
||||
forwarder=None,
|
||||
params=None,
|
||||
forward=None,
|
||||
when=None,
|
||||
),
|
||||
]
|
||||
|
||||
@@ -403,7 +446,14 @@ async def test_build_callback_automations_with_forwarder(
|
||||
),
|
||||
)
|
||||
mock_build_callback.assert_called_once_with(
|
||||
parent, "add_on_state_callback", [], conf, forwarder=TriggerOnTrueForwarder
|
||||
parent,
|
||||
"add_on_state_callback",
|
||||
[],
|
||||
conf,
|
||||
forwarder=TriggerOnTrueForwarder,
|
||||
params=None,
|
||||
forward=None,
|
||||
when=None,
|
||||
)
|
||||
|
||||
|
||||
@@ -437,7 +487,14 @@ async def test_build_callback_automations_mixed_entries(
|
||||
assert mock_build_callback.call_count == 3
|
||||
assert mock_build_callback.call_args_list == [
|
||||
call(
|
||||
parent, "add_on_state_callback", [(bool, "x")], conf_state, forwarder=None
|
||||
parent,
|
||||
"add_on_state_callback",
|
||||
[(bool, "x")],
|
||||
conf_state,
|
||||
forwarder=None,
|
||||
params=None,
|
||||
forward=None,
|
||||
when=None,
|
||||
),
|
||||
call(
|
||||
parent,
|
||||
@@ -445,6 +502,9 @@ async def test_build_callback_automations_mixed_entries(
|
||||
[],
|
||||
conf_press,
|
||||
forwarder=TriggerOnTrueForwarder,
|
||||
params=None,
|
||||
forward=None,
|
||||
when=None,
|
||||
),
|
||||
call(
|
||||
parent,
|
||||
@@ -452,6 +512,9 @@ async def test_build_callback_automations_mixed_entries(
|
||||
[],
|
||||
conf_release,
|
||||
forwarder=TriggerOnFalseForwarder,
|
||||
params=None,
|
||||
forward=None,
|
||||
when=None,
|
||||
),
|
||||
]
|
||||
|
||||
@@ -477,7 +540,14 @@ async def test_build_callback_automations_skips_missing_keys(
|
||||
),
|
||||
)
|
||||
mock_build_callback.assert_called_once_with(
|
||||
parent, "add_on_state_callback", [], conf, forwarder=TriggerOnTrueForwarder
|
||||
parent,
|
||||
"add_on_state_callback",
|
||||
[],
|
||||
conf,
|
||||
forwarder=TriggerOnTrueForwarder,
|
||||
params=None,
|
||||
forward=None,
|
||||
when=None,
|
||||
)
|
||||
|
||||
|
||||
@@ -495,7 +565,14 @@ async def test_build_callback_automations_defaults(
|
||||
(CallbackAutomation("on_press", "add_on_press_callback"),),
|
||||
)
|
||||
mock_build_callback.assert_called_once_with(
|
||||
parent, "add_on_press_callback", [], conf, forwarder=None
|
||||
parent,
|
||||
"add_on_press_callback",
|
||||
[],
|
||||
conf,
|
||||
forwarder=None,
|
||||
params=None,
|
||||
forward=None,
|
||||
when=None,
|
||||
)
|
||||
|
||||
|
||||
@@ -512,9 +589,59 @@ class MockCodegen(NamedTuple):
|
||||
new_pvariable: MagicMock
|
||||
register_parented: AsyncMock
|
||||
add_global: MagicMock
|
||||
add: MagicMock
|
||||
calls: MagicMock # new_pvariable and add_global attached, to check their order
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_build_automation() -> Generator[AsyncMock]:
|
||||
with patch("esphome.automation.build_automation", new_callable=AsyncMock) as mock:
|
||||
yield mock
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_trigger_automations_with_parent(
|
||||
mock_build_automation: AsyncMock,
|
||||
) -> None:
|
||||
"""Each entry's Trigger class is instantiated with the parent and built with its args."""
|
||||
parent = MockObj("var", "->")
|
||||
on_conf = {"trigger_id": ID("trig_1"), "then": []}
|
||||
set_conf = {"trigger_id": ID("trig_2"), "then": []}
|
||||
config = {"on_turn_on": [on_conf], "on_speed_set": [set_conf]}
|
||||
with patch("esphome.codegen.new_Pvariable") as new_pvariable:
|
||||
new_pvariable.side_effect = lambda id_, *args: MockObj(str(id_), "->")
|
||||
await build_trigger_automations(
|
||||
parent,
|
||||
config,
|
||||
(
|
||||
("on_turn_on", []),
|
||||
("on_turn_off", []),
|
||||
("on_speed_set", [(cg.int_, "x")]),
|
||||
),
|
||||
)
|
||||
assert [c.args for c in new_pvariable.call_args_list] == [
|
||||
(ID("trig_1"), parent),
|
||||
(ID("trig_2"), parent),
|
||||
]
|
||||
calls = mock_build_automation.call_args_list
|
||||
assert [(str(c.args[0]), c.args[1], c.args[2]) for c in calls] == [
|
||||
("trig_1", [], on_conf),
|
||||
("trig_2", [(cg.int_, "x")], set_conf),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_trigger_automations_without_parent(
|
||||
mock_build_automation: AsyncMock,
|
||||
) -> None:
|
||||
"""A None parent instantiates the Trigger class with no constructor arguments."""
|
||||
conf = {"trigger_id": ID("trig_1"), "then": []}
|
||||
with patch("esphome.codegen.new_Pvariable") as new_pvariable:
|
||||
await build_trigger_automations(None, {"on_boot": [conf]}, (("on_boot", []),))
|
||||
new_pvariable.assert_called_once_with(ID("trig_1"))
|
||||
mock_build_automation.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_cg() -> Generator[MockCodegen]:
|
||||
"""Patch the codegen calls the shared builders make."""
|
||||
@@ -525,6 +652,7 @@ def mock_cg() -> Generator[MockCodegen]:
|
||||
"esphome.codegen.register_parented", new_callable=AsyncMock
|
||||
) as register_parented,
|
||||
patch("esphome.cpp_generator.add_global") as add_global,
|
||||
patch("esphome.codegen.add") as add,
|
||||
):
|
||||
get_variable.return_value = PARENT_OBJ
|
||||
new_pvariable.return_value = NEW_OBJ
|
||||
@@ -532,7 +660,7 @@ def mock_cg() -> Generator[MockCodegen]:
|
||||
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
|
||||
get_variable, new_pvariable, register_parented, add_global, add, calls
|
||||
)
|
||||
|
||||
|
||||
@@ -927,6 +1055,26 @@ async def test_apply_literal_with_length_is_plain_on_every_platform(
|
||||
assert "progmem_string" not in text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("platform", "rendered"),
|
||||
[("esp32", 'StringRef("ON", 2)'), ("esp8266", 'ESPHOME_F("ON")')],
|
||||
)
|
||||
async def test_apply_string_ref_literal_stays_in_flash_on_esp8266(
|
||||
registries: tuple[Registry, Registry],
|
||||
mock_cg: MockCodegen,
|
||||
platform: str,
|
||||
rendered: str,
|
||||
) -> None:
|
||||
fields = (
|
||||
ApplyField(
|
||||
"payload", "set_payload", cg.std_string, const_fn=string_ref_literal
|
||||
),
|
||||
)
|
||||
await _run_apply_action(registries, fields, {"payload": "ON"}, platform=platform)
|
||||
assert f"::{PARENT_OBJ}->set_payload({rendered});" in _apply_definition(mock_cg)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_apply_condition_predicate(
|
||||
registries: tuple[Registry, Registry], mock_cg: MockCodegen
|
||||
@@ -1036,3 +1184,120 @@ async def test_templatable_bytes_rejects_oversized_payload() -> None:
|
||||
await templatable_bytes(
|
||||
[0] * 0x10000, [], var.set_code_template, var.set_code_static, "payload"
|
||||
)
|
||||
|
||||
|
||||
TRIGGER_CONF = {CONF_AUTOMATION_ID: ID("automation_1"), CONF_THEN: []}
|
||||
|
||||
|
||||
def _squash(expr: object) -> str:
|
||||
"""Generated text with its whitespace folded, for one-line assertions."""
|
||||
return " ".join(str(expr).split())
|
||||
|
||||
|
||||
def test_parent_ref_is_global_scoped() -> None:
|
||||
"""A parent named through parent_ref cannot be shadowed by a trigger argument."""
|
||||
assert (
|
||||
str(parent_ref(MockObj("sel", "->")).option_at(RawExpression("i")))
|
||||
== "::sel->option_at(i)"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_trigger_callback_forwards_params(mock_cg: MockCodegen) -> None:
|
||||
"""Without forward or when the callback passes its parameters straight to trigger()."""
|
||||
text = str(
|
||||
await build_trigger_callback(
|
||||
[(cg.bool_, "state")], TRIGGER_CONF, [(cg.bool_, "state")]
|
||||
)
|
||||
)
|
||||
assert text.startswith("[](const std::remove_cvref_t<bool> & state) -> void {")
|
||||
assert f"::{NEW_OBJ}->trigger(state);" in text
|
||||
assert "if (" not in text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_trigger_callback_reshapes_and_filters(mock_cg: MockCodegen) -> None:
|
||||
"""A filter returns early, forward picks the trigger args, a string constant is a plain literal."""
|
||||
when = ApplyCall("payload == {}", (("payload", cg.std_string),))
|
||||
params = [(cg.std_string, "topic"), (cg.std_string, "payload")]
|
||||
text = _squash(
|
||||
await build_trigger_callback(
|
||||
[(cg.std_string, "x")],
|
||||
{**TRIGGER_CONF, "payload": "hi"},
|
||||
params,
|
||||
forward=["payload"],
|
||||
when=when,
|
||||
)
|
||||
)
|
||||
assert "& topic, const std::remove_cvref_t<std::string> & payload) -> void" in text
|
||||
assert f'if (!(payload == "hi")) return; ::{NEW_OBJ}->trigger(payload);' in text
|
||||
# An absent optional key skips the filter, as for any ApplyCall.
|
||||
text = str(
|
||||
await build_trigger_callback(
|
||||
[(cg.std_string, "x")], TRIGGER_CONF, params, forward=["payload"], when=when
|
||||
)
|
||||
)
|
||||
assert "if (" not in text
|
||||
|
||||
|
||||
@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}"):
|
||||
await build_trigger_callback(
|
||||
[], {**TRIGGER_CONF, "mode": 1}, [(cg.int_, "mode")], forward=[], when=when
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_parent_callback_automation(mock_cg: MockCodegen) -> None:
|
||||
"""A no-argument callback registers a lambda that hands the parent to the automation."""
|
||||
parent = MockObj("fan", "->")
|
||||
await build_parent_callback_automation(
|
||||
parent,
|
||||
"add_on_state_callback",
|
||||
(cg.RawExpression("Fan *"), "x"),
|
||||
TRIGGER_CONF,
|
||||
)
|
||||
assert _squash(mock_cg.add.call_args.args[0]) == (
|
||||
f"fan->add_on_state_callback([]() -> void {{ ::{NEW_OBJ}->trigger(::fan); }})"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_callback_automation_lambda(mock_cg: MockCodegen) -> None:
|
||||
"""Reshaping keywords switch the registration to the lambda; forwarder cannot join them."""
|
||||
parent = MockObj("sel", "->")
|
||||
await build_callback_automation(
|
||||
parent,
|
||||
"add_cb",
|
||||
[(cg.std_string, "x"), (cg.size_t, "i")],
|
||||
TRIGGER_CONF,
|
||||
params=[(cg.size_t, "index")],
|
||||
forward=[parent_ref(parent).option_at(RawExpression("index")), "index"],
|
||||
)
|
||||
assert _squash(mock_cg.add.call_args.args[0]) == (
|
||||
"sel->add_cb([](const std::remove_cvref_t<size_t> & index) -> void { "
|
||||
f"::{NEW_OBJ}->trigger(::sel->option_at(index), index); }})"
|
||||
)
|
||||
with pytest.raises(ValueError, match="forwarder"):
|
||||
await build_callback_automation(
|
||||
parent,
|
||||
"add_cb",
|
||||
[],
|
||||
TRIGGER_CONF,
|
||||
forwarder=TriggerOnTrueForwarder,
|
||||
when="x",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_callback_automation_forwarder(mock_cg: MockCodegen) -> None:
|
||||
"""The forwarder path registers a pointer-sized TriggerForwarder, not a lambda."""
|
||||
await build_callback_automation(
|
||||
MockObj("parent", "->"), "add_cb", [(cg.bool_, "x")], TRIGGER_CONF
|
||||
)
|
||||
assert str(mock_cg.add.call_args.args[0]) == (
|
||||
f"parent->add_cb(TriggerForwarder<bool>{{{NEW_OBJ}}})"
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user