ashv2 10.0.2

Implementation of the Asynchronous Serial Host (ASH) protocol.
Documentation
use std::io;
use std::io::ErrorKind;
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use std::sync::atomic::Ordering::Relaxed;
use std::time::{Duration, Instant};

use async_serialport::Writer;
use log::{debug, error, info, trace, warn};
use tokio::sync::mpsc::{Receiver, WeakSender};

use self::buffer::Buffer;
use self::transmission::Transmission;
use crate::actor::message::Message;
use crate::frame::{Ack, Data, Error, Nak, RST, Rst, RstAck};
use crate::seq::Seq;
use crate::status::Status;
use crate::types::{MAX_FRAME_SIZE, Payload};
use crate::{T_RSTACK_MAX_MILLIS, T_RX_ACK_MAX_MILLIS, TX_K};

mod buffer;
mod transmission;

/// Maximum time to wait for RST ACK frame after sending RST frame.
const T_RSTACK_MAX: Duration = Duration::from_millis(T_RSTACK_MAX_MILLIS);

const T_RX_ACK_MAX: Duration = Duration::from_millis(T_RX_ACK_MAX_MILLIS);

const TRANSMITTER_CHANNEL_CLOSED: &str = "ASHv2 transmitter channel is closed";

/// `ASHv2` transmitter.
#[derive(Debug)]
pub struct Transmitter {
    buffer: Buffer<Writer>,
    messages: Receiver<Message>,
    requeue: WeakSender<Message>,
    status: Status,
    last_rst_sent: Option<Instant>,
    transmissions: heapless::Vec<Transmission, TX_K>,
    frame_number: Seq,
    ack_number: Seq,
}

impl Transmitter {
    /// Creates a new `ASHv2` transmitter.
    #[must_use]
    pub const fn new(
        writer: Writer,
        messages: Receiver<Message>,
        requeue: WeakSender<Message>,
    ) -> Self {
        Self {
            buffer: Buffer::new(writer),
            messages,
            requeue,
            status: Status::Uninitialized,
            last_rst_sent: None,
            transmissions: heapless::Vec::new(),
            frame_number: Seq::Zero,
            ack_number: Seq::Zero,
        }
    }
}

impl Transmitter {
    /// Runs the transmitter, processing messages from the channel.
    pub async fn run(mut self, running: Arc<AtomicBool>) {
        trace!("Starting transmitter with frame size: {MAX_FRAME_SIZE}");
        self.reset().await.unwrap_or_else(|error| {
            error!("Failed to send initial RST frame: {error}");
        });

        while let Some(message) = self.messages.recv().await {
            trace!("Received message: {message}");

            if matches!(message, Message::Terminate) {
                debug!("Terminating transmitter");
                break;
            }

            if let Err(error) = self.handle_message(message).await {
                error!("Resetting connection due to I/O error: {error}");
                self.status = Status::Failed;
            }
        }

        running.store(false, Relaxed);
        info!("Transmitter loop terminated.");
    }

    async fn handle_message(&mut self, message: Message) -> io::Result<()> {
        if self.status != Status::Connected {
            if let Message::RstAck(ack) = message {
                return self.handle_rst_ack(ack).await;
            }

            trace!("Received message before connection was established. Re-queueing.");
            self.requeue(message).await?;

            // Only log if the connection has failed, not if it hasn't been established yet.
            if self.status == Status::Failed {
                warn!("ASHv2 Connection failed. Resetting...");
            }

            return self.reset().await;
        }

        match message {
            Message::Payload {
                payload,
                response_tx: response,
            } => self.handle_payload(payload, response).await,
            Message::Ack(ack_num) => self.send_ack(ack_num).await,
            Message::Nak(ack_num) => self.send_nak(ack_num).await,
            Message::Rst(rst) => self.handle_rst(rst).await,
            Message::RstAck(rst_ack) => self.handle_rst_ack(rst_ack).await,
            Message::Error(error) => self.handle_error(error).await,
            Message::AckSentFrame(frame_num) => {
                self.ack_sent_frames(frame_num);
                Ok(())
            }
            Message::NakSentFrame(frame_num) => self.nak_sent_frames(frame_num).await,
            Message::Terminate => {
                error!(
                    "Termination signal received. This should have already been handed in the main loop ."
                );
                Ok(())
            }
        }
    }

    async fn handle_payload(
        &mut self,
        payload: Box<Payload>,
        response: tokio::sync::oneshot::Sender<io::Result<()>>,
    ) -> io::Result<()> {
        if self.transmissions.is_full() {
            warn!("Insufficient space in transmission queue for payload, requeueing.");
            return self
                .requeue(Message::Payload {
                    payload,
                    response_tx: response,
                })
                .await;
        }

        let data = Data::new(self.next_frame_number(), self.ack_number, *payload);
        // With a sliding windows size > 1 the NCP may enter an "ERROR: Assert" state when sending
        // fragmented messages if each DATA frame's ACK number is not increased.
        self.ack_number.increment();
        response
            .send(self.transmit(data.into()).await)
            .unwrap_or_else(|_| {
                error!("Failed to send transmit result through response channel.");
            });
        Ok(())
    }

    async fn send_ack(&mut self, ack_num: Seq) -> io::Result<()> {
        self.ack_number = ack_num;
        self.buffer.write_frame(Ack::new(ack_num, false)).await
    }

