use super::protocol::ProtocolFrame;
use super::Result;
#[derive(Clone)]
struct PendingData {
frame: ProtocolFrame,
send_time: u64,
retransmit_count: usize,
acknowledged: bool,
}
pub struct Publisher {
enable_nack: bool,
retransmit_timeout: core::time::Duration,
max_retransmits: usize,
next_seq_num: u32,
pending_data: alloc::vec::Vec<PendingData>,
}
impl Clone for Publisher {
fn clone(&self) -> Self {
Self {
enable_nack: self.enable_nack,
retransmit_timeout: self.retransmit_timeout,
max_retransmits: self.max_retransmits,
next_seq_num: self.next_seq_num,
pending_data: self.pending_data.clone(),
}
}
}
impl Publisher {
pub fn new(
enable_nack: bool,
retransmit_timeout: core::time::Duration,
max_retransmits: usize,
) -> Result<Self> {
Ok(Self {
enable_nack,
retransmit_timeout,
max_retransmits,
next_seq_num: 1, pending_data: alloc::vec::Vec::new(),
})
}
pub fn create_frame(&mut self, topic_id: u16, data: &[u8]) -> Result<ProtocolFrame> {
let seq_num = self.next_seq_num;
self.next_seq_num += 1;
let frame = ProtocolFrame::new_data_frame(seq_num, topic_id, data)?;
if self.enable_nack {
self.pending_data.push(PendingData {
frame: frame.clone(),
send_time: Self::get_current_time(),
retransmit_count: 0,
acknowledged: false,
});
}
Ok(frame)
}
pub fn handle_nack(&mut self, seq_num: u32, topic_id: u16) -> Result<Vec<ProtocolFrame>> {
if !self.enable_nack {
return Ok(Vec::new());
}
let mut frames_to_retransmit = Vec::new();
for pending in &mut self.pending_data {
if pending.frame.seq_num() == seq_num && pending.frame.topic_id() == topic_id {
if pending.retransmit_count < self.max_retransmits {
pending.retransmit_count += 1;
pending.send_time = Self::get_current_time();
frames_to_retransmit.push(pending.frame.clone());
} else {
pending.acknowledged = true;
}
break;
}
}
Ok(frames_to_retransmit)
}
pub fn check_timeouts(&mut self) -> Result<Vec<ProtocolFrame>> {
if !self.enable_nack {
return Ok(Vec::new());
}
let current_time = Self::get_current_time();
let timeout_ms = self.retransmit_timeout.as_millis() as u64;
let mut frames_to_retransmit = Vec::new();
for pending in &mut self.pending_data {
if !pending.acknowledged {
let elapsed = current_time - pending.send_time;
if elapsed > timeout_ms {
if pending.retransmit_count < self.max_retransmits {
pending.retransmit_count += 1;
pending.send_time = current_time;
frames_to_retransmit.push(pending.frame.clone());
} else {
pending.acknowledged = true;
}
}
}
}
self.pending_data.retain(|pending| !pending.acknowledged);
Ok(frames_to_retransmit)
}
pub fn create_heartbeat_frame(&self) -> ProtocolFrame {
ProtocolFrame::new_heartbeat_frame()
}
fn get_current_time() -> u64 {
#[cfg(feature = "posix")]
{
let now = std::time::SystemTime::now();
let duration = now
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or(std::time::Duration::ZERO);
return duration.as_millis() as u64;
}
#[cfg(feature = "baremetal")]
{
return 0u64;
}
#[cfg(not(any(feature = "posix", feature = "baremetal")))]
{
return 0u64;
}
}
pub fn clear_pending(&mut self) {
self.pending_data.clear();
}
pub fn pending_count(&self) -> usize {
self.pending_data.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::pubsub::protocol::FrameType;
#[test]
fn test_create_frame() {
let mut publisher =
Publisher::new(true, core::time::Duration::from_millis(100), 3).unwrap();
let frame = publisher.create_frame(0, b"test data").unwrap();
assert_eq!(frame.frame_type(), FrameType::Data);
assert_eq!(frame.seq_num(), 1);
assert_eq!(frame.topic_id(), 0);
assert_eq!(frame.payload(), b"test data");
}
#[test]
fn test_handle_nack() {
let mut publisher =
Publisher::new(true, core::time::Duration::from_millis(100), 3).unwrap();
let original_frame = publisher.create_frame(0, b"test data").unwrap();
let retransmit_frames = publisher
.handle_nack(original_frame.seq_num(), original_frame.topic_id())
.unwrap();
assert_eq!(retransmit_frames.len(), 1);
assert_eq!(retransmit_frames[0].seq_num(), original_frame.seq_num());
assert_eq!(retransmit_frames[0].topic_id(), original_frame.topic_id());
assert_eq!(retransmit_frames[0].payload(), original_frame.payload());
}
#[test]
fn test_check_timeouts() {
let mut publisher = Publisher::new(true, core::time::Duration::from_millis(1), 3).unwrap();
let original_frame = publisher.create_frame(0, b"test data").unwrap();
std::thread::sleep(std::time::Duration::from_millis(2));
let retransmit_frames = publisher.check_timeouts().unwrap();
assert_eq!(retransmit_frames.len(), 1);
assert_eq!(retransmit_frames[0].seq_num(), original_frame.seq_num());
}
#[test]
fn test_create_heartbeat_frame() {
let publisher = Publisher::new(false, core::time::Duration::from_millis(100), 3).unwrap();
let frame = publisher.create_heartbeat_frame();
assert_eq!(frame.frame_type(), FrameType::Heartbeat);
assert_eq!(frame.seq_num(), 0);
assert_eq!(frame.topic_id(), 0);
assert_eq!(frame.payload(), b"");
}
}