diff --git a/examples/calculator.rs b/examples/calculator.rs index 63ca26d..91e7377 100644 --- a/examples/calculator.rs +++ b/examples/calculator.rs @@ -1,9 +1,12 @@ use std::{sync::mpsc::Receiver, thread, time::Duration}; -use seq_ex::sync::{MpscGuard, MpscSeqEx, MpscTransport, PacketType, RecvSuccess}; +use seq_ex::{ + sync::{MpscSeqEx, MpscTransport}, + Packet, +}; #[derive(Clone)] -enum Packet { +enum Payload { Add(f32), Sub(f32), Mul(f32), @@ -16,28 +19,20 @@ fn drop_packet() -> bool { rand_core::OsRng.next_u32() & 1 > 0 } -fn process(_: MpscGuard<'_, Packet>, recv_packet: Packet, _: Option, value: &mut f32) { - use Packet::*; - match recv_packet { - Add(n) => *value = *value + n, - Sub(n) => *value = *value - n, - Mul(n) => *value = *value * n, - Div(n) => *value = *value / n, - Mod(n) => *value = *value % n, - } -} - -fn receive(recv: &Receiver>, seq: &MpscSeqEx, transport: &MpscTransport, value: &mut f32) { +fn receive(recv: &Receiver>, seq: &MpscSeqEx, transport: &MpscTransport, value: &mut f32) { while let Ok(packet) = recv.try_recv() { - if !drop_packet() { - match packet { - PacketType::Ack(reply_no) => { - let _ = seq.receive_ack(reply_no); - } - PacketType::Payload(seq_no, reply_no, payload) => { - for RecvSuccess { guard, packet, send_data } in seq.receive_all(transport, seq_no, reply_no, payload) { - process(guard, packet, send_data, value); - } + if drop_packet() { + continue; + } + for recv_data in seq.receive_all(transport, packet) { + use Payload::*; + if let Some((_, recv_packet)) = recv_data.consume().0 { + match recv_packet { + Add(n) => *value = *value + n, + Sub(n) => *value = *value - n, + Mul(n) => *value = *value * n, + Div(n) => *value = *value / n, + Mod(n) => *value = *value % n, } } } @@ -52,15 +47,15 @@ fn main() { let mut value = 0.0; let mut remote_value = value; - seq1.send(&transport1, Packet::Add(1.0)); + seq1.send(&transport1, Payload::Add(1.0)); value += 1.0; - seq1.send(&transport1, Packet::Sub(2.0)); + seq1.send(&transport1, Payload::Sub(2.0)); value -= 2.0; - seq1.send(&transport1, Packet::Mul(3.0)); + seq1.send(&transport1, Payload::Mul(3.0)); value *= 3.0; - seq1.send(&transport1, Packet::Div(4.0)); + seq1.send(&transport1, Payload::Div(4.0)); value /= 4.0; - seq1.send(&transport1, Packet::Mod(5.0)); + seq1.send(&transport1, Payload::Mod(5.0)); value %= 5.0; for _ in 0..16 { diff --git a/examples/file_download.rs b/examples/file_download.rs index 8d2be99..5356b69 100644 --- a/examples/file_download.rs +++ b/examples/file_download.rs @@ -11,14 +11,14 @@ use std::{ use rand_core::{OsRng, RngCore}; use seq_ex::{ - sync::{PacketType, RecvSuccess, ReplyGuard, SeqExSync}, - SeqNo, TransportLayer, + sync::{RecvOk, SeqExSync}, + Packet, TransportLayer, }; use serde::{Deserialize, Serialize}; const FILE_CHUNK_SIZE: usize = 1000; #[derive(Clone, Debug, Serialize, Deserialize)] -enum Packet { +enum Payload { RequestFile { filename: String }, ConfirmRequestFile { filesize: u64 }, FileDownload { filename: String, file_chunk: Vec }, @@ -32,25 +32,17 @@ struct Transport { struct Peer { filesystem: Arc>>>, transport: Transport, - seqex: Arc>, + seqex: Arc>, receiver: Receiver>, } -impl TransportLayer for &Transport { +impl TransportLayer for &Transport { fn time(&mut self) -> i64 { self.time.elapsed().as_millis() as i64 } - fn send(&mut self, seq_no: SeqNo, reply_no: Option, payload: &Packet) { - let p = PacketType::Payload(seq_no, reply_no, payload.clone()); - if let Ok(p) = serde_json::to_vec(&p) { - let _ = self.sender.send(p); - } - } - - fn send_ack(&mut self, reply_no: SeqNo) { - let p = PacketType::::Ack(reply_no); - if let Ok(p) = serde_json::to_vec(&p) { + fn send(&mut self, packet: Packet<&Payload>) { + if let Ok(p) = serde_json::to_vec(&packet) { let _ = self.sender.send(p); } } @@ -60,14 +52,15 @@ fn drop_packet() -> bool { OsRng.next_u32() >= (u32::MAX / 4 * 3) } -fn process(peer: &Peer, guard: ReplyGuard<'_, &Transport, Packet, Packet>, recv_packet: Packet, sent_packet: Option) { - match (recv_packet, sent_packet) { - (Packet::RequestFile { filename }, None) => { +fn process(peer: &Peer, recv_data: RecvOk<'_, &Transport, Payload, Payload, Payload>) { + use Payload::*; + match recv_data.consume() { + (Some((guard, RequestFile { filename })), None) => { let filesystem = peer.filesystem.clone(); let transport = peer.transport.clone(); let seqex = peer.seqex.clone(); if let Some(file) = filesystem.read().unwrap().get(&filename) { - guard.reply(Packet::ConfirmRequestFile { filesize: file.len() as u64 }); + guard.reply(ConfirmRequestFile { filesize: file.len() as u64 }); } thread::spawn(move || { let filesystem = filesystem.read().unwrap(); @@ -77,21 +70,18 @@ fn process(peer: &Peer, guard: ReplyGuard<'_, &Transport, Packet, Packet>, recv_ let mut i = 0; while i < file.len() { let j = file.len().min(i + FILE_CHUNK_SIZE); - seqex.send( - &transport, - Packet::FileDownload { filename: filename.clone(), file_chunk: file[i..j].to_vec() }, - ); + seqex.send_locked(&transport, FileDownload { filename: filename.clone(), file_chunk: file[i..j].to_vec() }); i = j; } } }); } - (Packet::ConfirmRequestFile { filesize }, Some(Packet::RequestFile { filename })) => { + (Some((_, ConfirmRequestFile { filesize })), Some(RequestFile { filename })) => { let mut filesystem = peer.filesystem.write().unwrap(); let file = Vec::with_capacity(filesize as usize); filesystem.insert(filename, file); } - (Packet::FileDownload { filename, file_chunk }, None) => { + (Some((_, FileDownload { filename, file_chunk })), None) => { let mut filesystem = peer.filesystem.write().unwrap(); if let Some(file) = filesystem.get_mut(&filename) { if file.len() + file_chunk.len() <= file.capacity() { @@ -99,9 +89,10 @@ fn process(peer: &Peer, guard: ReplyGuard<'_, &Transport, Packet, Packet>, recv_ } } } - _ => { - assert!(false); + (Some(a), b) => { + print!("Unsolicited packet received: {:?}", RecvOk::new(Some(a), b)); } + _ => {} } } @@ -110,17 +101,10 @@ fn receive(peer: &Peer) { if drop_packet() { continue; } - let parsed_packet = serde_json::from_slice::>(&packet); - match parsed_packet { - Ok(PacketType::Ack(reply_no)) => { - let _ = peer.seqex.receive_ack(reply_no); + if let Ok(parsed_packet) = serde_json::from_slice::>(&packet) { + for recv_data in peer.seqex.receive_all(&peer.transport, parsed_packet) { + process(peer, recv_data); } - Ok(PacketType::Payload(seq_no, reply_no, payload)) => { - for RecvSuccess { guard, packet, send_data } in peer.seqex.receive_all(&peer.transport, seq_no, reply_no, payload) { - process(peer, guard, packet, send_data); - } - } - _ => {} } } } @@ -153,9 +137,9 @@ fn main() { receiver: recv2, }; - peer1.seqex.send(&peer1.transport, Packet::RequestFile { filename: "File1".to_string() }); - peer1.seqex.send(&peer1.transport, Packet::RequestFile { filename: "File3".to_string() }); - peer1.seqex.send(&peer1.transport, Packet::RequestFile { filename: "File2".to_string() }); + peer1.seqex.send(&peer1.transport, Payload::RequestFile { filename: "File1".to_string() }); + peer1.seqex.send(&peer1.transport, Payload::RequestFile { filename: "File3".to_string() }); + peer1.seqex.send(&peer1.transport, Payload::RequestFile { filename: "File2".to_string() }); for _ in 0..300 { receive(&peer1); diff --git a/examples/hello_world.rs b/examples/hello_world.rs index 7202448..3e284ca 100644 --- a/examples/hello_world.rs +++ b/examples/hello_world.rs @@ -1,55 +1,46 @@ use std::sync::mpsc::Receiver; -use seq_ex::sync::{MpscGuard, MpscSeqEx, MpscTransport, PacketType, RecvSuccess}; +use seq_ex::{ + sync::{MpscSeqEx, MpscTransport, RecvOk}, + Packet, +}; #[derive(Clone, Debug)] -enum Packet { +enum Payload { Hello, Space, World, Exclamation, } -use Packet::*; +use Payload::*; -fn process(guard: MpscGuard<'_, Packet>, recv_packet: Packet, send_packet: Option) { - match (recv_packet, send_packet) { - (Hello, None) => { - print!("Hello"); - guard.reply(Space); - } - (Space, Some(Hello)) => { - print!(" "); - guard.reply(World); - } - (World, Some(Space)) => { - print!("World"); - guard.reply(Exclamation); - } - (Exclamation, Some(World)) => { - print!("!"); - } - (a, None) => { - print!("Unsolicited packet received: {:?}", a); - } - (a, Some(b)) => { - print!("Incorrect reply received: {:?}, was a reply to: {:?}", a, b); - } - } -} - -fn receive(recv: &Receiver>, seq: &MpscSeqEx, transport: &MpscTransport) { - match recv.recv().unwrap() { - PacketType::Ack(reply_no) => { - let result = seq.receive_ack(reply_no); - if let Ok(Exclamation) = result { +fn receive(recv: &Receiver>, seq: &MpscSeqEx, transport: &MpscTransport) { + let packet = recv.recv().unwrap(); + for recv_data in seq.receive_all(transport, packet) { + match recv_data.consume() { + (Some((guard, Hello)), None) => { + print!("Hello"); + guard.reply(Space); + } + (Some((guard, Space)), Some(Hello)) => { + print!(" "); + guard.reply(World); + } + (Some((guard, World)), Some(Space)) => { + print!("World"); + guard.reply(Exclamation); + } + (Some((_, Exclamation)), Some(World)) => { + print!("!"); + } + (None, Some(Exclamation)) => { // Our Hello World exchange ends right here. print!("\n"); } - } - PacketType::Payload(seq_no, reply_no, payload) => { - for RecvSuccess { guard, packet, send_data } in seq.receive_all(transport, seq_no, reply_no, payload) { - process(guard, packet, send_data) + (Some(a), b) => { + print!("Unsolicited packet received: {:?}", RecvOk::new(Some(a), b)); } + _ => {} } } } @@ -61,7 +52,7 @@ fn main() { let seq2 = MpscSeqEx::default(); // We begin a "Hello World" exchange right here. - seq1.send(&transport1, Packet::Hello); + seq1.send(&transport1, Payload::Hello); receive(&recv2, &seq2, &transport2); receive(&recv1, &seq1, &transport1); diff --git a/src/seq_queue.rs b/src/seq_queue.rs index 01ca110..9f4f081 100644 --- a/src/seq_queue.rs +++ b/src/seq_queue.rs @@ -34,8 +34,6 @@ //! ## Examples //! -use crate::TransportLayer; - /// A 32-bit sequence number. Packets transported with SEP are expected to contain at least one /// sequence number, and sometimes two. /// All packets will either have a seq_no, a reply_no, or both. @@ -67,11 +65,13 @@ pub struct SeqEx { /// when `reply_raw` or `ack_raw` are called with it, they are guaranteed not to fail. /// To accomplish this we must track all issued reply numbers. concurrent_replies_total: usize, + is_locked: bool, } struct RecvEntry { seq_no: SeqNo, reply_no: Option, + locked: bool, data: RecvData, } @@ -82,16 +82,123 @@ struct SendEntry { data: SendData, } -/// The error type for when a packet has been received, but for whatever reason could not be -/// immediately processed. -#[derive(Debug, Clone)] -pub enum Error { +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum DirectError { + /// The packet is out-of-sequence. It was either received too soon or too late and so it would be + /// invalid to process it right now. No action needs to be taken by the caller. + OutOfSequence, + /// The Send Window is currently full. The received packet cannot be processed right now because + /// it could cause the send window to overflow. + WindowIsFull(Packet), + WindowIsLocked(Packet), + ResendAck(SeqNo), +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum PumpError { /// The packet is out-of-sequence. It was either received too soon or too late and so it would be /// invalid to process it right now. No action needs to be taken by the caller. OutOfSequence, /// The Send Window is currently full. The received packet cannot be processed right now because /// it could cause the send window to overflow. No action needs to be taken by the caller. WindowIsFull, + WindowIsLocked, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] +pub enum Packet { + Payload(SeqNo, RecvData), + LockPayload(SeqNo, RecvData), + Reply(SeqNo, SeqNo, RecvData), + LockReply(SeqNo, SeqNo, RecvData), + Ack(SeqNo), +} +use Packet::*; +impl Packet { + pub fn new_with_data(seq_no: SeqNo, reply_no: Option, is_locking: bool, data: RecvData) -> Self { + Self::new(Some(seq_no), reply_no, is_locking, Some(data)).unwrap() + } + pub fn new(seq_no: Option, reply_no: Option, is_locking: bool, data: Option) -> Option { + match (seq_no, reply_no, is_locking, data) { + (Some(s), None, false, Some(d)) => Some(Payload(s, d)), + (Some(s), None, true, Some(d)) => Some(LockPayload(s, d)), + (Some(s), Some(r), false, Some(d)) => Some(Reply(s, r, d)), + (Some(s), Some(r), true, Some(d)) => Some(LockReply(s, r, d)), + (None, Some(r), false, None) => Some(Ack(r)), + _ => None, + } + } + pub fn as_ref(&self) -> Packet<&RecvData> { + match self { + Payload(seq_no, data) => Payload(*seq_no, data), + LockPayload(seq_no, data) => LockPayload(*seq_no, data), + Reply(seq_no, reply_no, data) => Reply(*seq_no, *reply_no, data), + LockReply(seq_no, reply_no, data) => LockReply(*seq_no, *reply_no, data), + Ack(reply_no) => Ack(*reply_no), + } + } + pub fn map(self, f: impl FnOnce(RecvData) -> SendData) -> Packet { + match self { + Payload(seq_no, data) => Payload(seq_no, f(data)), + LockPayload(seq_no, data) => LockPayload(seq_no, f(data)), + Reply(seq_no, reply_no, data) => Reply(seq_no, reply_no, f(data)), + LockReply(seq_no, reply_no, data) => LockReply(seq_no, reply_no, f(data)), + Ack(reply_no) => Ack(reply_no), + } + } + pub fn payload(self) -> Option { + match self { + Payload(_, data) | LockPayload(_, data) | Reply(_, _, data) | LockReply(_, _, data) => Some(data), + Ack(_) => None, + } + } + pub fn is_locking(&self) -> bool { + matches!(self, LockPayload(..) | LockReply(..)) + } + pub fn set_locking(&mut self, locking: bool) { + let mut tmp = Ack(0); + core::mem::swap(&mut tmp, self); + match tmp { + Payload(seq_no, data) | LockPayload(seq_no, data) => { + *self = if locking { + LockPayload(seq_no, data) + } else { + Payload(seq_no, data) + } + } + Reply(seq_no, reply_no, data) | LockReply(seq_no, reply_no, data) => { + *self = if locking { + LockReply(seq_no, reply_no, data) + } else { + Reply(seq_no, reply_no, data) + } + } + Ack(reply_no) => *self = Ack(reply_no), + } + } +} +impl Packet<&RecvData> { + pub fn cloned(&self) -> Packet { + self.map(|d| d.clone()) + } +} + +pub enum RecvOkRaw { + Payload { + reply_no: SeqNo, + locked: bool, + recv_data: RecvData, + }, + Reply { + reply_no: SeqNo, + locked: bool, + recv_data: RecvData, + send_data: SendData, + }, + Ack { + send_data: SendData, + }, } /// An iterator over all packets in the send window. It will iterate over all packets currently @@ -109,6 +216,12 @@ pub struct Iter<'a, SendData>(core::slice::Iter<'a, Option>> /// the packet. pub struct IterMut<'a, SendData>(core::slice::IterMut<'a, Option>>); +#[derive(Clone, Debug)] +pub struct ServiceIter { + idx: usize, + next_time: i64, +} + impl SeqEx { /// Creates a new instance of `SeqEx` for a new remote peer. /// An instance of `SeqEx` expects to communicate with only exactly one other remote instance @@ -129,6 +242,7 @@ impl SeqEx { send_window: core::array::from_fn(|_| None), concurrent_replies: core::array::from_fn(|_| 0), concurrent_replies_total: 0, + is_locked: false, } } fn send_window_slot_mut(&mut self, seq_no: SeqNo) -> &mut Option> { @@ -137,6 +251,38 @@ impl SeqEx { fn send_window_slot(&self, seq_no: SeqNo) -> &Option> { &self.send_window[seq_no as usize % self.send_window.len()] } + fn is_full_inner(&self, reserve_one: bool) -> bool { + if self.concurrent_replies_total >= self.concurrent_replies.len() { + return true; + } + for i in 0..self.concurrent_replies_total as u32 + 1 + reserve_one as u32 { + let slot = self.send_window_slot(self.next_send_seq_no.wrapping_add(i)); + if slot.is_some() { + return true; + } + } + false + } + fn take_send(&mut self, reply_no: SeqNo) -> Option { + let slot = self.send_window_slot_mut(reply_no); + if slot.as_ref().map_or(false, |e| e.seq_no == reply_no) { + slot.take().map(|e| e.data) + } else { + None + } + } + + fn remove_reservation(&mut self, reply_no: SeqNo) -> bool { + for i in 0..self.concurrent_replies_total { + if self.concurrent_replies[i] == reply_no { + // swap remove + self.concurrent_replies_total -= 1; + self.concurrent_replies[i] = self.concurrent_replies[self.concurrent_replies_total]; + return true; + } + } + false + } /// Returns whether or not the send window is full. /// If the send window is full calls to `SeqEx::send` will always fail. @@ -171,35 +317,56 @@ impl SeqEx { /// `current_time` should be a timestamp of the current time, using whatever units of time the /// user would like. However this choice of units must be consistent with the units of the /// `retry_interval`. `current_time` does not have to be monotonically increasing. - #[must_use = "The queue might be full causing the packet to not be sent"] - pub fn try_send(&mut self, mut app: impl TransportLayer, packet_data: SendData) -> Result<(), SendData> { + fn try_send_direct_inner(&mut self, current_time: i64) -> Option<(&mut Option>, SeqNo, i64)> { if self.is_full() { - return Err(packet_data); + return None; } let seq_no = self.next_send_seq_no; self.next_send_seq_no = self.next_send_seq_no.wrapping_add(1); - let current_time = app.time(); let next_resend_time = current_time + self.resend_interval; if self.next_service_timestamp > next_resend_time { self.next_service_timestamp = next_resend_time; } let slot = self.send_window_slot_mut(seq_no); debug_assert!(slot.is_none()); - let entry = slot.insert(SendEntry { seq_no, reply_no: None, next_resend_time, data: packet_data }); - - app.send(entry.seq_no, entry.reply_no, &entry.data); - Ok(()) + Some((slot, seq_no, next_resend_time)) + } + pub fn try_send_direct(&mut self, packet_data: SendData, current_time: i64) -> Result, SendData> { + if let Some((slot, seq_no, next_resend_time)) = self.try_send_direct_inner(current_time) { + let entry = slot.insert(SendEntry { seq_no, reply_no: None, next_resend_time, data: packet_data }); + Ok(Packet::Payload(entry.seq_no, &entry.data)) + } else { + Err(packet_data) + } + } + pub fn try_send_direct_with(&mut self, packet_data: impl FnOnce(SeqNo) -> SendData, current_time: i64) -> Result, ()> { + if let Some((slot, seq_no, next_resend_time)) = self.try_send_direct_inner(current_time) { + let entry = slot.insert(SendEntry { + seq_no, + reply_no: None, + next_resend_time, + data: packet_data(seq_no), + }); + Ok(Packet::Payload(entry.seq_no, &entry.data)) + } else { + Err(()) + } } /// If this returns `Ok` then `try_send` might succeed on next call. - pub fn receive_raw>( - &mut self, - mut app: impl TransportLayer, - seq_no: SeqNo, - reply_no: Option, - packet: P, - ) -> Result<(SeqNo, P, Option), Error> { + pub fn receive_raw_and_direct>(&mut self, packet: Packet

) -> Result, DirectError

