From ae7d96a2be643b53bcb6b05ee5db6a0391e35f5e Mon Sep 17 00:00:00 2001 From: "J. Nick Koston" Date: Fri, 25 Sep 2026 15:46:46 +0100 Subject: [PATCH] [core][mqtt] Build trigger callbacks as stateless lambdas --- AGENTS.md | 18 ++ esphome/automation.py | 184 +++++++++-- esphome/codegen.py | 1 + esphome/components/mqtt/__init__.py | 87 +++--- esphome/components/mqtt/mqtt_client.cpp | 70 ++--- esphome/components/mqtt/mqtt_client.h | 90 +++--- esphome/components/mqtt/mqtt_component.cpp | 8 - esphome/components/mqtt/mqtt_component.h | 8 +- esphome/core/helpers.h | 32 +- esphome/cpp_generator.py | 7 + tests/components/mqtt/common-triggers.yaml | 25 ++ .../mqtt/test-triggers.esp32-idf.yaml | 3 + tests/unit_tests/test_automation.py | 287 +++++++++++++++++- 13 files changed, 625 insertions(+), 195 deletions(-) create mode 100644 tests/components/mqtt/common-triggers.yaml create mode 100644 tests/components/mqtt/test-triggers.esp32-idf.yaml diff --git a/AGENTS.md b/AGENTS.md index 6a60fc47b17..bfa0124e234 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -384,6 +384,22 @@ file does, and it is the authority when they disagree. The most useful starting ) ``` + When the parent callback's parameters are not the automation's arguments, or the trigger should only fire for some values, pass `params`, `forward` and `when`; the helper then generates a capture-less lambda instead of a forwarder (stored inline by `Callback` and `std::function`, nothing is allocated). `params` are the callback's parameters as `[(type, name)]`, `forward` the expressions passed to `trigger()` (default: the parameter names; name the parent with `automation.parent_ref(var)`), `when` a filter: a string, or an `ApplyCall` compared against config values that is skipped when its keys are absent. Compose `forward` from `MockObj` calls, not from f-strings of C++: + ```python + # select: the callback carries the index, the automation also gets the option text + parent = automation.parent_ref(var) + index = cg.RawExpression("index") + await automation.build_callback_automation( + var, + "add_on_state_callback", + [(cg.StringRef, "x"), (cg.size_t, "i")], + conf, + params=[(cg.size_t, "index")], + forward=[cg.StringRef(parent.option_at(index)), index], + ) + ``` + A callback with no arguments whose automation receives the parent is one line, `automation.build_parent_callback_automation(var, "add_on_state_callback", (Fan.operator("ptr"), "x"), conf)`. When the registration takes extra arguments (mqtt's topic and qos), `automation.build_trigger_callback(args, conf, params=..., forward=..., when=...)` returns the lambda for the component to register itself. Several callbacks on one parent go in a module-level `_CALLBACK_AUTOMATIONS` tuple of `automation.CallbackAutomation(conf_key, callback_method, args, ...)` entries, applied with `automation.build_callback_automations(var, config, _CALLBACK_AUTOMATIONS)`; each entry takes the same optional `forwarder`, `params`, `forward` and `when`. + **C++ -- no trigger class needed.** The callback registration method must be templatized to accept both `std::function` and lightweight forwarder structs (which avoid heap allocation): ```cpp class MyComponent : public Component { @@ -405,6 +421,8 @@ file does, and it is the authority when they disagree. The most useful starting Use `build_automation()` with a `Trigger` subclass only when the forwarder needs **mutable state beyond a single `Automation*` pointer** (e.g. edge detection tracking previous state, timing logic). + Several such triggers on one parent go in a module-level `_TRIGGER_AUTOMATIONS` tuple of `(conf_key, args)` pairs applied with `automation.build_trigger_automations(var, config, _TRIGGER_AUTOMATIONS)`; it instantiates each entry's class from its `CONF_TRIGGER_ID` with `var` (or nothing when `var` is `None`) and builds the automations. + **Python:** ```python TurnOnTrigger = my_ns.class_("TurnOnTrigger", automation.Trigger.template()) diff --git a/esphome/automation.py b/esphome/automation.py index c8d0e4c5743..20578e06a53 100644 --- a/esphome/automation.py +++ b/esphome/automation.py @@ -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, Lambda from esphome.cpp_generator import ( - FlashStringLiteral, + Expression, LambdaExpression, MockObj, MockObjClass, @@ -219,9 +219,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: @@ -233,6 +231,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_), ...)``. @@ -339,9 +344,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: @@ -373,7 +385,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]: @@ -451,7 +463,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 @@ -470,6 +482,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: @@ -495,22 +528,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) @@ -1030,23 +1057,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. @@ -1058,6 +1139,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. @@ -1066,22 +1150,56 @@ async def build_callback_automation( TriggerForwarder. 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.""" @@ -1090,6 +1208,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( @@ -1111,4 +1232,7 @@ async def build_callback_automations( entry.args, conf, forwarder=entry.forwarder, + params=entry.params, + forward=entry.forward, + when=entry.when, ) diff --git a/esphome/codegen.py b/esphome/codegen.py index 4de1d8d5c19..2490804c219 100644 --- a/esphome/codegen.py +++ b/esphome/codegen.py @@ -38,6 +38,7 @@ from esphome.cpp_generator import ( # noqa: F401 new_variable, process_lambda, progmem_array, + progmem_string, safe_exp, set_cpp_standard, statement, diff --git a/esphome/components/mqtt/__init__.py b/esphome/components/mqtt/__init__.py index b6badb4ef96..ea79d1a401c 100644 --- a/esphome/components/mqtt/__init__.py +++ b/esphome/components/mqtt/__init__.py @@ -53,7 +53,6 @@ from esphome.const import ( CONF_SUBSCRIBE_QOS, CONF_TOPIC, CONF_TOPIC_PREFIX, - CONF_TRIGGER_ID, CONF_USE_ABBREVIATIONS, CONF_USERNAME, CONF_WILL_MESSAGE, @@ -118,18 +117,6 @@ MQTTMessage = mqtt_ns.struct("MQTTMessage") MQTTClientDisconnectReason = mqtt_ns.enum("MQTTClientDisconnectReason") MQTTClientComponent = mqtt_ns.class_("MQTTClientComponent", cg.Component) MQTTPublishJsonAction = mqtt_ns.class_("MQTTPublishJsonAction", automation.Action) -MQTTMessageTrigger = mqtt_ns.class_( - "MQTTMessageTrigger", automation.Trigger.template(cg.std_string), cg.Component -) -MQTTJsonMessageTrigger = mqtt_ns.class_( - "MQTTJsonMessageTrigger", automation.Trigger.template(cg.JsonObjectConst) -) -MQTTConnectTrigger = mqtt_ns.class_( - "MQTTConnectTrigger", automation.Trigger.template(cg.bool_) -) -MQTTDisconnectTrigger = mqtt_ns.class_( - "MQTTDisconnectTrigger", automation.Trigger.template(MQTTClientDisconnectReason) -) MQTTComponent = mqtt_ns.class_("MQTTComponent", cg.Component) MQTTAlarmControlPanelComponent = mqtt_ns.class_( @@ -283,21 +270,10 @@ CONFIG_SCHEMA = cv.All( cv.Optional( CONF_REBOOT_TIMEOUT, default="15min" ): cv.positive_time_period_milliseconds, - cv.Optional(CONF_ON_CONNECT): automation.validate_automation( - { - cv.GenerateID(CONF_TRIGGER_ID): cv.declare_id(MQTTConnectTrigger), - } - ), - cv.Optional(CONF_ON_DISCONNECT): automation.validate_automation( - { - cv.GenerateID(CONF_TRIGGER_ID): cv.declare_id( - MQTTDisconnectTrigger - ), - } - ), + cv.Optional(CONF_ON_CONNECT): automation.validate_automation(), + cv.Optional(CONF_ON_DISCONNECT): automation.validate_automation(), cv.Optional(CONF_ON_MESSAGE): automation.validate_automation( { - cv.GenerateID(CONF_TRIGGER_ID): cv.declare_id(MQTTMessageTrigger), cv.Required(CONF_TOPIC): cv.subscribe_topic, cv.Optional(CONF_QOS, default=0): cv.mqtt_qos, cv.Optional(CONF_PAYLOAD): cv.string_strict, @@ -305,9 +281,6 @@ CONFIG_SCHEMA = cv.All( ), cv.Optional(CONF_ON_JSON_MESSAGE): automation.validate_automation( { - cv.GenerateID(CONF_TRIGGER_ID): cv.declare_id( - MQTTJsonMessageTrigger - ), cv.Required(CONF_TOPIC): cv.subscribe_topic, cv.Optional(CONF_QOS, default=0): cv.mqtt_qos, } @@ -342,6 +315,18 @@ def exp_mqtt_message(config): ) +_CALLBACK_AUTOMATIONS = ( + automation.CallbackAutomation( + CONF_ON_CONNECT, "set_on_connect", [(cg.bool_, "session_present")] + ), + automation.CallbackAutomation( + CONF_ON_DISCONNECT, + "set_on_disconnect", + [(MQTTClientDisconnectReason, "reason")], + ), +) + + @coroutine_with_priority(CoroPriority.WEB) async def to_code(config): var = cg.new_Pvariable(config[CONF_ID]) @@ -460,29 +445,37 @@ async def to_code(config): cg.add_define("USE_MQTT_IDF_ENQUEUE") # end esp-idf + # The client queues subscriptions until it connects, so they can be made at construction. for conf in config.get(CONF_ON_MESSAGE, []): - trig = cg.new_Pvariable(conf[CONF_TRIGGER_ID], conf[CONF_TOPIC]) - cg.add(trig.set_qos(conf[CONF_QOS])) - if CONF_PAYLOAD in conf: - cg.add(trig.set_payload(conf[CONF_PAYLOAD])) - await cg.register_component(trig, conf) - await automation.build_automation(trig, [(cg.std_string, "x")], conf) + callback = await automation.build_trigger_callback( + [(cg.std_string, "x")], + conf, + params=[(cg.std_string, "topic"), (cg.std_string, "payload")], + forward=["payload"], + # The length is compared first; the literal stays in flash on ESP8266. + when=automation.ApplyCall( + "StringRef(payload) == {}", + ((CONF_PAYLOAD, cg.std_string, automation.string_ref_literal),), + ), + ) + cg.add( + var.subscribe(cg.progmem_string(conf[CONF_TOPIC]), callback, conf[CONF_QOS]) + ) for conf in config.get(CONF_ON_JSON_MESSAGE, []): - trig = cg.new_Pvariable(conf[CONF_TRIGGER_ID], conf[CONF_TOPIC], conf[CONF_QOS]) - await automation.build_automation(trig, [(cg.JsonObjectConst, "x")], conf) - - for conf in config.get(CONF_ON_CONNECT, []): - trigger = cg.new_Pvariable(conf[CONF_TRIGGER_ID], var) - await automation.build_automation( - trigger, [(cg.bool_, "session_present")], conf + callback = await automation.build_trigger_callback( + [(cg.JsonObjectConst, "x")], + conf, + params=[(cg.std_string, "topic"), (cg.JsonObject, "root")], + forward=["root"], + ) + cg.add( + var.subscribe_json( + cg.progmem_string(conf[CONF_TOPIC]), callback, conf[CONF_QOS] + ) ) - for conf in config.get(CONF_ON_DISCONNECT, []): - trigger = cg.new_Pvariable(conf[CONF_TRIGGER_ID], var) - await automation.build_automation( - trigger, [(MQTTClientDisconnectReason, "reason")], conf - ) + await automation.build_callback_automations(var, config, _CALLBACK_AUTOMATIONS) cg.add(var.set_publish_nan_as_none(config[CONF_PUBLISH_NAN_AS_NONE])) diff --git a/esphome/components/mqtt/mqtt_client.cpp b/esphome/components/mqtt/mqtt_client.cpp index 2ecab47904a..2fdb29485c9 100644 --- a/esphome/components/mqtt/mqtt_client.cpp +++ b/esphome/components/mqtt/mqtt_client.cpp @@ -461,37 +461,24 @@ void MQTTClientComponent::resubscribe_subscriptions_() { } } -void MQTTClientComponent::subscribe(const std::string &topic, mqtt_callback_t callback, uint8_t qos) { - MQTTSubscription subscription{ - .topic = topic, +void MQTTClientComponent::add_subscription_(std::string &&topic, mqtt_callback_t callback, bool boxed, uint8_t qos) { + auto &subscription = this->subscriptions_.emplace_back(MQTTSubscription{ + .topic = std::move(topic), .qos = qos, - .callback = std::move(callback), .subscribed = false, + .boxed = boxed, + .callback = callback, .resubscribe_timeout = 0, - }; + }); this->resubscribe_subscription_(&subscription); - this->subscriptions_.push_back(subscription); -} - -void MQTTClientComponent::subscribe_json(const std::string &topic, const mqtt_json_callback_t &callback, uint8_t qos) { - auto f = [callback](const std::string &topic, const std::string &payload) { - json::parse_json(payload, [topic, callback](JsonObject root) -> bool { - callback(topic, root); - return true; - }); - }; - MQTTSubscription subscription{ - .topic = topic, - .qos = qos, - .callback = f, - .subscribed = false, - .resubscribe_timeout = 0, - }; - this->resubscribe_subscription_(&subscription); - this->subscriptions_.push_back(subscription); } void MQTTClientComponent::unsubscribe(const std::string &topic) { + if (this->dispatching_) { + // Would erase, and for a boxed callback free, the subscription that may be running right now. + ESP_LOGE(TAG, "Cannot unsubscribe from '%s' inside a subscription callback", topic.c_str()); + return; + } bool ret = this->mqtt_backend_.unsubscribe(topic.c_str()); yield(); if (ret) { @@ -505,6 +492,8 @@ void MQTTClientComponent::unsubscribe(const std::string &topic) { auto it = subscriptions_.begin(); while (it != subscriptions_.end()) { if (it->topic == topic) { + if (it->boxed) + it->callback.free_boxed(); it = subscriptions_.erase(it); } else { ++it; @@ -660,10 +649,16 @@ void MQTTClientComponent::on_message(const std::string &topic, const std::string // in simple tests but will cause crashes with complex automations. this->defer([this, topic, payload]() { #endif - for (auto &subscription : this->subscriptions_) { + // Indexed with the count taken up front: a callback may subscribe, which can reallocate the + // vector, and the new entry waits for the next message. + this->dispatching_ = true; + const size_t count = this->subscriptions_.size(); + for (size_t i = 0; i < count; i++) { + const auto &subscription = this->subscriptions_[i]; if (topic_match(topic.c_str(), subscription.topic.c_str())) - subscription.callback(topic, payload); + subscription.callback.call(topic, payload); } + this->dispatching_ = false; #ifdef USE_ESP8266 }); #endif @@ -762,29 +757,6 @@ void MQTTClientComponent::set_on_disconnect(mqtt_on_disconnect_callback_t &&call MQTTClientComponent *global_mqtt_client = nullptr; // NOLINT(cppcoreguidelines-avoid-non-const-global-variables) -// MQTTMessageTrigger -MQTTMessageTrigger::MQTTMessageTrigger(std::string topic) : topic_(std::move(topic)) {} -void MQTTMessageTrigger::setup() { - global_mqtt_client->subscribe( - this->topic_, - [this](const std::string &topic, const std::string &payload) { - if (this->payload_.has_value() && payload != *this->payload_) { - return; - } - - this->trigger(payload); - }, - this->qos_); -} -void MQTTMessageTrigger::dump_config() { - ESP_LOGCONFIG(TAG, - "MQTT Message Trigger:\n" - " Topic: '%s'\n" - " QoS: %u", - this->topic_.c_str(), this->qos_); -} -float MQTTMessageTrigger::get_setup_priority() const { return setup_priority::AFTER_CONNECTION; } - } // namespace esphome::mqtt #endif // USE_MQTT diff --git a/esphome/components/mqtt/mqtt_client.h b/esphome/components/mqtt/mqtt_client.h index ced9c84e100..f16a285eed3 100644 --- a/esphome/components/mqtt/mqtt_client.h +++ b/esphome/components/mqtt/mqtt_client.h @@ -31,19 +31,16 @@ namespace esphome::mqtt { using mqtt_on_connect_callback_t = std::function; using mqtt_on_disconnect_callback_t = std::function; -/** Callback for MQTT subscriptions. - * - * First parameter is the topic, the second one is the payload. - */ -using mqtt_callback_t = std::function; -using mqtt_json_callback_t = std::function; +/// Callback for MQTT subscriptions: the topic, then the payload. +using mqtt_callback_t = Callback; /// internal struct for MQTT subscriptions. struct MQTTSubscription { std::string topic; uint8_t qos; - mqtt_callback_t callback; bool subscribed; + bool boxed; // made by create_boxed(), freed on unsubscribe + mqtt_callback_t callback; uint32_t resubscribe_timeout; }; @@ -165,12 +162,21 @@ class MQTTClientComponent final : public Component { bool is_log_message_enabled() const; /** Subscribe to an MQTT topic and call callback when a message is received. + * + * A callback that fits in a pointer (capture at most `this`) is stored inline; a larger one is + * boxed on the heap and freed by unsubscribe(). * * @param topic The topic. Wildcards are currently not supported. * @param callback The callback function. * @param qos The QoS of this subscription. */ - void subscribe(const std::string &topic, mqtt_callback_t callback, uint8_t qos = 0); + template void subscribe(std::string topic, F &&callback, uint8_t qos = 0) { + if constexpr (mqtt_callback_t::fits_inline()) { + this->add_subscription_(std::move(topic), mqtt_callback_t::create(std::forward(callback)), false, qos); + } else { + this->add_subscription_(std::move(topic), mqtt_callback_t::create_boxed(std::forward(callback)), true, qos); + } + } /** Subscribe to a MQTT topic and automatically parse JSON payload. * @@ -181,11 +187,30 @@ class MQTTClientComponent final : public Component { * received. * @param qos The QoS of this subscription. */ - void subscribe_json(const std::string &topic, const mqtt_json_callback_t &callback, uint8_t qos = 0); + template void subscribe_json(std::string topic, F &&callback, uint8_t qos = 0) { + using DecayF = std::decay_t; + if constexpr (std::is_invocable_v) { + this->subscribe( + std::move(topic), + [cb = std::forward(callback)](const std::string &topic, const std::string &payload) { + call_json(cb, topic, payload); + }, + qos); + } else { + // A mutable callback keeps its state; the wrapper is then boxed rather than copied per call. + this->subscribe( + std::move(topic), + [cb = std::forward(callback)](const std::string &topic, const std::string &payload) mutable { + call_json(cb, topic, payload); + }, + qos); + } + } /** Unsubscribe from an MQTT topic. * * If multiple existing subscriptions to the same topic exist, all of them will be removed. + * Not allowed from inside a subscription callback; such a call is logged and ignored. * * @param topic The topic to unsubscribe from. * Must match the topic in the original subscribe or subscribe_json call exactly. @@ -285,6 +310,14 @@ class MQTTClientComponent final : public Component { bool subscribe_(const char *topic, uint8_t qos); void resubscribe_subscription_(MQTTSubscription *sub); + // parse_json runs its callback before returning, so the captures can be references. + template static void call_json(F &callback, const std::string &topic, const std::string &payload) { + json::parse_json(payload, [&topic, &callback](JsonObject root) -> bool { + callback(topic, root); + return true; + }); + } + void add_subscription_(std::string &&topic, mqtt_callback_t callback, bool boxed, uint8_t qos); void resubscribe_subscriptions_(); MQTTCredentials credentials_; @@ -336,48 +369,11 @@ class MQTTClientComponent final : public Component { bool publish_nan_as_none_{false}; bool wait_for_connection_{false}; + bool dispatching_{false}; }; extern MQTTClientComponent *global_mqtt_client; // NOLINT(cppcoreguidelines-avoid-non-const-global-variables) -class MQTTMessageTrigger final : public Trigger, public Component { - public: - explicit MQTTMessageTrigger(std::string topic); - - void set_qos(uint8_t qos) { this->qos_ = qos; } - void set_payload(const std::string &payload) { this->payload_ = payload; } - void setup() override; - void dump_config() override; - float get_setup_priority() const override; - - protected: - std::string topic_; - uint8_t qos_{0}; - optional payload_; -}; - -class MQTTJsonMessageTrigger final : public Trigger { - public: - explicit MQTTJsonMessageTrigger(const std::string &topic, uint8_t qos) { - global_mqtt_client->subscribe_json( - topic, [this](const std::string &topic, JsonObject root) { this->trigger(root); }, qos); - } -}; - -class MQTTConnectTrigger final : public Trigger { - public: - explicit MQTTConnectTrigger(MQTTClientComponent *client) { - client->set_on_connect([this](bool session_present) { this->trigger(session_present); }); - } -}; - -class MQTTDisconnectTrigger final : public Trigger { - public: - explicit MQTTDisconnectTrigger(MQTTClientComponent *client) { - client->set_on_disconnect([this](MQTTClientDisconnectReason reason) { this->trigger(reason); }); - } -}; - template class MQTTPublishJsonAction final : public Action { public: MQTTPublishJsonAction(MQTTClientComponent *parent) : parent_(parent) {} diff --git a/esphome/components/mqtt/mqtt_component.cpp b/esphome/components/mqtt/mqtt_component.cpp index 59a5d02d977..6263ea933fd 100644 --- a/esphome/components/mqtt/mqtt_component.cpp +++ b/esphome/components/mqtt/mqtt_component.cpp @@ -342,14 +342,6 @@ bool MQTTComponent::is_discovery_enabled() const { return this->discovery_enabled_ && global_mqtt_client->is_discovery_enabled(); } -void MQTTComponent::subscribe(const std::string &topic, mqtt_callback_t callback, uint8_t qos) { - global_mqtt_client->subscribe(topic, std::move(callback), qos); -} - -void MQTTComponent::subscribe_json(const std::string &topic, const mqtt_json_callback_t &callback, uint8_t qos) { - global_mqtt_client->subscribe_json(topic, callback, qos); -} - MQTTComponent::MQTTComponent() = default; float MQTTComponent::get_setup_priority() const { return setup_priority::AFTER_CONNECTION; } diff --git a/esphome/components/mqtt/mqtt_component.h b/esphome/components/mqtt/mqtt_component.h index b4ae6244048..bd99d668fe5 100644 --- a/esphome/components/mqtt/mqtt_component.h +++ b/esphome/components/mqtt/mqtt_component.h @@ -259,7 +259,9 @@ class MQTTComponent : public Component { * @param callback The callback that will be called when a message with matching topic is received. * @param qos The MQTT quality of service. Defaults to 0. */ - void subscribe(const std::string &topic, mqtt_callback_t callback, uint8_t qos = 0); + template void subscribe(std::string topic, F &&callback, uint8_t qos = 0) { + global_mqtt_client->subscribe(std::move(topic), std::forward(callback), qos); + } /** Subscribe to a MQTT topic and automatically parse JSON payload. * @@ -270,7 +272,9 @@ class MQTTComponent : public Component { * received. * @param qos The MQTT quality of service. Defaults to 0. */ - void subscribe_json(const std::string &topic, const mqtt_json_callback_t &callback, uint8_t qos = 0); + template void subscribe_json(std::string topic, F &&callback, uint8_t qos = 0) { + global_mqtt_client->subscribe_json(std::move(topic), std::forward(callback), qos); + } protected: /// Helper method to get the discovery topic for this component into a buffer. diff --git a/esphome/core/helpers.h b/esphome/core/helpers.h index 6d00e187999..fb140f9ae4c 100644 --- a/esphome/core/helpers.h +++ b/esphome/core/helpers.h @@ -1719,11 +1719,19 @@ template struct Callback { /// Invoke the callback. Only valid on Callbacks created via create(), never on default-constructed instances. void call(Ts... args) const { this->fn_(this->ctx_, std::forward(args)...); } + /// Whether create() stores F inline in the ctx pointer. The inline path invokes a copy, so a + /// callable that mutates its captures (a mutable lambda) is kept whole on the heap instead. + template static constexpr bool fits_inline() { + using DecayF = std::decay_t; + return sizeof(DecayF) <= sizeof(void *) && std::is_trivially_copyable_v && + std::is_invocable_v; + } + /// Create from any callable. Small trivially-copyable callables (like [this] lambdas) /// are stored inline in the ctx pointer without heap allocation. template static Callback create(F &&callable) { using DecayF = std::decay_t; - if constexpr (sizeof(DecayF) <= sizeof(void *) && std::is_trivially_copyable_v) { + if constexpr (fits_inline()) { // Small trivial callable (e.g. [this]() { this->method(); }) - store inline in ctx. // Safe under C++20 (P0593R6): byte copy into aligned storage implicitly // creates objects of implicit-lifetime types (trivially copyable qualifies). @@ -1746,6 +1754,28 @@ template struct Callback { return {[](void *c, Ts... args) { (*static_cast(c))(args...); }, static_cast(stored)}; } } + + /// Heap home of a callable stored by create_boxed(); the header knows how to free it. + struct Box { + void (*free)(Box *box); + }; + + /// Store any callable on the heap with a deleter, for an owner that can later free_boxed() it. + template static Callback create_boxed(F &&callable) { + struct Boxed : Box { + std::decay_t fn; + }; + auto *box = new Boxed{{[](Box *b) { delete static_cast(b); }}, // NOLINT(cppcoreguidelines-owning-memory) + std::forward(callable)}; + return {[](void *c, Ts... args) { static_cast(static_cast(c))->fn(args...); }, + static_cast(box)}; + } + + /// Free the callable of a Callback made by create_boxed(); the Callback must not be called afterwards. + void free_boxed() const { + auto *box = static_cast(this->ctx_); + box->free(box); + } }; /// Grow a CallbackManager's backing array to exactly size+1. Defined in helpers.cpp. diff --git a/esphome/cpp_generator.py b/esphome/cpp_generator.py index b0c3533e040..923a37c75bb 100644 --- a/esphome/cpp_generator.py +++ b/esphome/cpp_generator.py @@ -264,6 +264,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",) diff --git a/tests/components/mqtt/common-triggers.yaml b/tests/components/mqtt/common-triggers.yaml new file mode 100644 index 00000000000..dadfc6f53b9 --- /dev/null +++ b/tests/components/mqtt/common-triggers.yaml @@ -0,0 +1,25 @@ +# on_message with a payload filter, a subscription callback larger than a pointer with its unsubscribe, +# and mutable callbacks, which keep their state instead of being copied per call. +mqtt: + id: mqtt_client + on_message: + - topic: livingroom/ota_mode + payload: "ON" + qos: 1 + then: + - logger.log: Got livingroom/ota_mode ON + +esphome: + on_boot: + - lambda: |- + std::string tag = "boxed"; + id(mqtt_client).subscribe("some/topic/boxed", [tag](const std::string &topic, const std::string &payload) { + ESP_LOGD("test", "%s %s", tag.c_str(), payload.c_str()); + }); + id(mqtt_client).unsubscribe("some/topic/boxed"); + id(mqtt_client).subscribe("some/topic/counted", [n = 0](const std::string &topic, const std::string &payload) mutable { + ESP_LOGD("test", "message %d: %s", ++n, payload.c_str()); + }); + id(mqtt_client).subscribe_json("some/topic/counted_json", [n = 0](const std::string &topic, JsonObject root) mutable { + ESP_LOGD("test", "json %d", ++n); + }); diff --git a/tests/components/mqtt/test-triggers.esp32-idf.yaml b/tests/components/mqtt/test-triggers.esp32-idf.yaml new file mode 100644 index 00000000000..b34863e8737 --- /dev/null +++ b/tests/components/mqtt/test-triggers.esp32-idf.yaml @@ -0,0 +1,3 @@ +packages: + common: !include common.yaml + triggers: !include common-triggers.yaml diff --git a/tests/unit_tests/test_automation.py b/tests/unit_tests/test_automation.py index 3a902a429f6..aac117b5a52 100644 --- a/tests/unit_tests/test_automation.py +++ b/tests/unit_tests/test_automation.py @@ -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, register_apply_action, register_apply_condition, register_bare_action, @@ -28,10 +33,11 @@ from esphome.automation import ( register_parented_condition, register_simple_action, register_simple_condition, + string_ref_literal, ) 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 @@ -325,7 +331,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, ) @@ -345,10 +358,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, ) @@ -376,9 +403,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, ), ] @@ -401,7 +444,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, ) @@ -435,7 +485,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, @@ -443,6 +500,9 @@ async def test_build_callback_automations_mixed_entries( [], conf_press, forwarder=TriggerOnTrueForwarder, + params=None, + forward=None, + when=None, ), call( parent, @@ -450,6 +510,9 @@ async def test_build_callback_automations_mixed_entries( [], conf_release, forwarder=TriggerOnFalseForwarder, + params=None, + forward=None, + when=None, ), ] @@ -475,7 +538,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, ) @@ -493,7 +563,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, ) @@ -510,9 +587,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.""" @@ -523,6 +650,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 @@ -530,7 +658,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 ) @@ -925,6 +1053,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 @@ -989,3 +1137,120 @@ async def test_apply_condition_string_lambda_paths( text = _apply_definition(mock_cg) assert expected in text assert ("-> std::string {" in text) is called + + +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 & 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 & 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 & 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{{{NEW_OBJ}}})" + )