diff --git a/esphome/components/stepper/__init__.py b/esphome/components/stepper/__init__.py index 8e80187662..017fd95cd8 100644 --- a/esphome/components/stepper/__init__.py +++ b/esphome/components/stepper/__init__.py @@ -1,3 +1,6 @@ +from collections.abc import Callable +from typing import Any + from esphome import automation import esphome.codegen as cg import esphome.config_validation as cv @@ -11,18 +14,13 @@ from esphome.const import ( CONF_TARGET, ) from esphome.core import CORE, CoroPriority, coroutine_with_priority +from esphome.types import SafeExpType IS_PLATFORM_COMPONENT = True stepper_ns = cg.esphome_ns.namespace("stepper") Stepper = stepper_ns.class_("Stepper") -SetTargetAction = stepper_ns.class_("SetTargetAction", automation.Action) -ReportPositionAction = stepper_ns.class_("ReportPositionAction", automation.Action) -SetSpeedAction = stepper_ns.class_("SetSpeedAction", automation.Action) -SetAccelerationAction = stepper_ns.class_("SetAccelerationAction", automation.Action) -SetDecelerationAction = stepper_ns.class_("SetDecelerationAction", automation.Action) - def validate_acceleration(value): value = cv.string(value) @@ -90,99 +88,53 @@ async def register_stepper(var, config): await setup_stepper_core_(var, config) -@automation.register_action( - "stepper.set_target", - SetTargetAction, - cv.Schema( - { - cv.Required(CONF_ID): cv.use_id(Stepper), - cv.Required(CONF_TARGET): cv.templatable(cv.int_), - } - ), - synchronous=True, +def _register_stepper_action( + name: str, + key: str, + validator: Callable[[Any], Any], + target: str, + type_: SafeExpType, + *extra: automation.ApplyCall, +) -> None: + automation.register_apply_action( + f"stepper.{name}", + cv.Schema( + { + cv.Required(CONF_ID): cv.use_id(Stepper), + cv.Required(key): cv.templatable(validator), + } + ), + automation.ApplyField(key, target, type_), + *extra, + ) + + +_register_stepper_action("set_target", CONF_TARGET, cv.int_, "set_target", cg.int32) +_register_stepper_action( + "report_position", CONF_POSITION, cv.int_, "report_position", cg.int32 ) -async def stepper_set_target_to_code(config, action_id, template_arg, args): - paren = await cg.get_variable(config[CONF_ID]) - var = cg.new_Pvariable(action_id, template_arg, paren) - template_ = await cg.templatable(config[CONF_TARGET], args, cg.int32) - cg.add(var.set_target(template_)) - return var - - -@automation.register_action( - "stepper.report_position", - ReportPositionAction, - cv.Schema( - { - cv.Required(CONF_ID): cv.use_id(Stepper), - cv.Required(CONF_POSITION): cv.templatable(cv.int_), - } - ), - synchronous=True, +_register_stepper_action( + "set_speed", + CONF_SPEED, + validate_speed, + "set_max_speed", + cg.float_, + automation.ApplyCall("on_update_speed()"), ) -async def stepper_report_position_to_code(config, action_id, template_arg, args): - paren = await cg.get_variable(config[CONF_ID]) - var = cg.new_Pvariable(action_id, template_arg, paren) - template_ = await cg.templatable(config[CONF_POSITION], args, cg.int32) - cg.add(var.set_position(template_)) - return var - - -@automation.register_action( - "stepper.set_speed", - SetSpeedAction, - cv.Schema( - { - cv.Required(CONF_ID): cv.use_id(Stepper), - cv.Required(CONF_SPEED): cv.templatable(validate_speed), - } - ), - synchronous=True, +_register_stepper_action( + "set_acceleration", + CONF_ACCELERATION, + validate_acceleration, + "set_acceleration", + cg.float_, ) -async def stepper_set_speed_to_code(config, action_id, template_arg, args): - paren = await cg.get_variable(config[CONF_ID]) - var = cg.new_Pvariable(action_id, template_arg, paren) - template_ = await cg.templatable(config[CONF_SPEED], args, cg.float_) - cg.add(var.set_speed(template_)) - return var - - -@automation.register_action( - "stepper.set_acceleration", - SetAccelerationAction, - cv.Schema( - { - cv.Required(CONF_ID): cv.use_id(Stepper), - cv.Required(CONF_ACCELERATION): cv.templatable(validate_acceleration), - } - ), - synchronous=True, +_register_stepper_action( + "set_deceleration", + CONF_DECELERATION, + validate_acceleration, + "set_deceleration", + cg.float_, ) -async def stepper_set_acceleration_to_code(config, action_id, template_arg, args): - paren = await cg.get_variable(config[CONF_ID]) - var = cg.new_Pvariable(action_id, template_arg, paren) - template_ = await cg.templatable(config[CONF_ACCELERATION], args, cg.float_) - cg.add(var.set_acceleration(template_)) - return var - - -@automation.register_action( - "stepper.set_deceleration", - SetDecelerationAction, - cv.Schema( - { - cv.Required(CONF_ID): cv.use_id(Stepper), - cv.Required(CONF_DECELERATION): cv.templatable(validate_acceleration), - } - ), - synchronous=True, -) -async def stepper_set_deceleration_to_code(config, action_id, template_arg, args): - paren = await cg.get_variable(config[CONF_ID]) - var = cg.new_Pvariable(action_id, template_arg, paren) - template_ = await cg.templatable(config[CONF_DECELERATION], args, cg.float_) - cg.add(var.set_deceleration(template_)) - return var @coroutine_with_priority(CoroPriority.CORE) diff --git a/esphome/components/stepper/stepper.h b/esphome/components/stepper/stepper.h index 06ef3bab37..8a8fbcd896 100644 --- a/esphome/components/stepper/stepper.h +++ b/esphome/components/stepper/stepper.h @@ -1,7 +1,6 @@ #pragma once #include "esphome/core/component.h" -#include "esphome/core/automation.h" namespace esphome::stepper { @@ -37,74 +36,4 @@ class Stepper { uint32_t last_step_{0}; }; -template class SetTargetAction final : public Action { - public: - explicit SetTargetAction(Stepper *parent) : parent_(parent) {} - - TEMPLATABLE_VALUE(int32_t, target) - - void play(const Ts &...x) override { this->parent_->set_target(this->target_.value(x...)); } - - protected: - Stepper *parent_; -}; - -template class ReportPositionAction final : public Action { - public: - explicit ReportPositionAction(Stepper *parent) : parent_(parent) {} - - TEMPLATABLE_VALUE(int32_t, position) - - void play(const Ts &...x) override { this->parent_->report_position(this->position_.value(x...)); } - - protected: - Stepper *parent_; -}; - -template class SetSpeedAction final : public Action { - public: - explicit SetSpeedAction(Stepper *parent) : parent_(parent) {} - - TEMPLATABLE_VALUE(float, speed); - - void play(const Ts &...x) override { - float speed = this->speed_.value(x...); - this->parent_->set_max_speed(speed); - this->parent_->on_update_speed(); - } - - protected: - Stepper *parent_; -}; - -template class SetAccelerationAction final : public Action { - public: - explicit SetAccelerationAction(Stepper *parent) : parent_(parent) {} - - TEMPLATABLE_VALUE(float, acceleration); - - void play(const Ts &...x) override { - float acceleration = this->acceleration_.value(x...); - this->parent_->set_acceleration(acceleration); - } - - protected: - Stepper *parent_; -}; - -template class SetDecelerationAction final : public Action { - public: - explicit SetDecelerationAction(Stepper *parent) : parent_(parent) {} - - TEMPLATABLE_VALUE(float, deceleration); - - void play(const Ts &...x) override { - float deceleration = this->deceleration_.value(x...); - this->parent_->set_deceleration(deceleration); - } - - protected: - Stepper *parent_; -}; - } // namespace esphome::stepper diff --git a/tests/components/stepper/common.yaml b/tests/components/stepper/common.yaml index fcf5759618..ba70fb7979 100644 --- a/tests/components/stepper/common.yaml +++ b/tests/components/stepper/common.yaml @@ -25,3 +25,12 @@ switch: - stepper.report_position: id: test_stepper position: 0 + - stepper.set_speed: + id: test_stepper + speed: 300 steps/s + - stepper.set_acceleration: + id: test_stepper + acceleration: !lambda return 150.0f; + - stepper.set_deceleration: + id: test_stepper + deceleration: 250 steps/s^2