> { + let locked = packet.is_locking(); + let (seq_no, reply_no, recv_data) = match packet { + Packet::Payload(seq_no, recv_data) | Packet::LockPayload(seq_no, recv_data) => (seq_no, None, recv_data), + Packet::Reply(seq_no, reply_no, recv_data) | Packet::LockReply(seq_no, reply_no, recv_data) => (seq_no, Some(reply_no), recv_data), + Packet::Ack(reply_no) => { + return self + .take_send(reply_no) + .map(|send_data| RecvOkRaw::Ack { send_data }) + .ok_or(DirectError::OutOfSequence) + } + }; // We only want to accept packets with sequence numbers in the range: // `self.pre_recv_seq_no < seq_no <= self.pre_recv_seq_no + self.recv_window.len()`. // To check that range we compute `seq_no - (self.pre_recv_seq_no + 1)` and check @@ -220,104 +387,105 @@ impl SeqEx { // resending the packet. for entry in self.send_window.iter().flatten() { if entry.reply_no == Some(seq_no) { - return Err(Error::OutOfSequence); + return Err(DirectError::OutOfSequence); } } for i in 0..self.concurrent_replies_total { if self.concurrent_replies[i] == seq_no { - return Err(Error::OutOfSequence); + return Err(DirectError::OutOfSequence); } } - app.send_ack(seq_no); - return Err(Error::OutOfSequence); + return Err(DirectError::ResendAck(seq_no)); } else if is_above_range { - return Err(Error::OutOfSequence); + return Err(DirectError::OutOfSequence); } // If the send window is full we cannot safely process received packets, // because there would be no way to reply. // We can only process this packet if processing it would make space in the send window. let is_full = self.is_full_inner(false); + // Check whether or not we've already received this packet let i = seq_no as usize % self.recv_window.len(); - if let Some(pre) = self.recv_window[i].as_mut() { + let is_in_window = if let Some(pre) = self.recv_window[i].as_mut() { if seq_no == pre.seq_no { - if is_next && !is_full { - self.recv_window[i] = None; - } else { - return if is_full { - Err(Error::WindowIsFull) - } else { - Err(Error::OutOfSequence) - }; - } + true } else { // This is currently unreachable due to the range check. // `self.pre_recv_seq_no < seq_no <= self.pre_recv_seq_no + self.recv_window.len()`. - return Err(Error::OutOfSequence); + return Err(DirectError::OutOfSequence); + } + } else { + false + }; + + if !is_next { + if !is_in_window { + self.recv_window[i] = Some(RecvEntry { seq_no, reply_no, locked, data: recv_data.into() }) + } + Err(DirectError::OutOfSequence) + } else if is_full { + Err(DirectError::WindowIsFull(Packet::new_with_data(seq_no, reply_no, locked, recv_data))) + } else if locked && self.is_locked { + Err(DirectError::WindowIsLocked(Packet::new_with_data(seq_no, reply_no, locked, recv_data))) + } else { + if is_in_window { + self.recv_window[i] = None; } - } - if is_next && !is_full { self.pre_recv_seq_no = seq_no; self.concurrent_replies[self.concurrent_replies_total] = seq_no; self.concurrent_replies_total += 1; - let data = reply_no.and_then(|r| self.take_send(r)); - Ok((seq_no, packet, data)) - } else { - self.recv_window[i] = Some(RecvEntry { seq_no, reply_no, data: packet.into() }); - if is_full { - Err(Error::WindowIsFull) + if locked { + debug_assert!(!self.is_locked); + self.is_locked = true; + } + Ok(if let Some(send_data) = reply_no.and_then(|r| self.take_send(r)) { + RecvOkRaw::Reply { reply_no: seq_no, recv_data, send_data, locked } } else { - Err(Error::OutOfSequence) - } + RecvOkRaw::Payload { reply_no: seq_no, recv_data, locked } + }) } } /// If this returns `Ok` then `try_send` might succeed on next call. - pub fn receive_ack(&mut self, reply_no: SeqNo) -> Result { - self.take_send(reply_no).ok_or(Error::OutOfSequence) - } - - fn is_full_inner(&self, reserve_one: bool) -> bool { - if self.concurrent_replies_total >= self.concurrent_replies.len() { - return true; - } - for i in 0..self.concurrent_replies_total as u32 + 1 + reserve_one as u32 { - let slot = self.send_window_slot(self.next_send_seq_no.wrapping_add(i)); - if slot.is_some() { - return true; - } - } - false - } - - fn take_send(&mut self, reply_no: SeqNo) -> Option { - let slot = self.send_window_slot_mut(reply_no); - if slot.as_ref().map_or(false, |e| e.seq_no == reply_no) { - slot.take().map(|e| e.data) - } else { - None - } - } - /// If this returns `Ok` then `try_send` might succeed on next call. - pub fn pump_raw(&mut self) -> Result<(SeqNo, RecvData, Option), Error> { + pub fn pump_raw(&mut self) -> Result, PumpError> { let next_seq_no = self.pre_recv_seq_no.wrapping_add(1); let i = next_seq_no as usize % self.recv_window.len(); - if self.recv_window[i].as_ref().map_or(false, |pre| pre.seq_no == next_seq_no) { - if self.is_full_inner(false) { - return Err(Error::WindowIsFull); + if let Some(entry) = &self.recv_window[i].as_ref() { + if entry.seq_no == next_seq_no { + if self.is_full_inner(false) { + return Err(PumpError::WindowIsFull); + } + + if !entry.locked || !self.is_locked { + let entry = self.recv_window[i].take().unwrap(); + self.pre_recv_seq_no = next_seq_no; + self.concurrent_replies[self.concurrent_replies_total] = entry.seq_no; + self.concurrent_replies_total += 1; + if entry.locked { + debug_assert!(!self.is_locked); + self.is_locked = true; + } + return Ok(if let Some(send_data) = entry.reply_no.and_then(|r| self.take_send(r)) { + RecvOkRaw::Reply { + reply_no: entry.seq_no, + locked: entry.locked, + recv_data: entry.data, + send_data, + } + } else { + RecvOkRaw::Payload { + reply_no: entry.seq_no, + locked: entry.locked, + recv_data: entry.data, + } + }); + } else { + return Err(PumpError::WindowIsLocked); + } } - - let entry = self.recv_window[i].take().unwrap(); - self.pre_recv_seq_no = next_seq_no; - self.concurrent_replies[self.concurrent_replies_total] = entry.seq_no; - self.concurrent_replies_total += 1; - let data = entry.reply_no.and_then(|r| self.take_send(r)); - Ok((entry.seq_no, entry.data, data)) - } else { - Err(Error::OutOfSequence) } + Err(PumpError::OutOfSequence) } - /// This function must be passed a reply number given by `receive_raw` or `pump_raw`, otherwise /// it will do nothing. This reply number can only be used to reply once. /// @@ -326,16 +494,20 @@ impl SeqEx { /// The identifier will tell the remote peer which packets contain fragments of the file, /// and since each fragment will be received in order it will be trivial for them to reconstruct /// the original file. - pub fn reply_raw(&mut self, mut app: impl TransportLayer, reply_no: SeqNo, packet_data: SendData) { + #[must_use] + pub fn reply_raw_and_direct(&mut self, reply_no: SeqNo, unlock: bool, packet_data: SendData, current_time: i64) -> Option> { if self.remove_reservation(reply_no) { let seq_no = self.next_send_seq_no; self.next_send_seq_no = self.next_send_seq_no.wrapping_add(1); - let current_time = app.time(); let next_resend_time = current_time + self.resend_interval; if self.next_service_timestamp > next_resend_time { self.next_service_timestamp = next_resend_time; } + if unlock { + debug_assert!(self.is_locked, "The window must be locked to attempt to unlock: double unlock detected."); + self.is_locked = false; + } let slot = self.send_window_slot_mut(seq_no); debug_assert!(slot.is_none()); let entry = slot.insert(SendEntry { @@ -345,46 +517,47 @@ impl SeqEx { data: packet_data, }); - app.send(entry.seq_no, entry.reply_no, &entry.data); + Some(Packet::Reply(entry.seq_no, reply_no, &entry.data)) + } else { + None } } - pub fn ack_raw(&mut self, mut app: impl TransportLayer, reply_no: SeqNo) { + #[must_use] + pub fn ack_raw_and_direct(&mut self, reply_no: SeqNo, unlock: bool) -> Option> { if self.remove_reservation(reply_no) { - // Acks are only sent once. There is code in `receive_raw` to handle resending - // an ack in the event that the first one here was dropped by the network. - app.send_ack(reply_no); - } - } - fn remove_reservation(&mut self, reply_no: SeqNo) -> bool { - for i in 0..self.concurrent_replies_total { - if self.concurrent_replies[i] == reply_no { - // swap remove - self.concurrent_replies_total -= 1; - self.concurrent_replies[i] = self.concurrent_replies[self.concurrent_replies_total]; - return true; + if unlock { + debug_assert!(self.is_locked, "The window must be locked to attempt to unlock: double unlock detected."); + self.is_locked = false; } + Some(Packet::Ack(reply_no)) + } else { + None } - false } - pub fn service(&mut self, mut app: impl TransportLayer) -> i64 { - let current_time = app.time(); - let real_interval = if self.next_service_timestamp <= current_time { - let next_resend_time = current_time + self.resend_interval; - let mut next_activity = i64::MAX; - for entry in self.send_window.iter_mut().flatten() { - if entry.next_resend_time <= current_time { - entry.next_resend_time = next_resend_time; - app.send(entry.seq_no, entry.reply_no, &entry.data); + pub fn service_direct<'a>(&'a mut self, current_time: i64, iter: &mut Option) -> Option> { + if self.next_service_timestamp <= current_time { + let iter = iter.get_or_insert(ServiceIter { idx: 0, next_time: i64::MAX }); + while let Some(entry) = self.send_window.get(iter.idx) { + iter.idx += 1; + if let Some(entry) = entry { + if entry.next_resend_time <= current_time { + let entry = self.send_window[iter.idx - 1].as_mut().unwrap(); + entry.next_resend_time = current_time + self.resend_interval; + iter.next_time = iter.next_time.min(entry.next_resend_time); + return Some(if let Some(reply_no) = entry.reply_no { + Packet::Reply(entry.seq_no, reply_no, &entry.data) + } else { + Packet::Payload(entry.seq_no, &entry.data) + }); + } else { + iter.next_time = iter.next_time.min(entry.next_resend_time); + } } - next_activity = next_activity.min(entry.next_resend_time); } - self.next_service_timestamp = next_activity; - next_activity - current_time - } else { - self.next_service_timestamp - current_time - }; - self.resend_interval.min(real_interval) + self.next_service_timestamp = iter.next_time; + } + None } pub fn iter(&self) -> Iter<'_, SendData> { @@ -428,10 +601,6 @@ macro_rules! iterator { } None } - - fn size_hint(&self) -> (usize, Option) { - (0, Some(self.0.len())) - } } impl<'a, SendData> DoubleEndedIterator for $iter<'a, SendData> { fn next_back(&mut self) -> Option { diff --git a/src/single_thread.rs b/src/single_thread.rs index 0faae3c..813754a 100644 --- a/src/single_thread.rs +++ b/src/single_thread.rs @@ -1,10 +1,11 @@ -use crate::{Error, SeqEx, SeqNo, TransportLayer, DEFAULT_WINDOW_CAP}; +use crate::{DirectError, Packet, PumpError, RecvOkRaw, SeqEx, SeqNo, TransportLayer, DEFAULT_WINDOW_CAP}; -pub struct ReplyGuard<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP>( - &'a mut SeqEx, - TL, - SeqNo, -); +pub struct ReplyGuard<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { + seq: &'a mut SeqEx, + app: Option, + reply_no: SeqNo, + locked: bool, +} impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> ReplyGuard<'a, TL, SendData, RecvData, CAP> { /// If you need to reply more than once, say to fragment a large file, then include in your /// first reply some identifier, and then `send` all fragments with the same included identifier. @@ -12,35 +13,204 @@ impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> Rep /// and since each fragment will be received in order it will be trivial for them to reconstruct /// the original file. pub fn reply(self, packet_data: SendData) { - self.0.reply_raw(self.1.clone(), self.2, packet_data); + self.reply_inner(false, |_, _| packet_data) + } + pub fn reply_locked(self, packet_data: SendData) { + self.reply_inner(true, |_, _| packet_data) + } + pub fn reply_with(self, packet_data: impl FnOnce(SeqNo, SeqNo) -> SendData) { + self.reply_inner(false, packet_data) + } + pub fn reply_locked_with(self, packet_data: impl FnOnce(SeqNo, SeqNo) -> SendData) { + self.reply_inner(true, packet_data) + } + fn reply_inner(mut self, locked: bool, packet_data: impl FnOnce(SeqNo, SeqNo) -> SendData) { + let mut app = None; + core::mem::swap(&mut app, &mut self.app); + let seq_no = self.seq.seq_no(); + self.seq + .reply_raw(app.unwrap(), self.reply_no, self.locked, locked, packet_data(seq_no, self.reply_no)); core::mem::forget(self); } } impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> Drop for ReplyGuard<'a, TL, SendData, RecvData, CAP> { fn drop(&mut self) { - self.0.ack_raw(self.1.clone(), self.2) + if let Some(app) = &mut self.app { + if let Some(p) = self.seq.ack_raw_and_direct(self.reply_no, self.locked) { + app.send(p) + } + } + } +} +impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> std::fmt::Debug for ReplyGuard<'a, TL, SendData, RecvData, CAP> { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ReplyGuard").field("reply_no", &self.reply_no).finish() } } -pub struct RecvSuccess<'a, TL: TransportLayer, P: Into, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { - pub guard: ReplyGuard<'a, TL, SendData, RecvData, CAP>, - pub packet: P, - pub send_data: Option, +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum Error { + /// The packet is out-of-sequence. It was either received too soon or too late and so it would be + /// invalid to process it right now. No action needs to be taken by the caller. + OutOfSequence, + /// The Send Window is currently full. The received packet cannot be processed right now because + /// it could cause the send window to overflow. + WindowIsFull(Packet), + WindowIsLocked(Packet), } +pub enum RecvOk<'a, TL: TransportLayer, P, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { + Payload { + reply_guard: ReplyGuard<'a, TL, SendData, RecvData, CAP>, + recv_data: P, + }, + Reply { + reply_guard: ReplyGuard<'a, TL, SendData, RecvData, CAP>, + recv_data: P, + send_data: SendData, + }, + Ack { + send_data: SendData, + }, +} +macro_rules! impl_recvok { + ($recv:tt, $seq_ex:ty) => { + #[cfg(feature = "std")] + impl<'a, TL: TransportLayer, P: std::fmt::Debug, SendData: std::fmt::Debug, RecvData, const CAP: usize> std::fmt::Debug + for $recv<'a, TL, P, SendData, RecvData, CAP> + { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Payload { reply_guard, recv_data } => f + .debug_struct("Payload") + .field("reply_guard", reply_guard) + .field("recv_data", recv_data) + .finish(), + Self::Reply { reply_guard, recv_data, send_data } => f + .debug_struct("Reply") + .field("reply_guard", reply_guard) + .field("recv_data", recv_data) + .field("send_data", send_data) + .finish(), + Self::Ack { send_data } => f.debug_struct("Ack").field("send_data", send_data).finish(), + } + } + } + impl<'a, TL: TransportLayer, P, SendData, RecvData, const CAP: usize> $recv<'a, TL, P, SendData, RecvData, CAP> { + pub fn from_raw(seq: $seq_ex, app: TL, value: RecvOkRaw) -> Self { + match value { + RecvOkRaw::Payload { reply_no, locked, recv_data } => Self::Payload { + reply_guard: ReplyGuard { seq, app: Some(app), reply_no, locked }, + recv_data, + }, + RecvOkRaw::Reply { reply_no, locked, recv_data, send_data } => Self::Reply { + reply_guard: ReplyGuard { seq, app: Some(app), reply_no, locked }, + recv_data, + send_data, + }, + RecvOkRaw::Ack { send_data } => Self::Ack { send_data }, + } + } + pub fn consume(self) -> (Option<(ReplyGuard<'a, TL, SendData, RecvData, CAP>, P)>, Option) { + match self { + Self::Payload { reply_guard, recv_data } => (Some((reply_guard, recv_data)), None), + Self::Reply { reply_guard, recv_data, send_data } => (Some((reply_guard, recv_data)), Some(send_data)), + Self::Ack { send_data } => (None, Some(send_data)), + } + } + pub fn new(recv_data: Option<(ReplyGuard<'a, TL, SendData, RecvData, CAP>, P)>, send_data: Option) -> Option { + match (recv_data, send_data) { + (Some((reply_guard, recv_data)), None) => Some(Self::Payload { reply_guard, recv_data }), + (Some((reply_guard, recv_data)), Some(send_data)) => Some(Self::Reply { reply_guard, recv_data, send_data }), + (None, Some(send_data)) => Some(Self::Ack { send_data }), + (None, None) => None, + } + } + } + impl<'a, TL: TransportLayer, P: Into, SendData, RecvData, const CAP: usize> $recv<'a, TL, P, SendData, RecvData, CAP> { + pub fn into(self) -> $recv<'a, TL, RecvData, SendData, RecvData, CAP> { + match self { + Self::Payload { reply_guard, recv_data } => $recv::Payload { reply_guard, recv_data: recv_data.into() }, + Self::Reply { reply_guard, recv_data, send_data } => $recv::Reply { reply_guard, recv_data: recv_data.into(), send_data }, + Self::Ack { send_data } => $recv::Ack { send_data }, + } + } + } + }; +} +impl_recvok!(RecvOk, &'a mut SeqEx); +pub(crate) use impl_recvok; + impl SeqEx { + pub fn try_send(&mut self, mut app: impl TransportLayer, locked: bool, packet_data: SendData) -> Result<(), SendData> { + match self.try_send_direct(packet_data, app.time()) { + Ok(mut p) => { + p.set_locking(locked); + app.send(p); + Ok(()) + } + Err(e) => Err(e), + } + } + pub fn try_send_with( + &mut self, + mut app: impl TransportLayer, + locked: bool, + packet_data: impl FnOnce(SeqNo) -> SendData, + ) -> Result<(), ()> { + match self.try_send_direct_with(packet_data, app.time()) { + Ok(mut p) => { + p.set_locking(locked); + app.send(p); + Ok(()) + } + Err(e) => Err(e), + } + } + /// If this returns `Ok` then `try_send` might succeed on next call. + pub fn receive_raw>( + &mut self, + mut app: impl TransportLayer, + packet: Packet

, + ) -> Result, Error

