use embassy_sync::{
blocking_mutex::raw::CriticalSectionRawMutex,
pubsub::WaitResult,
watch::{Receiver, Sender, Watch},
};
use pidsk_controller::PdGainsf32;
use static_cell::StaticCell;
use imu_sensors::{AccFullScale, AccUnits, GyroFullScale, GyroUnits, ImuDevice, ImuError};
use motor_mixers::MotorMixerMessage;
use sensor_fusion::{MadgwickFilterf32, SensorFusion};
use simple_bitset::BitSet64;
#[cfg(feature = "rpm_filters")]
use motor_mixers::RpmNotchFilterBankConfig;
use crate::{
boards::targets::BoardImu,
config::{FastConfigItem, FastConfigSubscriber, fast_config_subscriber},
flight::{FilterAccGyro, FlightController, ImuFilterBank, ImuFilterBankConfig, RcControls, VehicleControl},
tasks::{
GyroPidMessage, SetpointMessage,
motor_mixer::MOTOR_MIXER_SIGNAL,
rx::{RxMessageReceiver, rx_message_receiver},
},
};
#[cfg(feature = "gps")]
use crate::tasks::gps::GPS_YAW_HEADING_SIGNAL;
const GYRO_PID_WATCH_COUNT: usize = 3;
static GYRO_PID_WATCH: Watch<CriticalSectionRawMutex, GyroPidMessage, GYRO_PID_WATCH_COUNT> = Watch::new();
type GyroPidSender = Sender<'static, CriticalSectionRawMutex, GyroPidMessage, GYRO_PID_WATCH_COUNT>;
pub fn gyro_pid_sender() -> GyroPidSender {
GYRO_PID_WATCH.sender()
}
#[allow(unused)]
pub type GyroPidReceiver = Receiver<'static, CriticalSectionRawMutex, GyroPidMessage, GYRO_PID_WATCH_COUNT>;
#[allow(unused)]
#[allow(clippy::expect_used)]
pub fn gyro_pid_receiver() -> GyroPidReceiver {
GYRO_PID_WATCH.receiver().expect("gyro_pid receiver failed")
}
const SETPOINT_WATCH_COUNT: usize = 3;
static SETPOINT_WATCH: Watch<CriticalSectionRawMutex, SetpointMessage, SETPOINT_WATCH_COUNT> = Watch::new();
type SetpointSender = Sender<'static, CriticalSectionRawMutex, SetpointMessage, SETPOINT_WATCH_COUNT>;
pub fn setpoint_sender() -> SetpointSender {
SETPOINT_WATCH.sender()
}
pub type SetpointReceiver = Receiver<'static, CriticalSectionRawMutex, SetpointMessage, SETPOINT_WATCH_COUNT>;
#[allow(unused)]
#[allow(clippy::expect_used)]
pub fn setpoint_receiver() -> SetpointReceiver {
SETPOINT_WATCH.receiver().expect("setpoint receiver failed")
}
static GYRO_PID_CTX: StaticCell<GyroPidContext<BoardImu>> = StaticCell::new();
#[allow(unused)]
pub struct GyroPidContext<I: ImuDevice> {
pub imu: I,
pub rx_receiver: RxMessageReceiver,
pub gyro_pid_sender: GyroPidSender,
pub setpoint_sender: SetpointSender,
pub fast_config_subscriber: FastConfigSubscriber,
pub imu_filters: ImuFilterBank,
pub sensor_fusion: MadgwickFilterf32,
pub flight_controller: FlightController,
pub rc_controls: RcControls,
pub rc_modes: BitSet64,
pub gyro_pid_send_count: u32,
pub gyro_pid_denominator: u32,
}
pub fn init(
imu: BoardImu,
imu_filter_bank_config: ImuFilterBankConfig,
#[cfg(feature = "rpm_filters")] rpm_notch_filter_bank_config: RpmNotchFilterBankConfig,
#[cfg(feature = "rpm_filters")] looptime_seconds: f32,
) -> &'static mut GyroPidContext<BoardImu> {
let ctx = GyroPidContext {
imu,
rx_receiver: rx_message_receiver(),
gyro_pid_sender: gyro_pid_sender(),
setpoint_sender: setpoint_sender(),
fast_config_subscriber: fast_config_subscriber(),
#[cfg(feature = "rpm_filters")]
imu_filters: ImuFilterBank::with_config_and_notch(
imu_filter_bank_config,
rpm_notch_filter_bank_config,
looptime_seconds,
),
#[cfg(not(feature = "rpm_filters"))]
imu_filters: ImuFilterBank::with_config(imu_filter_bank_config),
sensor_fusion: MadgwickFilterf32::new(),
flight_controller: FlightController::new(),
rc_controls: RcControls::new(),
rc_modes: BitSet64::new(),
gyro_pid_send_count: 0,
gyro_pid_denominator: 10,
};
GYRO_PID_CTX.init(ctx)
}
#[embassy_executor::task]
pub async fn run(ctx: &'static mut GyroPidContext<BoardImu>) {
log::info!(" GYRO_PID: task started");
let sample_rates = ctx.imu.init(8000, GyroFullScale::Max, GyroUnits::Rps, AccFullScale::Max, AccUnits::G).await;
let (gyro_rate_hz, _acc_rate_hz) = match sample_rates {
Ok((gyro_rate_hz, acc_rate_hz)) => (gyro_rate_hz, acc_rate_hz),
Err(_err) => (1000, 1000),
};
#[allow(clippy::cast_precision_loss)]
let delta_t: f32 = 1.0 / (gyro_rate_hz as f32);
let mut ticker = embassy_time::Ticker::every(embassy_time::Duration::from_hz(u64::from(gyro_rate_hz)));
let mut loop_count: u32 = 0;
loop {
ticker.next().await;
if let Err(err) = gyro_pid_loop_iteration(ctx, delta_t).await {
log::error!("GYRO_PID iteration failed: {err:?}");
}
if loop_count.is_multiple_of(1000) {
log::info!(" GYRO_PID: loop {loop_count}");
}
loop_count = loop_count.wrapping_add(1);
}
}
async fn gyro_pid_loop_iteration(ctx: &mut GyroPidContext<BoardImu>, delta_t: f32) -> Result<(), ImuError> {
let (acc, gyro_rps) = ctx.imu.read_acc_gyro().await?;
let gyro_rps_unfiltered = gyro_rps;
let (acc, gyro_rps) = ctx.imu_filters.update(acc, gyro_rps, delta_t);
#[cfg(feature = "gps")]
if let Some(gps_yaw_heading) = GPS_YAW_HEADING_SIGNAL.try_take() {
_ = ctx.sensor_fusion.correct_yaw(gps_yaw_heading.yaw_heading_radians, gps_yaw_heading.delta_t);
}
let orientation = ctx.sensor_fusion.fuse_acc_gyro(acc, gyro_rps, delta_t);
if let Some(rx_message) = ctx.rx_receiver.try_changed() {
ctx.rc_controls = rx_message.rc_controls;
ctx.rc_modes = rx_message.rc_modes;
}
let (motor_commands, setpoints_updated) =
ctx.flight_controller.calculate_motor_commands(gyro_rps, orientation, delta_t, ctx.rc_controls, ctx.rc_modes);
MOTOR_MIXER_SIGNAL.signal(MotorMixerMessage::from(motor_commands));
ctx.gyro_pid_send_count += 1;
#[cfg(not(any(feature = "blackbox", feature = "osd")))]
{
_ = gyro_rps_unfiltered;
_ = setpoints_updated;
}
#[cfg(any(feature = "blackbox", feature = "osd"))]
if ctx.gyro_pid_send_count >= ctx.gyro_pid_denominator {
ctx.gyro_pid_send_count = 0;
let roll_errors = ctx.flight_controller.roll.rate_pid.error();
let pitch_errors = ctx.flight_controller.pitch.rate_pid.error();
let yaw_errors = ctx.flight_controller.yaw_rate_pid.error();
let pid_errors_p = [roll_errors.p, pitch_errors.p, yaw_errors.p];
let pid_errors_i = [roll_errors.i, pitch_errors.i, yaw_errors.i];
let pid_errors_d = [roll_errors.d, pitch_errors.d];
let time_us = embassy_time::Instant::now().as_micros();
let gyro_pid_message = GyroPidMessage {
orientation,
motor_commands,
acc,
gyro_rps,
gyro_rps_unfiltered,
pid_errors_p,
pid_errors_i,
pid_errors_d,
time_us,
};
ctx.gyro_pid_sender.send(gyro_pid_message);
if setpoints_updated {
let pid_errors_s = [roll_errors.s, pitch_errors.s, yaw_errors.s];
let pid_errors_k = [roll_errors.k, pitch_errors.k, yaw_errors.k];
let setpoint_message = SetpointMessage {
time_us,
rc_modes: ctx.rc_modes,
setpoints: [
ctx.rc_controls.roll_stick_dps,
ctx.rc_controls.pitch_stick_dps,
ctx.rc_controls.yaw_stick_dps,
ctx.rc_controls.throttle_stick,
],
pid_errors_s,
pid_errors_k,
rc_commands: ctx.rc_controls.controls_pwm,
#[cfg(feature = "dshot_telemetry")]
motor_rpm_d2: [0i16; SetpointMessage::MAX_SUPPORTED_MOTOR_COUNT],
#[cfg(feature = "servos")]
servos: [0i16; SetpointMessage::MAX_SUPPORTED_SERVO_COUNT],
gps_state_flags: 0,
failsafe_phase: 0,
rx_signal_received: true,
rx_flight_channel_is_valid: true,
};
ctx.setpoint_sender.send(setpoint_message);
}
}
if let Some(wait_result) = ctx.fast_config_subscriber.try_next_message()
&& let WaitResult::Message(fast_config_item) = wait_result
{
adjust_pid_gains(&mut ctx.flight_controller, fast_config_item);
}
Ok(())
}
fn adjust_pid_gains(flight_controller: &mut FlightController, fast_config_item: FastConfigItem) {
match fast_config_item {
FastConfigItem::RollRate(pid_config) => {
let gains = FlightController::calculate_gains(pid_config);
flight_controller.roll.rate_pid.set_gains(gains);
flight_controller.roll.rate_pid.switch_integration_off();
flight_controller.roll.rate_pid.set_setpoint(0.0);
}
FastConfigItem::PitchRate(pid_config) => {
let gains = FlightController::calculate_gains(pid_config);
flight_controller.pitch.rate_pid.set_gains(gains);
flight_controller.pitch.rate_pid.switch_integration_off();
flight_controller.pitch.rate_pid.set_setpoint(0.0);
}
FastConfigItem::YawRate(pid_config) => {
let gains = FlightController::calculate_gains(pid_config);
flight_controller.yaw_rate_pid.set_gains(gains);
flight_controller.yaw_rate_pid.switch_integration_off();
flight_controller.yaw_rate_pid.set_setpoint(0.0);
}
FastConfigItem::RollAngle(pid_config) => {
let gains = FlightController::calculate_gains(pid_config);
let pd_gains = PdGainsf32 { kp: gains.kp, kd: gains.kd };
flight_controller.roll.angle_pid.set_gains(pd_gains);
flight_controller.roll.angle_pid.set_setpoint(0.0);
}
FastConfigItem::PitchAngle(pid_config) => {
let gains = FlightController::calculate_gains(pid_config);
let pd_gains = PdGainsf32 { kp: gains.kp, kd: gains.kd };
flight_controller.pitch.angle_pid.set_gains(pd_gains);
flight_controller.pitch.angle_pid.set_setpoint(0.0);
}
}
}