diff --git a/examples/file_download.rs b/examples/file_download.rs index 7a31388..034352b 100644 --- a/examples/file_download.rs +++ b/examples/file_download.rs @@ -52,7 +52,7 @@ fn drop_packet() -> bool { OsRng.next_u32() >= (u32::MAX / 4 * 3) } -fn process(peer: &Peer, recv_data: RecvOk<'_, &Transport, Payload, Payload, Payload>) { +fn process(peer: &Peer, recv_data: RecvOk<'_, &Transport, Payload, Payload>) { use Payload::*; match recv_data.consume() { (Some((guard, RequestFile { filename })), None) => { diff --git a/src/lib.rs b/src/lib.rs index 6a331a3..ee1159c 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -13,5 +13,5 @@ pub use single_thread::*; #[cfg(feature = "std")] pub mod sync; -#[cfg(feature = "tokio")] -pub mod tokio; +//#[cfg(feature = "tokio")] +//pub mod tokio; diff --git a/src/seq_queue.rs b/src/seq_queue.rs index efbb6b5..03e5ae0 100644 --- a/src/seq_queue.rs +++ b/src/seq_queue.rs @@ -93,7 +93,7 @@ struct SendEntry { } #[derive(Debug, Clone, PartialEq, Eq)] -pub enum DirectRecvError { +pub enum TryRecvError { DroppedTooEarly, DroppedDuplicate, DroppedDuplicateResendAck(SeqNo), @@ -101,36 +101,36 @@ pub enum DirectRecvError { WaitingForReply, } #[cfg(feature = "std")] -impl std::fmt::Display for DirectRecvError { +impl std::fmt::Display for TryRecvError { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { - DirectRecvError::DroppedTooEarly => write!(f, "packet arrived too early"), - DirectRecvError::DroppedDuplicate => write!(f, "packet was a duplicate"), - DirectRecvError::DroppedDuplicateResendAck(_) => write!(f, "packet was a duplicate, resending ack"), - DirectRecvError::WaitingForRecv => write!(f, "can't process until another packet is received"), - DirectRecvError::WaitingForReply => write!(f, "can't process until a reply is finished"), + TryRecvError::DroppedTooEarly => write!(f, "packet arrived too early"), + TryRecvError::DroppedDuplicate => write!(f, "packet was a duplicate"), + TryRecvError::DroppedDuplicateResendAck(_) => write!(f, "packet was a duplicate, resending ack"), + TryRecvError::WaitingForRecv => write!(f, "can't process until another packet is received"), + TryRecvError::WaitingForReply => write!(f, "can't process until a reply is finished"), } } } #[cfg(feature = "std")] -impl std::error::Error for DirectRecvError {} +impl std::error::Error for TryRecvError {} #[derive(Debug, Clone, PartialEq, Eq)] -pub enum PumpError { +pub enum TryError { WaitingForRecv, WaitingForReply, } #[cfg(feature = "std")] -impl std::fmt::Display for PumpError { +impl std::fmt::Display for TryError { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { - PumpError::WaitingForRecv => write!(f, "can't process until another packet is received"), - PumpError::WaitingForReply => write!(f, "can't process until a reply is finished"), + TryError::WaitingForRecv => write!(f, "can't process until another packet is received"), + TryError::WaitingForReply => write!(f, "can't process until a reply is finished"), } } } #[cfg(feature = "std")] -impl std::error::Error for PumpError {} +impl std::error::Error for TryError {} #[derive(Clone, Copy, Debug, PartialEq, Eq)] #[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] @@ -282,18 +282,18 @@ impl SeqEx { &mut self.send_window[seq_no as usize % self.send_window.len()] } #[inline] - fn is_full_inner(&self, is_for_send: bool, reply_no: Option) -> Result<(), PumpError> { + fn is_full_inner(&self, reply_no: Option) -> Result<(), TryError> { if self.concurrent_replies_total >= self.concurrent_replies.len() - 2 { - return Err(PumpError::WaitingForReply); + return Err(TryError::WaitingForReply); } let reply_idx = reply_no.map_or(self.send_window.len(), |r| r as usize % self.send_window.len()); - for i in 0..1 + is_for_send as u32 + self.concurrent_replies_total as u32 { + for i in 0..1 + self.concurrent_replies_total as u32 { let idx = self.next_send_seq_no.wrapping_add(i) as usize % self.send_window.len(); if self.send_window[idx].is_some() && reply_idx != idx { - return if i < 1 + is_for_send as u32 { - Err(PumpError::WaitingForRecv) + return if i == 0 { + Err(TryError::WaitingForRecv) } else { - Err(PumpError::WaitingForReply) + Err(TryError::WaitingForReply) } } } @@ -348,7 +348,7 @@ impl SeqEx { /// `retry_interval`. `current_time` does not have to be monotonically increasing. /// /// Can mutate `next_service_timestamp`. - pub fn try_send_direct(&mut self, current_time: i64, seq_cst: bool, packet_data: SendData) -> Result, (PumpError, SendData)> { + pub fn try_send_direct(&mut self, current_time: i64, seq_cst: bool, packet_data: SendData) -> Result, (TryError, SendData)> { let mut tmp = Some(packet_data); self.try_send_direct_with(current_time, seq_cst, |_| tmp.take().unwrap()) .map_err(|e| e.0) @@ -360,8 +360,8 @@ impl SeqEx { current_time: i64, seq_cst: bool, packet_data: F, - ) -> Result, (PumpError, F)> { - if let Err(e) = self.is_full_inner(true, None) { + ) -> Result, (TryError, F)> { + if let Err(e) = self.is_full_inner(None) { return Err((e, packet_data)); } @@ -404,7 +404,7 @@ impl SeqEx { } /// If this returns `Ok` then `try_send` might succeed on next call. - pub fn receive_raw_and_direct>(&mut self, packet: Packet

) -> Result<(RecvOkRaw, bool), DirectRecvError> { + pub fn receive_raw_and_direct>(&mut self, packet: Packet

) -> Result<(RecvOkRaw, bool), TryRecvError> { let seq_cst = packet.is_seq_cst(); let (seq_no, reply_no, recv_data) = match packet { Payload(seq_no, recv_data) | SeqCstPayload(seq_no, recv_data) => (seq_no, None, recv_data), @@ -413,7 +413,7 @@ impl SeqEx { return self .take_send(reply_no) .map(|send_data| (RecvOkRaw::Ack { send_data }, false)) - .ok_or(DirectRecvError::DroppedDuplicate) + .ok_or(TryRecvError::DroppedDuplicate) } }; // We only want to accept packets with sequence numbers in the range: @@ -436,17 +436,17 @@ impl SeqEx { // resending the packet. for entry in self.send_window.iter().flatten() { if entry.reply_no == Some(seq_no) { - return Err(DirectRecvError::DroppedDuplicate); + return Err(TryRecvError::DroppedDuplicate); } } for i in 0..self.concurrent_replies_total { if self.concurrent_replies[i] == seq_no { - return Err(DirectRecvError::DroppedDuplicate); + return Err(TryRecvError::DroppedDuplicate); } } - return Err(DirectRecvError::DroppedDuplicateResendAck(seq_no)); + return Err(TryRecvError::DroppedDuplicateResendAck(seq_no)); } else if is_above_range { - return Err(DirectRecvError::DroppedTooEarly); + return Err(TryRecvError::DroppedTooEarly); } // Check whether or not we've already received this packet @@ -462,19 +462,19 @@ impl SeqEx { // 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 (wait_recv, wait_reply) = match self.is_full_inner(false, reply_no) { + let (wait_recv, wait_reply) = match self.is_full_inner(reply_no) { Ok(()) => (seq_cst && !is_next, seq_cst && self.is_locked), - Err(PumpError::WaitingForRecv) => (true, false), - Err(PumpError::WaitingForReply) => (false, true), + Err(TryError::WaitingForRecv) => (true, false), + Err(TryError::WaitingForReply) => (false, true), }; if wait_recv || wait_reply { if !is_duplicate { self.recv_window[i] = RecvEntry::Occupied { seq_no, reply_no, seq_cst, data: recv_data.into() } } return if wait_recv { - Err(DirectRecvError::WaitingForRecv) + Err(TryRecvError::WaitingForRecv) } else { - Err(DirectRecvError::WaitingForReply) + Err(TryRecvError::WaitingForReply) }; } @@ -494,19 +494,19 @@ impl SeqEx { self.is_locked = true; } if let Some(send_data) = reply_no.and_then(|r| self.take_send(r)) { - Ok((RecvOkRaw::Reply { reply_no: seq_no, recv_data, send_data, seq_cst }, do_pump, true)) + Ok((RecvOkRaw::Reply { reply_no: seq_no, recv_data, send_data, seq_cst }, do_pump)) } else { - Ok((RecvOkRaw::Payload { reply_no: seq_no, recv_data, seq_cst }, do_pump, false)) + Ok((RecvOkRaw::Payload { reply_no: seq_no, recv_data, seq_cst }, do_pump)) } } /// If this returns `Ok` then `try_send` might succeed on next call. - pub fn try_pump_raw(&mut self) -> Result<(RecvOkRaw, bool), PumpError> { + pub fn try_pump_raw(&mut self) -> Result<(RecvOkRaw, bool), TryError> { let next_seq_no = self.next_recv_seq_no; let i = next_seq_no as usize % self.recv_window.len(); if let RecvEntry::Occupied { seq_no, reply_no, seq_cst, .. } = &self.recv_window[i] { debug_assert_eq!(*seq_no, next_seq_no); // We cannot safely reserve a reply no if the window is full. - self.is_full_inner(false, *reply_no)?; + self.is_full_inner(*reply_no)?; if !*seq_cst || !self.is_locked { let mut entry = RecvEntry::Empty; @@ -530,10 +530,10 @@ impl SeqEx { unreachable!(); } } else { - return Err(PumpError::WaitingForReply); + return Err(TryError::WaitingForReply); } } - Err(PumpError::WaitingForRecv) + Err(TryError::WaitingForRecv) } /// 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. diff --git a/src/single_thread.rs b/src/single_thread.rs index bdea3bb..ca68ad9 100644 --- a/src/single_thread.rs +++ b/src/single_thread.rs @@ -1,4 +1,4 @@ -use crate::{DirectRecvError, Packet, PumpError, RecvOkRaw, SeqEx, SeqNo, TransportLayer, DEFAULT_WINDOW_CAP}; +use crate::{TryRecvError, Packet, TryError, RecvOkRaw, SeqEx, SeqNo, TransportLayer, DEFAULT_WINDOW_CAP}; pub struct ReplyGuard<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { seq: &'a mut SeqEx, @@ -73,14 +73,14 @@ impl std::fmt::Display for RecvError { #[cfg(feature = "std")] impl std::error::Error for RecvError {} -pub enum RecvOk<'a, TL: TransportLayer, P, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { +pub enum RecvOk<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { Payload { reply_guard: ReplyGuard<'a, TL, SendData, RecvData, CAP>, - recv_data: P, + recv_data: RecvData, }, Reply { reply_guard: ReplyGuard<'a, TL, SendData, RecvData, CAP>, - recv_data: P, + recv_data: RecvData, send_data: SendData, }, Ack { @@ -90,8 +90,8 @@ pub enum RecvOk<'a, TL: TransportLayer, P, SendData, RecvData, const C 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> + impl<'a, TL: TransportLayer, SendData: std::fmt::Debug, RecvData: std::fmt::Debug, const CAP: usize> std::fmt::Debug + for $recv<'a, TL, SendData, RecvData, CAP> { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { @@ -110,8 +110,8 @@ macro_rules! impl_recvok { } } } - impl<'a, TL: TransportLayer, P, SendData, RecvData, const CAP: usize> $recv<'a, TL, P, SendData, RecvData, CAP> { - fn from_raw(seq: $seq_ex, app: TL, value: RecvOkRaw) -> Self { + impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> $recv<'a, TL, SendData, RecvData, CAP> { + fn from_raw(seq: $seq_ex, app: TL, value: RecvOkRaw) -> Self { match value { RecvOkRaw::Payload { reply_no, seq_cst, recv_data } => Self::Payload { reply_guard: ReplyGuard { seq, app: Some(app), reply_no, is_holding_lock: seq_cst }, @@ -125,14 +125,14 @@ macro_rules! impl_recvok { RecvOkRaw::Ack { send_data } => Self::Ack { send_data }, } } - pub fn consume(self) -> (Option<(ReplyGuard<'a, TL, SendData, RecvData, CAP>, P)>, Option) { + pub fn consume(self) -> (Option<(ReplyGuard<'a, TL, SendData, RecvData, CAP>, RecvData)>, 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 { + pub fn new(recv_data: Option<(ReplyGuard<'a, TL, SendData, RecvData, CAP>, RecvData)>, 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 }), @@ -141,15 +141,6 @@ macro_rules! impl_recvok { } } } - 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); @@ -157,7 +148,7 @@ pub(crate) use impl_recvok; impl SeqEx { /// Can mutate `next_service_timestamp`. - pub fn try_send(&mut self, mut app: impl TransportLayer, seq_cst: bool, packet_data: SendData) -> Result<(), SendData> { + pub fn try_send(&mut self, mut app: impl TransportLayer, seq_cst: bool, packet_data: SendData) -> Result<(), (TryError, SendData)> { match self.try_send_direct(app.time(), seq_cst, packet_data) { Ok(p) => { app.send(p); @@ -172,7 +163,7 @@ impl SeqEx { mut app: impl TransportLayer, seq_cst: bool, packet_data: F, - ) -> Result<(), F> { + ) -> Result<(), (TryError, F)> { match self.try_send_direct_with(app.time(), seq_cst, packet_data) { Ok(p) => { app.send(p); @@ -186,17 +177,17 @@ impl SeqEx { &mut self, mut app: impl TransportLayer, packet: Packet

, - ) -> Result, RecvError> { + ) -> Result<(RecvOkRaw, bool), RecvError> { match self.receive_raw_and_direct(packet) { Ok(a) => Ok(a), - Err(DirectRecvError::DroppedDuplicateResendAck(reply_no)) => { + Err(TryRecvError::DroppedDuplicateResendAck(reply_no)) => { app.send(Packet::Ack(reply_no)); Err(RecvError::DroppedDuplicate) } - Err(DirectRecvError::DroppedTooEarly) => Err(RecvError::DroppedTooEarly), - Err(DirectRecvError::DroppedDuplicate) => Err(RecvError::DroppedDuplicate), - Err(DirectRecvError::WaitingForRecv) => Err(RecvError::WaitingForRecv), - Err(DirectRecvError::WaitingForReply) => Err(RecvError::WaitingForReply), + Err(TryRecvError::DroppedTooEarly) => Err(RecvError::DroppedTooEarly), + Err(TryRecvError::DroppedDuplicate) => Err(RecvError::DroppedDuplicate), + Err(TryRecvError::WaitingForRecv) => Err(RecvError::WaitingForRecv), + Err(TryRecvError::WaitingForReply) => Err(RecvError::WaitingForReply), } } /// Can mutate `next_service_timestamp`. @@ -229,14 +220,14 @@ impl SeqEx { } self.resend_interval.min(self.next_service_timestamp - current_time) } - pub fn receive, P: Into>( + pub fn receive>( &mut self, app: TL, - packet: Packet

, - ) -> Result, RecvError> { - self.receive_raw(app.clone(), packet).map(|r| RecvOk::from_raw(self, app, r)) + packet: Packet, + ) -> Result<(RecvOk<'_, TL, SendData, RecvData, CAP>, bool), RecvError> { + self.receive_raw(app.clone(), packet).map(|(r, do_pump)| (RecvOk::from_raw(self, app, r), do_pump)) } - pub fn try_pump>(&mut self, app: TL) -> Result, PumpError> { - self.try_pump_raw().map(|r| RecvOk::from_raw(self, app, r)) + pub fn try_pump>(&mut self, app: TL) -> Result<(RecvOk<'_, TL, SendData, RecvData, CAP>, bool), TryError> { + self.try_pump_raw().map(|(r, do_pump)| (RecvOk::from_raw(self, app, r), do_pump)) } } diff --git a/src/sync.rs b/src/sync.rs index c7b5e20..c22d889 100644 --- a/src/sync.rs +++ b/src/sync.rs @@ -1,19 +1,26 @@ use std::{ sync::{ mpsc::{channel, Receiver, Sender}, - Condvar, Mutex, + Condvar, Mutex, MutexGuard, }, time::Instant, }; use crate::{ - Packet, PumpError, RecvError, RecvOkRaw, SeqEx, SeqNo, TransportLayer, DEFAULT_INITIAL_SEQ_NO, DEFAULT_RESEND_INTERVAL_MS, DEFAULT_WINDOW_CAP, + Packet, TryError, RecvError, RecvOkRaw, SeqEx, SeqNo, TransportLayer, DEFAULT_INITIAL_SEQ_NO, DEFAULT_RESEND_INTERVAL_MS, DEFAULT_WINDOW_CAP, }; pub struct SeqExSync { - seq_ex: Mutex<(SeqEx, usize, bool)>, - send_block: Condvar, - reply_block: Condvar, + inner: Mutex>, + wait_on_recv: Condvar, + wait_on_reply_sender: Condvar, + wait_on_reply_receiver: Condvar, +} +struct SeqExInner { + seq: SeqEx, + recv_waiters: usize, + reply_sender_waiters: bool, + reply_receiver_waiters: bool, } pub struct ReplyGuard<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { @@ -25,15 +32,11 @@ pub struct ReplyGuard<'a, TL: TransportLayer, SendData, RecvData, cons impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> ReplyGuard<'a, TL, SendData, RecvData, CAP> { fn reply_with(mut self, seq_cst: bool, packet_data: impl FnOnce(SeqNo, SeqNo) -> SendData) { let app = self.app.take().unwrap(); - let mut seq = self.seq.seq_ex.lock().unwrap(); - let seq_no = seq.0.seq_no(); - seq.0 + let mut inner = self.seq.inner.lock().unwrap(); + let seq_no = inner.seq.seq_no(); + inner.seq .reply_raw(app, self.reply_no, self.is_holding_lock, seq_cst, packet_data(seq_no, self.reply_no)); - if seq.2 { - seq.2 = false; - drop(seq); - self.seq.reply_block.notify_all(); - } + self.seq.notify_reply(inner); } /// 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. @@ -54,13 +57,9 @@ impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> Rep impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> Drop for ReplyGuard<'a, TL, SendData, RecvData, CAP> { fn drop(&mut self) { if let Some(app) = self.app.take() { - let mut seq = self.seq.seq_ex.lock().unwrap(); - seq.0.ack_raw(app, self.reply_no, self.is_holding_lock); - if seq.2 { - seq.2 = false; - drop(seq); - self.seq.reply_block.notify_all(); - } + let mut inner = self.seq.inner.lock().unwrap(); + inner.seq.ack_raw(app, self.reply_no, self.is_holding_lock); + self.seq.notify_reply(inner); } } } @@ -73,14 +72,14 @@ impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> std } } -pub enum RecvOk<'a, TL: TransportLayer, P, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { +pub enum RecvOk<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { Payload { reply_guard: ReplyGuard<'a, TL, SendData, RecvData, CAP>, - recv_data: P, + recv_data: RecvData, }, Reply { reply_guard: ReplyGuard<'a, TL, SendData, RecvData, CAP>, - recv_data: P, + recv_data: RecvData, send_data: SendData, }, Ack { @@ -89,138 +88,167 @@ pub enum RecvOk<'a, TL: TransportLayer, P, SendData, RecvData, const C } crate::impl_recvok!(RecvOk, &'a SeqExSync); -pub struct RecvIter<'a, TL: TransportLayer, P, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { +pub struct RecvIter<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { seq: Option<&'a SeqExSync>, app: TL, - first: Option>, + first: Option>, blocking: bool, } impl SeqExSync { pub fn new(retry_interval: i64, initial_seq_no: SeqNo) -> Self { Self { - seq_ex: Mutex::new((SeqEx::new(retry_interval, initial_seq_no), 0, false)), - send_block: Condvar::default(), - reply_block: Condvar::default(), + inner: Mutex::new(SeqExInner { seq: SeqEx::new(retry_interval, initial_seq_no), recv_waiters: 0, reply_sender_waiters: false, reply_receiver_waiters: false }), + wait_on_recv: Condvar::default(), + wait_on_reply_receiver: Condvar::default(), + wait_on_reply_sender: Condvar::default(), + } + } + fn notify_reply(&self, mut inner: MutexGuard<'_, SeqExInner>) { + if inner.reply_receiver_waiters { + inner.reply_receiver_waiters = false; + drop(inner); + self.wait_on_reply_receiver.notify_all(); + } else if inner.reply_sender_waiters { + inner.reply_sender_waiters = false; + drop(inner); + self.wait_on_reply_sender.notify_all(); + } + } + fn notify_recv(&self, mut inner: MutexGuard<'_, SeqExInner>) { + if inner.recv_waiters > 0 { + inner.recv_waiters -= 1; + drop(inner); + self.wait_on_recv.notify_one(); } } - pub fn receive, P: Into>( + pub fn try_receive>( &self, app: TL, - packet: Packet

, - ) -> Result, RecvError> { - let mut seq = self.seq_ex.lock().unwrap(); - match seq.0.receive_raw(app.clone(), packet) { - Ok(r) => { - if seq.1 > 0 { - seq.1 -= 1; - drop(seq); - self.send_block.notify_one(); - } - Ok(RecvOk::from_raw(self, app, r)) + packet: Packet, + ) -> Result<(RecvOk<'_, TL, SendData, RecvData, CAP>, bool), RecvError> { + let mut inner = self.inner.lock().unwrap(); + match inner.seq.receive_raw(app.clone(), packet) { + Ok((r, do_pump)) => { + self.notify_recv(inner); + Ok((RecvOk::from_raw(self, app, r), do_pump)) } Err(e) => Err(e), } } - pub fn try_pump>(&self, app: TL) -> Result, PumpError> { - let mut seq = self.seq_ex.lock().unwrap(); - // Enforce that only one thread may pump at a time. - if seq.2 { - return Err(PumpError::WaitingForReply); + pub fn receive>( + &self, + app: TL, + packet: Packet, + ) -> Option<(RecvOk<'_, TL, SendData, RecvData, CAP>, bool)> { + let result = self.try_receive(app.clone(), packet); + if let Err(RecvError::WaitingForReply) = result { + self.pump(app) + } else { + result.ok() } - match seq.0.try_pump_raw() { - Ok(r) => { - if seq.1 > 0 { - seq.1 -= 1; - drop(seq); - self.send_block.notify_one(); - } - Ok(RecvOk::from_raw(self, app, r)) + } + pub fn try_pump>(&self, app: TL) -> Result<(RecvOk<'_, TL, SendData, RecvData, CAP>, bool), TryError> { + let mut inner = self.inner.lock().unwrap(); + match inner.seq.try_pump_raw() { + Ok((r, do_pump)) => { + self.notify_recv(inner); + Ok((RecvOk::from_raw(self, app, r), do_pump)) } Err(e) => Err(e), } } - pub fn pump>(&self, app: TL) -> Option> { - let mut seq = self.seq_ex.lock().unwrap(); - // Enforce that only one thread may pump at a time. - if seq.2 { + pub fn pump>(&self, app: TL) -> Option<(RecvOk<'_, TL, SendData, RecvData, CAP>, bool)> { + let mut inner = self.inner.lock().unwrap(); + // Enforce that only one thread may wait to pump at a time. + if inner.reply_receiver_waiters { return None; } loop { - match seq.0.try_pump_raw() { - Ok(r) => { - if seq.1 > 0 { - seq.1 -= 1; - drop(seq); - self.send_block.notify_one(); - } - return Some(RecvOk::from_raw(self, app, r)); + match inner.seq.try_pump_raw() { + Ok((r, do_pump)) => { + self.notify_recv(inner); + return Some((RecvOk::from_raw(self, app, r), do_pump)); } - Err(PumpError::WaitingForRecv) => return None, - Err(PumpError::WaitingForReply) => { - seq.2 = true; - seq = self.reply_block.wait(seq).unwrap(); + Err(TryError::WaitingForRecv) => return None, + Err(TryError::WaitingForReply) => { + inner.reply_receiver_waiters = true; + inner = self.wait_on_reply_receiver.wait(inner).unwrap(); } } } } - fn receive_all_inner, P: Into>( + pub fn receive_all>( &self, app: TL, - blocking: bool, - packet: Packet

, - ) -> RecvIter<'_, TL, P, SendData, RecvData, CAP> { - match self.receive(app.clone(), packet) { - Ok(r) => RecvIter { seq: Some(self), app, first: Some(r), blocking }, - Err(RecvError::WaitingForReply) if blocking => RecvIter { seq: Some(self), app, first: None, blocking }, - Err(_) => RecvIter { seq: None, app, first: None, blocking }, + packet: Packet, + ) -> RecvIter<'_, TL, SendData, RecvData, CAP> { + let ret = self.receive(app.clone(), packet); + if let Some((first, do_pump)) = ret { + RecvIter { seq: do_pump.then_some(self), app, first: Some(first), blocking: true } + } else { + RecvIter { seq: None, app, first: None, blocking: true } } } - pub fn receive_all, P: Into>( + pub fn try_receive_all>( &self, app: TL, - packet: Packet

, - ) -> RecvIter<'_, TL, P, SendData, RecvData, CAP> { - self.receive_all_inner(app, true, packet) - } - pub fn try_receive_all, P: Into>( - &self, - app: TL, - packet: Packet

, - ) -> RecvIter<'_, TL, P, SendData, RecvData, CAP> { - self.receive_all_inner(app, false, packet) + packet: Packet, + ) -> RecvIter<'_, TL, SendData, RecvData, CAP> { + let ret = self.try_receive(app.clone(), packet); + if let Ok((first, do_pump)) = ret { + RecvIter { seq: do_pump.then_some(self), app, first: Some(first), blocking: false } + } else { + RecvIter { seq: None, app, first: None, blocking: false } + } } - pub fn try_send_with, F: FnOnce(SeqNo) -> SendData>(&self, app: TL, seq_cst: bool, packet_data: F) -> Result<(), F> { - let mut seq = self.seq_ex.lock().unwrap(); - seq.0.try_send_with(app, seq_cst, packet_data) + pub fn try_send_with, F: FnOnce(SeqNo) -> SendData>(&self, app: TL, seq_cst: bool, packet_data: F) -> Result<(), (TryError, F)> { + let mut inner = self.inner.lock().unwrap(); + inner.seq.try_send_with(app, seq_cst, packet_data) } - pub fn try_send>(&self, app: TL, seq_cst: bool, packet_data: SendData) -> Result<(), SendData> { - let mut seq = self.seq_ex.lock().unwrap(); - seq.0.try_send(app, seq_cst, packet_data) + pub fn try_send>(&self, app: TL, seq_cst: bool, packet_data: SendData) -> Result<(), (TryError, SendData)> { + let mut inner = self.inner.lock().unwrap(); + inner.seq.try_send(app, seq_cst, packet_data) } pub fn send_with>(&self, app: TL, seq_cst: bool, mut packet_data: impl FnOnce(SeqNo) -> SendData) { - let mut seq = self.seq_ex.lock().unwrap(); - while let Err(p) = seq.0.try_send_with(app.clone(), seq_cst, packet_data) { + let mut inner = self.inner.lock().unwrap(); + while let Err((e, p)) = inner.seq.try_send_with(app.clone(), seq_cst, packet_data) { packet_data = p; - seq.1 += 1; - seq = self.send_block.wait(seq).unwrap(); + match e { + TryError::WaitingForRecv => { + inner.recv_waiters += 1; + inner = self.wait_on_recv.wait(inner).unwrap(); + } + TryError::WaitingForReply => { + inner.reply_sender_waiters = true; + inner = self.wait_on_reply_sender.wait(inner).unwrap(); + } + } } } pub fn send>(&self, app: TL, seq_cst: bool, mut packet_data: SendData) { - let mut seq = self.seq_ex.lock().unwrap(); - while let Err(p) = seq.0.try_send(app.clone(), seq_cst, packet_data) { + let mut inner = self.inner.lock().unwrap(); + while let Err((e, p)) = inner.seq.try_send(app.clone(), seq_cst, packet_data) { packet_data = p; - seq.1 += 1; - seq = self.send_block.wait(seq).unwrap(); + match e { + TryError::WaitingForRecv => { + inner.recv_waiters += 1; + inner = self.wait_on_recv.wait(inner).unwrap(); + } + TryError::WaitingForReply => { + inner.reply_sender_waiters = true; + inner = self.wait_on_reply_sender.wait(inner).unwrap(); + } + } } } pub fn service>(&self, app: TL) -> i64 { - self.seq_ex.lock().unwrap().0.service(app) + self.inner.lock().unwrap().seq.service(app) } } impl Default for SeqExSync { @@ -229,23 +257,26 @@ impl Default for SeqExSync, P, SendData, RecvData, const CAP: usize> RecvIter<'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 RecvIter<'a, TL, P, SendData, RecvData, CAP> +impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> Iterator + for RecvIter<'a, TL, SendData, RecvData, CAP> { - type Item = RecvOk<'a, TL, RecvData, SendData, RecvData, CAP>; + type Item = RecvOk<'a, TL, SendData, RecvData, CAP>; fn next(&mut self) -> Option { - if let Some(g) = self.first.take() { - Some(g.into()) + if let Some(item) = self.first.take() { + Some(item) } else if let Some(origin) = self.seq { - if self.blocking { + let ret = if self.blocking { origin.pump(self.app.clone()) } else { origin.try_pump(self.app.clone()).ok() + }; + if let Some((item, do_pump)) = ret { + if !do_pump { + self.seq = None; + } + Some(item) + } else { + None } } else { None diff --git a/src/tokio.rs b/src/tokio.rs index ee68bab..6be57c4 100644 --- a/src/tokio.rs +++ b/src/tokio.rs @@ -4,7 +4,7 @@ use tokio::{ time, }; -use crate::{Packet, PumpError, RecvOkRaw, SeqEx, SeqNo, TransportLayer, DEFAULT_INITIAL_SEQ_NO, DEFAULT_RESEND_INTERVAL_MS, DEFAULT_WINDOW_CAP}; +use crate::{Packet, TryError, RecvOkRaw, SeqEx, SeqNo, TransportLayer, DEFAULT_INITIAL_SEQ_NO, DEFAULT_RESEND_INTERVAL_MS, DEFAULT_WINDOW_CAP}; type SendData = (oneshot::Sender>, Payload); @@ -211,7 +211,7 @@ impl SeqExTokio { &self, app: TL, blocking: bool, - ) -> Result<(ReplyGuard<'_, TL, Payload, CAP>, Payload), PumpError> { + ) -> Result<(ReplyGuard<'_, TL, Payload, CAP>, Payload), TryError> { let mut seq = self.seq_ex.lock().unwrap(); // Enforce that only one thread may pump at a time. loop { @@ -225,23 +225,23 @@ impl SeqExTokio { return Ok(ret); } } - Err(PumpError::WaitingForRecv) => return Err(PumpError::WaitingForRecv), - Err(PumpError::WaitingForReply) => { + Err(TryError::WaitingForRecv) => return Err(TryError::WaitingForRecv), + Err(TryError::WaitingForReply) => { seq.2 |= blocking; - return Err(PumpError::WaitingForReply); + return Err(TryError::WaitingForReply); } } } } - pub fn try_pump>(&self, app: TL) -> Result<(ReplyGuard<'_, TL, Payload, CAP>, Payload), PumpError> { + pub fn try_pump>(&self, app: TL) -> Result<(ReplyGuard<'_, TL, Payload, CAP>, Payload), TryError> { self.try_pump_inner(app, false) } pub async fn pump>(&self, app: TL) -> Option<(ReplyGuard<'_, TL, Payload, CAP>, Payload)> { loop { match self.try_pump_inner(app.clone(), true) { Ok(ret) => return Some(ret), - Err(PumpError::WaitingForRecv) => return None, - Err(PumpError::WaitingForReply) => { + Err(TryError::WaitingForRecv) => return None, + Err(TryError::WaitingForReply) => { self.reply_block.notified().await; } }