> { + match self.receive_raw_and_direct(packet) { + Ok(a) => Ok(a), + Err(DirectError::ResendAck(reply_no)) => { + app.send(Packet::Ack(reply_no)); + Err(Error::OutOfSequence) + } + Err(DirectError::OutOfSequence) => Err(Error::OutOfSequence), + Err(DirectError::WindowIsFull(p)) => Err(Error::WindowIsFull(p)), + Err(DirectError::WindowIsLocked(p)) => Err(Error::WindowIsLocked(p)), + } + } + pub fn reply_raw(&mut self, mut app: impl TransportLayer, reply_no: SeqNo, unlock: bool, locked_packet: bool, packet_data: SendData) { + if let Some(mut p) = self.reply_raw_and_direct(reply_no, unlock, packet_data, app.time()) { + p.set_locking(locked_packet); + app.send(p) + } + } + pub fn ack_raw(&mut self, mut app: impl TransportLayer, reply_no: SeqNo, unlock: bool) { + if let Some(p) = self.ack_raw_and_direct(reply_no, unlock) { + app.send(p) + } + } + pub fn service(&mut self, mut app: impl TransportLayer) -> i64 { + let current_time = app.time(); + let mut iter = None; + while let Some(p) = self.service_direct(current_time, &mut iter) { + app.send(p) + } + self.resend_interval.min(self.next_service_timestamp - current_time) + } pub fn receive, P: Into>( &mut self, app: TL, - seq_no: SeqNo, - reply_no: Option, - packet: P, - ) -> Result, Error> { - self.receive_raw(app.clone(), seq_no, reply_no, packet) - .map(|(reply_no, packet, send_data)| RecvSuccess { guard: ReplyGuard(self, app, reply_no), packet, send_data }) + packet: Packet

, + ) -> Result, Error

