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;
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";
#[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 {
#[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 {
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?;
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);
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
}
async fn handle_rst(&mut self, rst: Rst) -> io::Result<()> {
error!("Received RST frame: {rst}, resetting connection.");
self.status = Status::Failed;
self.reset().await
}
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(())
}
}
async fn handle_error(&mut self, error: Error) -> io::Result<()> {
warn!("Transmitter encountered error: {error}, resetting connection.");
self.status = Status::Failed;
self.reset().await
}
fn ack_sent_frames(&mut self, ack_num: Seq) {
self.transmissions
.retain(|transmission| !transmission.is_timed_out(T_RX_ACK_MAX));
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()
);
}
}
async fn nak_sent_frames(&mut self, nak_num: Seq) -> io::Result<()> {
self.transmissions
.retain(|transmission| !transmission.is_timed_out(T_RX_ACK_MAX));
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(())
}
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"))
}
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
}
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)
})
}
}