    async fn send_nak(&mut self, ack_num: Seq) -> io::Result<()> {
        self.buffer.write_frame(Nak::new(ack_num, false)).await
    }

    /// Handle RST frame received from the NCP.
    async fn handle_rst(&mut self, rst: Rst) -> io::Result<()> {
        error!("Received RST frame: {rst}, resetting connection.");
        self.status = Status::Failed;
        self.reset().await
    }

    /// Handle RST ACK frame received from the NCP.
    async fn handle_rst_ack(&mut self, rst_ack: RstAck) -> io::Result<()> {
        trace!("Received RST ACK frame: {rst_ack}, connection reset acknowledged.");

        if !rst_ack.is_ash_v2() {
            error!("Received RST ACK frame with invalid ASH version: {rst_ack}.");
            return Ok(());
        }

        if let Some(timestamp) = self.last_rst_sent.take() {
            if timestamp.elapsed() < T_RSTACK_MAX {
                debug!("Connection established successfully.");
                self.status = Status::Connected;
                Ok(())
            } else {
                warn!("RST ACK received after timeout. Resetting connection again.");
                self.reset().await
            }
        } else {
            warn!("Received unexpected RST ACK frame: {rst_ack}.");
            Ok(())
        }
    }

    /// Handle errors received from the NCP.
    async fn handle_error(&mut self, error: Error) -> io::Result<()> {
        warn!("Transmitter encountered error: {error}, resetting connection.");
        self.status = Status::Failed;
        self.reset().await
    }

    /// Remove `DATA` frames from the queue that have been acknowledged by the NCP.
    fn ack_sent_frames(&mut self, ack_num: Seq) {
        // Remove timed-out transmissions.
        self.transmissions
            .retain(|transmission| !transmission.is_timed_out(T_RX_ACK_MAX));

        // Remove acknowledged transmissions.
        while let Some(transmission) = self
            .transmissions
            .iter()
            .position(|transmission| transmission.frame_num().next() == ack_num)
            .map(|index| self.transmissions.remove(index))
        {
            trace!(
                "ACKed frame {transmission} after {:?}",
                transmission.elapsed()
            );
        }
    }

    /// Retransmit `DATA` frames that have been `NAK`ed by the NCP.
    async fn nak_sent_frames(&mut self, nak_num: Seq) -> io::Result<()> {
        // Remove timed-out transmissions.
        self.transmissions
            .retain(|transmission| !transmission.is_timed_out(T_RX_ACK_MAX));

        // Retransmit NAK'ed transmission.
        if let Some(transmission) = self
            .transmissions
            .iter()
            .position(|transmission| transmission.frame_num() == nak_num)
            .map(|index| self.transmissions.remove(index))
        {
            debug!("Retransmitting NAK'ed frame #{}", transmission.frame_num());
            self.transmit(transmission).await?;
        }

        Ok(())
    }

    /// Send a `DATA` frame.
    async fn transmit(&mut self, mut transmission: Transmission) -> io::Result<()> {
        let data = transmission.data_for_transmit()?;
        trace!("Transmitting frame {data:#04X}");
        self.buffer.write_frame(data).await?;
        self.transmissions
            .insert(0, transmission)
            .map_err(|_| io::Error::new(ErrorKind::OutOfMemory, "Failed to enqueue retransmit"))
    }

    /// Send RST frame to reset the connection.
    async fn reset(&mut self) -> io::Result<()> {
        if let Some(timestamp) = self.last_rst_sent.take()
            && timestamp.elapsed() < T_RSTACK_MAX
        {
            debug!("Last RST sent {timestamp:?} ago, waiting before sending another...");
            self.last_rst_sent.replace(timestamp);
            return Ok(());
        }

        self.last_rst_sent.replace(Instant::now());
        self.buffer.write_frame(RST).await
    }

    /// Returns the next frame number.
    pub const fn next_frame_number(&mut self) -> Seq {
        let frame_number = self.frame_number;
        self.frame_number.increment();
        frame_number
    }

    fn reject_message(message: Message, kind: ErrorKind, reason: &'static str) -> io::Result<()> {
        if let Message::Payload { response_tx, .. } = message {
            response_tx
                .send(Err(io::Error::new(kind, reason)))
                .unwrap_or_else(|_| {
                    error!("Failed to send transmit result through response channel.");
                });
        }

        Err(io::Error::new(kind, reason))
    }

    async fn requeue(&self, message: Message) -> io::Result<()> {
        let Some(sender) = self.requeue.upgrade() else {
            return Self::reject_message(
                message,
                ErrorKind::BrokenPipe,
                TRANSMITTER_CHANNEL_CLOSED,
            );
        };

        sender.send(message).await.map_err(|error| {
            let message = error.0;
            if let Message::Payload { response_tx, .. } = message {
                response_tx
                    .send(Err(io::Error::new(
                        ErrorKind::BrokenPipe,
                        TRANSMITTER_CHANNEL_CLOSED,
                    )))
                    .unwrap_or_else(|_| {
                        error!("Failed to send transmit result through response channel.");
                    });
            }

            io::Error::new(ErrorKind::BrokenPipe, TRANSMITTER_CHANNEL_CLOSED)
        })
    }
}