> { + self.receive_raw(app.clone(), packet).map(|r| RecvOk::from_raw(self, app, r)) } - pub fn pump>(&mut self, app: TL) -> Result, Error> { - self.pump_raw() - .map(|(reply_no, packet, send_data)| RecvSuccess { guard: ReplyGuard(self, app, reply_no), packet, send_data }) + pub fn pump>(&mut self, app: TL) -> Result, PumpError> { + self.pump_raw().map(|r| RecvOk::from_raw(self, app, r)) } } diff --git a/src/sync.rs b/src/sync.rs index 33c23ef..ecdf918 100644 --- a/src/sync.rs +++ b/src/sync.rs @@ -7,17 +7,19 @@ use std::{ time::Instant, }; -use crate::{Error, SeqEx, SeqNo, TransportLayer, DEFAULT_INITIAL_SEQ_NO, DEFAULT_RESEND_INTERVAL_MS, DEFAULT_WINDOW_CAP}; +use crate::{Packet, PumpError, RecvOkRaw, SeqEx, SeqNo, TransportLayer, DEFAULT_INITIAL_SEQ_NO, DEFAULT_RESEND_INTERVAL_MS, DEFAULT_WINDOW_CAP}; pub struct SeqExSync { seq_ex: Mutex<(SeqEx, usize)>, send_block: Condvar, + recv_lock: Condvar, } pub struct ReplyGuard<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { - origin: &'a SeqExSync, - app: TL, + seq: &'a SeqExSync, + app: Option, reply_no: SeqNo, + locked: bool, } impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> ReplyGuard<'a, TL, SendData, RecvData, CAP> { /// If you need to reply more than once, say to fragment a large file, then include in your @@ -26,34 +28,76 @@ impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> Rep /// and since each fragment will be received in order it will be trivial for them to reconstruct /// the original file. pub fn reply(self, packet_data: SendData) { - let mut seq = self.origin.lock(); - seq.reply_raw(self.app.clone(), self.reply_no, packet_data); - core::mem::forget(self); + self.reply_inner(false, |_, _| packet_data) + } + pub fn reply_locked(self, packet_data: SendData) { + self.reply_inner(true, |_, _| packet_data) } pub fn reply_with(self, packet_data: impl FnOnce(SeqNo, SeqNo) -> SendData) { - let mut seq = self.origin.lock(); + self.reply_inner(false, packet_data) + } + pub fn reply_locked_with(self, packet_data: impl FnOnce(SeqNo, SeqNo) -> SendData) { + self.reply_inner(true, packet_data) + } + fn reply_inner(mut self, locked: bool, packet_data: impl FnOnce(SeqNo, SeqNo) -> SendData) { + let mut app = None; + core::mem::swap(&mut app, &mut self.app); + let mut seq = self.seq.lock(); let seq_no = seq.seq_no(); - seq.reply_raw(self.app.clone(), self.reply_no, packet_data(seq_no, self.reply_no)); + seq.reply_raw(app.unwrap(), self.reply_no, self.locked, locked, packet_data(seq_no, self.reply_no)); + drop(seq); + if self.locked { + self.seq.recv_lock.notify_all(); + } core::mem::forget(self); } } impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> Drop for ReplyGuard<'a, TL, SendData, RecvData, CAP> { fn drop(&mut self) { - let mut seq = self.origin.lock(); - seq.ack_raw(self.app.clone(), self.reply_no); + if let Some(app) = self.app.as_mut() { + let mut seq = self.seq.lock(); + if let Some(p) = seq.ack_raw_and_direct(self.reply_no, self.locked) { + app.send(p) + } + } + } +} +impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> std::fmt::Debug for ReplyGuard<'a, TL, SendData, RecvData, CAP> { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ReplyGuard").field("reply_no", &self.reply_no).finish() } } -pub struct RecvSuccess<'a, TL: TransportLayer, P: Into, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { - pub guard: ReplyGuard<'a, TL, SendData, RecvData, CAP>, - pub packet: P, - pub send_data: Option, +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum Error { + /// The packet is out-of-sequence. It was either received too soon or too late and so it would be + /// invalid to process it right now. No action needs to be taken by the caller. + OutOfSequence, + /// The Send Window is currently full. The received packet cannot be processed right now because + /// it could cause the send window to overflow. No action needs to be taken by the caller. + WindowIsFull, } -pub struct ReplyIter<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { - origin: Option<&'a SeqExSync>, +pub enum RecvOk<'a, TL: TransportLayer, P, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { + Payload { + reply_guard: ReplyGuard<'a, TL, SendData, RecvData, CAP>, + recv_data: P, + }, + Reply { + reply_guard: ReplyGuard<'a, TL, SendData, RecvData, CAP>, + recv_data: P, + send_data: SendData, + }, + Ack { + send_data: SendData, + }, +} +crate::impl_recvok!(RecvOk, &'a SeqExSync); + +pub struct ReplyIter<'a, TL: TransportLayer, P: Into, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { + seq: Option<&'a SeqExSync>, app: TL, - first: Option>, + first: Option>, } pub struct SeqExGuard<'a, SendData, RecvData, const CAP: usize>(MutexGuard<'a, (SeqEx, usize)>); @@ -75,90 +119,103 @@ impl SeqExSync { Self { seq_ex: Mutex::new((SeqEx::new(retry_interval, initial_seq_no), 0)), send_block: Condvar::default(), + recv_lock: Condvar::default(), } } pub fn receive, P: Into>( &self, app: TL, - seq_no: SeqNo, - reply_no: Option, - packet: P, - ) -> Result, Error> { + mut packet: Packet

, + ) -> Result, Error> { let mut seq = self.seq_ex.lock().unwrap(); - let ret = seq.0.receive_raw(app.clone(), seq_no, reply_no, packet); - if seq.1 > 0 && ret.is_ok() { - self.send_block.notify_one(); + loop { + match seq.0.receive_raw(app.clone(), packet) { + Ok(r) => { + if seq.1 > 0 { + self.send_block.notify_one(); + } + return Ok(RecvOk::from_raw(self, app, r)); + } + Err(crate::Error::OutOfSequence) => return Err(Error::OutOfSequence), + Err(crate::Error::WindowIsFull(_)) => return Err(Error::WindowIsFull), + Err(crate::Error::WindowIsLocked(p)) => { + seq = self.recv_lock.wait(seq).unwrap(); + packet = p; + } + } } - ret.map(|(reply_no, packet, send_data)| RecvSuccess { - guard: ReplyGuard { origin: self, app, reply_no }, - packet, - send_data, - }) } - pub fn pump>(&self, app: TL) -> Result, Error> { + pub fn pump>(&self, app: TL) -> Result, Error> { let mut seq = self.seq_ex.lock().unwrap(); - let ret = seq.0.pump_raw(); - if seq.1 > 0 && ret.is_ok() { - self.send_block.notify_one(); + loop { + match seq.0.pump_raw() { + Ok(r) => { + if seq.1 > 0 { + self.send_block.notify_one(); + } + return Ok(RecvOk::from_raw(self, app, r)); + } + Err(PumpError::OutOfSequence) => return Err(Error::OutOfSequence), + Err(PumpError::WindowIsFull) => return Err(Error::WindowIsFull), + Err(PumpError::WindowIsLocked) => { + seq = self.recv_lock.wait(seq).unwrap(); + } + } } - ret.map(|(reply_no, packet, send_data)| RecvSuccess { - guard: ReplyGuard { origin: self, app, reply_no }, - packet, - send_data, - }) } - pub fn receive_all>( + pub fn receive_all, P: Into>( &self, app: TL, - seq_no: SeqNo, - reply_no: Option, - packet: RecvData, - ) -> ReplyIter<'_, TL, SendData, RecvData, CAP> { - if let Ok(g) = self.receive(app.clone(), seq_no, reply_no, packet) { - ReplyIter { origin: Some(self), app, first: Some(g) } + packet: Packet

, + ) -> ReplyIter<'_, TL, P, SendData, RecvData, CAP> { + if let Ok(r) = self.receive(app.clone(), packet) { + ReplyIter { seq: Some(self), app, first: Some(r) } } else { - ReplyIter { origin: None, app, first: None } + ReplyIter { seq: None, app, first: None } } } - pub fn try_send>(&self, app: TL, packet_data: SendData) -> Result<(), SendData> { - let mut seq = self.lock(); - seq.try_send(app, packet_data) + pub fn try_send>(&self, app: TL, locked: bool, packet_data: SendData) -> Result<(), SendData> { + self.try_send_with(app, locked, |_| packet_data) } - fn send_inner>( - &self, - mut seq: MutexGuard<'_, (SeqEx, usize)>, - app: TL, - mut packet_data: SendData, - ) { - while let Err(p) = seq.0.try_send(app.clone(), packet_data) { + fn send_with_inner>(&self, app: TL, locked: bool, mut packet_data: impl FnMut(SeqNo) -> SendData) { + let mut seq = self.seq_ex.lock().unwrap(); + while let Err(()) = seq.0.try_send_with(app.clone(), locked, &mut packet_data) { + seq.1 += 1; + seq = self.send_block.wait(seq).unwrap(); + seq.1 -= 1; + } + } + fn send_inner>(&self, app: TL, locked: bool, mut packet_data: SendData) { + let mut seq = self.seq_ex.lock().unwrap(); + while let Err(p) = seq.0.try_send(app.clone(), locked, packet_data) { packet_data = p; seq.1 += 1; seq = self.send_block.wait(seq).unwrap(); seq.1 -= 1; } } - pub fn send>(&self, app: TL, packet_data: SendData) { - self.send_inner(self.seq_ex.lock().unwrap(), app, packet_data) + pub fn send(&self, app: impl TransportLayer, packet_data: SendData) { + self.send_inner(app, false, packet_data) } - pub fn try_send_with>(&self, app: TL, packet_data: impl FnOnce(SeqNo) -> SendData) -> Result<(), SendData> { + pub fn send_locked(&self, app: impl TransportLayer, packet_data: SendData) { + self.send_inner(app, true, packet_data) + } + pub fn send_with(&self, app: impl TransportLayer, packet_data: impl FnMut(SeqNo) -> SendData) { + self.send_with_inner(app, false, packet_data) + } + pub fn send_locked_with(&self, app: impl TransportLayer, packet_data: impl FnMut(SeqNo) -> SendData) { + self.send_with_inner(app, true, packet_data) + } + pub fn try_send_with>( + &self, + app: TL, + locked: bool, + packet_data: impl FnOnce(SeqNo) -> SendData, + ) -> Result<(), SendData> { let mut seq = self.lock(); let seq_no = seq.seq_no(); - seq.try_send(app, packet_data(seq_no)) - } - pub fn send_with>(&self, app: TL, packet_data: impl FnOnce(SeqNo) -> SendData) { - let seq = self.seq_ex.lock().unwrap(); - let seq_no = seq.0.seq_no(); - self.send_inner(seq, app, packet_data(seq_no)) - } - - pub fn receive_ack(&self, reply_no: SeqNo) -> Result { - let mut seq = self.seq_ex.lock().unwrap(); - let ret = seq.0.receive_ack(reply_no); - if seq.1 > 0 && ret.is_ok() { - self.send_block.notify_one(); - } - ret + seq.try_send(app, locked, packet_data(seq_no)) } pub fn service>(&self, app: TL) -> i64 { self.lock().service(app) @@ -174,12 +231,19 @@ impl Default for SeqExSync, SendData, RecvData, const CAP: usize> Iterator for ReplyIter<'a, TL, SendData, RecvData, CAP> { - type Item = RecvSuccess<'a, TL, RecvData, SendData, RecvData, CAP>; +impl<'a, TL: TransportLayer, P: Into, SendData, RecvData, const CAP: usize> ReplyIter<'a, TL, P, SendData, RecvData, CAP> { + pub fn take_first(&mut self) -> Option> { + self.first.take() + } +} +impl<'a, TL: TransportLayer, P: Into, SendData, RecvData, const CAP: usize> Iterator + for ReplyIter<'a, TL, P, SendData, RecvData, CAP> +{ + type Item = RecvOk<'a, TL, RecvData, SendData, RecvData, CAP>; fn next(&mut self) -> Option { if let Some(g) = self.first.take() { - Some(g) - } else if let Some(origin) = self.origin { + Some(g.into()) + } else if let Some(origin) = self.seq { origin.pump(self.app.clone()).ok() } else { None @@ -187,27 +251,20 @@ impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> Ite } } -#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] -#[derive(Clone)] -pub enum PacketType { - Payload(SeqNo, Option, Payload), - Ack(SeqNo), -} - -#[derive(Clone)] +#[derive(Clone, Debug)] pub struct MpscTransport { - pub channel: Sender>, + pub channel: Sender>, pub time: Instant, } pub type MpscGuard<'a, Packet> = ReplyGuard<'a, &'a MpscTransport, Packet, Packet>; pub type MpscSeqEx = SeqExSync; impl MpscTransport { - pub fn new() -> (Self, Receiver>) { + pub fn new() -> (Self, Receiver>) { let (send, recv) = channel(); (Self { channel: send, time: std::time::Instant::now() }, recv) } - pub fn from_sender(send: Sender>) -> Self { + pub fn from_sender(send: Sender>) -> Self { Self { channel: send, time: std::time::Instant::now() } } } @@ -216,10 +273,7 @@ impl TransportLayer for &MpscTransport { self.time.elapsed().as_millis() as i64 } - fn send(&mut self, seq_no: SeqNo, reply_no: Option, payload: &Payload) { - let _ = self.channel.send(PacketType::Payload(seq_no, reply_no, payload.clone())); - } - fn send_ack(&mut self, reply_no: SeqNo) { - let _ = self.channel.send(PacketType::Ack(reply_no)); + fn send(&mut self, packet: Packet<&Payload>) { + let _ = self.channel.send(packet.cloned()); } } diff --git a/src/tokio.rs b/src/tokio.rs index 814eac7..c2ec721 100644 --- a/src/tokio.rs +++ b/src/tokio.rs @@ -1,9 +1,8 @@ -use std:: - sync::{ - Mutex, MutexGuard, - } -; -use tokio::{task, sync::{Notify, oneshot, mpsc}, time}; +use std::sync::{Mutex, MutexGuard}; +use tokio::{ + sync::{mpsc, oneshot, Notify}, + task, time, +}; use crate::{Error, SeqEx, SeqNo, TransportLayer, DEFAULT_INITIAL_SEQ_NO, DEFAULT_RESEND_INTERVAL_MS, DEFAULT_WINDOW_CAP}; @@ -82,11 +81,7 @@ pub struct TokioTransport { impl TokioTransport { pub fn new> + Send + 'static>(app: TL, seq: S) -> Self { let (update_queue, mut recv) = mpsc::channel(4); - let ret = TokioTransport { - time: time::Instant::now(), - update_queue, - app, - }; + let ret = TokioTransport { time: time::Instant::now(), update_queue, app }; let task_tl = ret.clone(); task::spawn(async move { let mut update_ts = i64::MAX; @@ -96,7 +91,7 @@ impl TokioTransport { let mut do_update = diff <= 0; if diff > 0 { let sleep = time::sleep(time::Duration::from_millis(diff as u64)); - tokio::select!{ + tokio::select! { Some(up) = recv.recv() => { update_ts = up; } @@ -131,7 +126,7 @@ impl SeqExTokio { &'a self, app: &'a TokioTransport, mut seq: MutexGuard<'_, (SeqEx, Packet, CAP>, usize)>, - result: Result<(SeqNo, Packet, Option>), Error> + result: Result<(SeqNo, Packet, Option>), Error>, ) -> Option<(Packet, ReplyGuard<'_, TL, Packet, CAP>)> { if let Ok((reply_no, packet, send_data)) = result { if seq.1 > 0 { @@ -161,7 +156,10 @@ impl SeqExTokio { let result = seq.0.receive_raw(app, seq_no, reply_no, packet); self.process(app, seq, result) } - pub fn pump<'a, TL: TokioTransportLayer>(&'a self, app: &'a TokioTransport) -> Option<(Packet, ReplyGuard<'_, TL, Packet, CAP>)> { + pub fn pump<'a, TL: TokioTransportLayer>( + &'a self, + app: &'a TokioTransport, + ) -> Option<(Packet, ReplyGuard<'_, TL, Packet, CAP>)> { let mut seq = self.seq_ex.lock().unwrap(); let result = seq.0.pump_raw(); self.process(app, seq, result) @@ -197,7 +195,7 @@ impl SeqExTokio { mut seq: MutexGuard<'_, (SeqEx, Packet, CAP>, usize)>, app: &TokioTransport, mut tx: oneshot::Sender<(Packet, SeqNo)>, - mut packet: Packet + mut packet: Packet, ) { let mut pre_ts = seq.0.next_service_timestamp; while let Err(e) = seq.0.try_send(app, (tx, packet)) { @@ -214,7 +212,11 @@ impl SeqExTokio { } } /// If this future is dropped then the remote peer's reply to this packet will also be dropped. - pub async fn send<'a, TL: TokioTransportLayer>(&'a self, app: &'a TokioTransport, packet: Packet) -> Option<(Packet, ReplyGuard<'_, TL, Packet, CAP>)> { + pub async fn send<'a, TL: TokioTransportLayer>( + &'a self, + app: &'a TokioTransport, + packet: Packet, + ) -> Option<(Packet, ReplyGuard<'_, TL, Packet, CAP>)> { self.send_with(app, |_| packet).await } //pub fn try_send_with>(&self, app: TL, packet_data: impl FnOnce(SeqNo) -> SendData) -> Result<(), SendData> { @@ -222,7 +224,11 @@ impl SeqExTokio { // let seq_no = seq.seq_no(); // seq.try_send(app, packet_data(seq_no)) //} - pub async fn send_with<'a, TL: TokioTransportLayer>(&'a self, app: &'a TokioTransport, packet: impl FnOnce(SeqNo) -> Packet) -> Option<(Packet, ReplyGuard<'_, TL, Packet, CAP>)> { + pub async fn send_with<'a, TL: TokioTransportLayer>( + &'a self, + app: &'a TokioTransport, + packet: impl FnOnce(SeqNo) -> Packet, + ) -> Option<(Packet, ReplyGuard<'_, TL, Packet, CAP>)> { let (tx, rx) = oneshot::channel(); let seq = self.seq_ex.lock().unwrap(); let seq_no = seq.0.seq_no(); diff --git a/src/transport_layer.rs b/src/transport_layer.rs index 53bc10d..ef67169 100644 --- a/src/transport_layer.rs +++ b/src/transport_layer.rs @@ -1,4 +1,4 @@ -use crate::SeqNo; +use crate::Packet; /// A trait for giving an instance of SeqEx access to the transport layer. /// @@ -9,6 +9,5 @@ use crate::SeqNo; pub trait TransportLayer: Clone { fn time(&mut self) -> i64; - fn send(&mut self, seq_no: SeqNo, reply_no: Option, payload: &SendData); - fn send_ack(&mut self, reply_no: SeqNo); + fn send(&mut self, packet: Packet<&SendData>); }