use std::collections::VecDeque;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use tokio::sync::{mpsc, oneshot};
use crate::error::Error;
use crate::link::{CycleOutcome, Link};
use crate::protocol::{Cmd, RX_FRAME_BYTES, RxFrame, Seq, TX_FRAME_BYTES, TxFrame};
use crate::response::Response;
use autd3_cpu_wire::Mode;
use autd3_cpu_wire::payload::SetModePayload;
use zerocopy::FromBytes;
use super::completion::CompletionSender;
use super::config::{ClientConfig, MAX_DEVICES};
use super::pool::Slot;
pub(super) struct CmdMessage {
pub(super) frame: Slot,
pub(super) response_tx: CompletionSender,
pub(super) exclusive: bool,
}
struct Inflight {
seq: Seq,
frame: Slot,
acked: u128,
age: u32,
exclusive: bool,
response_tx: CompletionSender,
}
fn stage_frame(seq: Seq, frame: &Slot, tx_bufs: &mut [[u8; TX_FRAME_BYTES]]) {
for (device, buf) in tx_bufs.iter_mut().enumerate() {
buf[0] = seq.get();
buf[1] = frame.cmd_for(device).as_u8();
buf[2..].copy_from_slice(frame.payload_for(device));
}
}
enum HeadAction {
None,
Reset,
GiveUp,
}
#[derive(Default)]
struct ResyncState {
active: bool,
rounds: u32,
reset_tried: bool,
}
impl ResyncState {
fn reset(&mut self) {
*self = Self::default();
}
fn on_ack_progress(&mut self, pending: &mut VecDeque<Inflight>) {
if self.active {
tracing::debug!("ack progress during resync");
self.rounds = 0;
self.reset_tried = false;
if let Some(head) = pending.front_mut() {
head.age = 0;
}
}
}
fn advance_head(
&mut self,
pending: &mut VecDeque<Inflight>,
config: &ClientConfig,
) -> HeadAction {
let Some(head) = pending.front_mut() else {
if self.active {
self.reset();
tracing::debug!("resync complete");
}
return HeadAction::None;
};
head.age = head.age.saturating_add(1);
if head.age < config.timeout_cycles {
return HeadAction::None;
}
head.age = 0;
if !self.active {
self.active = true;
self.rounds = 0;
tracing::debug!(
seq = head.seq.get(),
"head frame unacked past timeout; entering resync"
);
return HeadAction::None;
}
self.rounds += 1;
if self.rounds < config.max_resync_rounds.get() {
return HeadAction::None;
}
self.rounds = 0;
if self.reset_tried {
HeadAction::GiveUp
} else {
self.reset_tried = true;
tracing::warn!(
seq = head.seq.get(),
"resync rounds exhausted; resetting sequence"
);
HeadAction::Reset
}
}
}
pub(super) fn run_rt_thread<L: Link>(
link: L,
cmd_rx: mpsc::Receiver<CmdMessage>,
config: ClientConfig,
hs_done_tx: oneshot::Sender<Result<(), String>>,
closed: Arc<AtomicBool>,
) {
autd3_rs_core::apply_thread_tuning(autd3_rs_core::RtThreadTuning {
priority: config.rt_priority,
policy: config.rt_policy,
affinity: config.rt_affinity,
});
let mut rt = RtThread::new(link, cmd_rx, config, closed);
match rt.handshake() {
Ok(()) => {}
Err(e) => {
let _ = hs_done_tx.send(Err(e));
return;
}
}
if hs_done_tx.send(Ok(())).is_err() {
return;
}
rt.run();
}
struct RtThread<L: Link> {
link: L,
cmd_rx: mpsc::Receiver<CmdMessage>,
config: ClientConfig,
closed: Arc<AtomicBool>,
all_acked: u128,
tx_bufs: Vec<[u8; TX_FRAME_BYTES]>,
rx_bufs: Vec<[u8; RX_FRAME_BYTES]>,
next_seq: Seq,
pending: VecDeque<Inflight>,
held_exclusive: Option<CmdMessage>,
resync: ResyncState,
stale_run: u32,
reset_remaining: u32,
stale_limit: u32,
}
enum StageOutcome {
Staged,
Disconnected,
}
impl<L: Link> RtThread<L> {
fn new(
link: L,
cmd_rx: mpsc::Receiver<CmdMessage>,
config: ClientConfig,
closed: Arc<AtomicBool>,
) -> Self {
let num_devices = link.num_devices();
let all_acked: u128 = if num_devices == MAX_DEVICES {
u128::MAX
} else {
(1u128 << num_devices) - 1
};
let stale_limit = config
.timeout_cycles
.saturating_mul(config.max_resync_rounds.get());
Self {
link,
cmd_rx,
pending: VecDeque::with_capacity(config.max_inflight.get()),
held_exclusive: None,
config,
closed,
all_acked,
tx_bufs: vec![[0u8; TX_FRAME_BYTES]; num_devices],
rx_bufs: vec![[0u8; RX_FRAME_BYTES]; num_devices],
next_seq: Seq::ZERO,
resync: ResyncState::default(),
stale_run: 0,
reset_remaining: 0,
stale_limit,
}
}
fn handshake(&mut self) -> Result<(), String> {
tracing::debug!(
cycles = self.config.reset_resend_cycles.get(),
low_latency = self.config.low_latency,
"starting handshake"
);
for buf in &mut self.tx_bufs {
TxFrame::new(Seq::ZERO, Cmd::Reset).write_to(buf);
}
for _ in 0..self.config.reset_resend_cycles.get() {
self.link
.cycle(&self.tx_bufs, &mut self.rx_bufs)
.map_err(|e| format!("handshake failed: {e}"))?;
}
self.next_seq = if self.config.low_latency {
self.negotiate_low_latency()?
} else {
Seq::ZERO
};
Ok(())
}
fn negotiate_low_latency(&mut self) -> Result<Seq, String> {
let mut frame = TxFrame::new(Seq::ZERO, Cmd::SetMode);
let (p, _) = SetModePayload::mut_from_prefix(&mut frame.payload).unwrap();
p.mode = Mode::LowLatency.as_u8();
for buf in &mut self.tx_bufs {
frame.write_to(buf);
}
let bound = self.config.timeout_cycles.max(2);
for _ in 0..bound {
let CycleOutcome { rx_valid } = self
.link
.cycle(&self.tx_bufs, &mut self.rx_bufs)
.map_err(|e| format!("handshake failed: {e}"))?;
if rx_valid && self.rx_bufs.iter().all(|rx| rx[0] == Seq::ZERO.get()) {
tracing::info!("low-latency mode established");
return Ok(Seq::new(1));
}
}
tracing::warn!("low-latency negotiation failed; staying in FIFO mode");
Ok(Seq::ZERO)
}
fn run(&mut self) {
let mut link_error = None;
loop {
if self.closed.load(Ordering::Acquire) {
break;
}
self.link.wait_next_cycle();
if matches!(self.stage_tx(), StageOutcome::Disconnected) {
break;
}
let rx_valid = match self.link.cycle(&self.tx_bufs, &mut self.rx_bufs) {
Ok(CycleOutcome { rx_valid }) => rx_valid,
Err(e) => {
tracing::error!("link cycle failed: {e}");
link_error = Some(format!("link cycle failed: {e}"));
break;
}
};
if self.reset_remaining > 0 {
self.advance_reset_phase();
} else if rx_valid {
self.handle_healthy();
} else {
self.handle_stale();
}
}
self.teardown(link_error.as_deref());
}
fn stage_tx(&mut self) -> StageOutcome {
if self.reset_remaining > 0 {
tracing::trace!(remaining = self.reset_remaining, "staging reset frame");
for buf in &mut self.tx_bufs {
TxFrame::new(Seq::ZERO, Cmd::Reset).write_to(buf);
}
return StageOutcome::Staged;
}
if self.resync.active {
if let Some(front) = self.pending.front() {
tracing::trace!(seq = front.seq.get(), "retransmitting head frame");
stage_frame(front.seq, &front.frame, &mut self.tx_bufs);
}
return StageOutcome::Staged;
}
if self.held_exclusive.is_some() {
if self.pending.is_empty() {
let msg = self.held_exclusive.take().expect("just checked is_some");
self.stage_new(msg);
}
return StageOutcome::Staged;
}
let exclusive_inflight = self.pending.front().is_some_and(|entry| entry.exclusive);
if !exclusive_inflight && self.pending.len() < self.config.max_inflight.get() {
match self.cmd_rx.try_recv() {
Ok(msg) if msg.exclusive && !self.pending.is_empty() => {
tracing::trace!("holding exclusive frame until pending drains");
self.held_exclusive = Some(msg);
}
Ok(msg) => self.stage_new(msg),
Err(mpsc::error::TryRecvError::Empty) => {}
Err(mpsc::error::TryRecvError::Disconnected) => return StageOutcome::Disconnected,
}
}
StageOutcome::Staged
}
fn stage_new(&mut self, msg: CmdMessage) {
let seq = self.next_seq;
self.next_seq = self.next_seq.next();
tracing::trace!(
seq = seq.get(),
cmd = ?msg.frame.cmd_for(0),
exclusive = msg.exclusive,
"staged frame"
);
stage_frame(seq, &msg.frame, &mut self.tx_bufs);
self.pending.push_back(Inflight {
seq,
frame: msg.frame,
acked: 0,
age: 0,
exclusive: msg.exclusive,
response_tx: msg.response_tx,
});
}
fn advance_reset_phase(&mut self) {
self.reset_remaining -= 1;
if self.reset_remaining == 0 {
let mut seq = Seq::ZERO;
for entry in &mut self.pending {
entry.seq = seq;
entry.acked = 0;
entry.age = 0;
seq = seq.next();
}
self.next_seq = seq;
self.resync.active = !self.pending.is_empty();
self.resync.rounds = 0;
tracing::debug!(
pending = self.pending.len(),
"sequence reset complete; replaying pending frames"
);
}
}
fn handle_healthy(&mut self) {
self.stale_run = 0;
self.route_acks();
match self.resync.advance_head(&mut self.pending, &self.config) {
HeadAction::None => {}
HeadAction::Reset => self.reset_remaining = self.config.reset_resend_cycles.get(),
HeadAction::GiveUp => {
tracing::warn!(
pending = self.pending.len(),
"sequence reset did not recover; failing pending frames with timeout"
);
self.fail_pending_timeout();
self.resync.reset();
}
}
}
fn handle_stale(&mut self) {
self.stale_run = self.stale_run.saturating_add(1);
tracing::trace!(run = self.stale_run, "no valid rx this cycle");
if self.stale_run >= self.stale_limit {
if !self.pending.is_empty() {
tracing::warn!(
pending = self.pending.len(),
cycles = self.stale_run,
"no valid rx from link; failing pending frames with timeout"
);
}
self.fail_pending_timeout();
self.resync.reset();
self.stale_run = 0;
}
}
fn route_acks(&mut self) {
if self.pending.is_empty() {
return;
}
let front_seq = self.pending.front().expect("non-empty").seq;
let back_seq = self.pending.back().expect("non-empty").seq;
let span = back_seq.distance_from(front_seq) as usize;
for (device, rx_buf) in self.rx_bufs.iter().enumerate() {
let rx = RxFrame::parse(rx_buf);
let ack_offset = rx.ack.distance_from(front_seq) as usize;
if ack_offset > span {
continue;
}
let bit = 1u128 << device;
for entry in self.pending.iter_mut().take(ack_offset + 1) {
if entry.acked & bit == 0 {
tracing::trace!(device, seq = entry.seq.get(), "device acked frame");
entry.acked |= bit;
entry.frame.record_data(device, rx.data);
}
}
}
let mut progressed = false;
while self
.pending
.front()
.is_some_and(|entry| entry.acked == self.all_acked)
{
let entry = self.pending.pop_front().expect("just checked");
tracing::trace!(seq = entry.seq.get(), "frame acked by all devices");
entry
.response_tx
.send(Ok(Response::from_slice(entry.frame.data())));
progressed = true;
}
if progressed {
self.resync.on_ack_progress(&mut self.pending);
}
}
fn fail_pending_timeout(&mut self) {
for entry in self.pending.drain(..) {
entry.response_tx.send(Err(Error::Timeout {
cycles: self.config.timeout_cycles,
}));
}
}
fn teardown(&mut self, link_error: Option<&str>) {
tracing::debug!(pending = self.pending.len(), "RT thread stopping");
let cause = || link_error.map_or(Error::RtClosed, |msg| Error::Link(msg.to_owned()));
self.cmd_rx.close();
if let Some(msg) = self.held_exclusive.take() {
msg.response_tx.send(Err(cause()));
}
for entry in self.pending.drain(..) {
entry.response_tx.send(Err(cause()));
}
while let Ok(msg) = self.cmd_rx.try_recv() {
msg.response_tx.send(Err(cause()));
}
}
}