From 63ea9c6ed24d784c57874774ac9c92477a8278ae Mon Sep 17 00:00:00 2001 From: Monica Moniot Date: Thu, 17 Aug 2023 22:25:22 -0400 Subject: [PATCH] integrated changes --- examples/calculator.rs | 8 ++--- examples/file_download.rs | 21 +++++-------- examples/hello_world.rs | 8 ++--- src/seq_queue.rs | 28 +++++++++++------ src/single_thread.rs | 65 +++++++++++++++++++++++++++++---------- src/sync.rs | 25 +++++++++------ src/tokio.rs | 40 ++++++++++++++---------- 7 files changed, 121 insertions(+), 74 deletions(-) diff --git a/examples/calculator.rs b/examples/calculator.rs index 63ca26d..fd4df9a 100644 --- a/examples/calculator.rs +++ b/examples/calculator.rs @@ -1,6 +1,6 @@ use std::{sync::mpsc::Receiver, thread, time::Duration}; -use seq_ex::sync::{MpscGuard, MpscSeqEx, MpscTransport, PacketType, RecvSuccess}; +use seq_ex::sync::{MpscGuard, MpscSeqEx, MpscTransport, PacketOwned, RecvSuccess}; #[derive(Clone)] enum Packet { @@ -27,14 +27,14 @@ fn process(_: MpscGuard<'_, Packet>, recv_packet: Packet, _: Option, val } } -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) => { + PacketOwned::Ack(reply_no) => { let _ = seq.receive_ack(reply_no); } - PacketType::Payload(seq_no, reply_no, payload) => { + PacketOwned::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); } diff --git a/examples/file_download.rs b/examples/file_download.rs index 8d2be99..3f0bd5c 100644 --- a/examples/file_download.rs +++ b/examples/file_download.rs @@ -11,8 +11,8 @@ use std::{ use rand_core::{OsRng, RngCore}; use seq_ex::{ - sync::{PacketType, RecvSuccess, ReplyGuard, SeqExSync}, - SeqNo, TransportLayer, + sync::{PacketOwned, RecvSuccess, ReplyGuard, SeqExSync}, + TransportLayer, }; use serde::{Deserialize, Serialize}; @@ -41,15 +41,8 @@ impl TransportLayer for &Transport { 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); + fn send(&mut self, packet: seq_ex::Packet<'_, Packet>) { + let p = PacketOwned::from(packet); if let Ok(p) = serde_json::to_vec(&p) { let _ = self.sender.send(p); } @@ -110,12 +103,12 @@ fn receive(peer: &Peer) { if drop_packet() { continue; } - let parsed_packet = serde_json::from_slice::>(&packet); + let parsed_packet = serde_json::from_slice::>(&packet); match parsed_packet { - Ok(PacketType::Ack(reply_no)) => { + Ok(PacketOwned::Ack(reply_no)) => { let _ = peer.seqex.receive_ack(reply_no); } - Ok(PacketType::Payload(seq_no, reply_no, payload)) => { + Ok(PacketOwned::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); } diff --git a/examples/hello_world.rs b/examples/hello_world.rs index 7202448..b2bc0d0 100644 --- a/examples/hello_world.rs +++ b/examples/hello_world.rs @@ -1,6 +1,6 @@ use std::sync::mpsc::Receiver; -use seq_ex::sync::{MpscGuard, MpscSeqEx, MpscTransport, PacketType, RecvSuccess}; +use seq_ex::sync::{MpscGuard, MpscSeqEx, MpscTransport, PacketOwned, RecvSuccess}; #[derive(Clone, Debug)] enum Packet { @@ -37,16 +37,16 @@ fn process(guard: MpscGuard<'_, Packet>, recv_packet: Packet, send_packet: Optio } } -fn receive(recv: &Receiver>, seq: &MpscSeqEx, transport: &MpscTransport) { +fn receive(recv: &Receiver>, seq: &MpscSeqEx, transport: &MpscTransport) { match recv.recv().unwrap() { - PacketType::Ack(reply_no) => { + PacketOwned::Ack(reply_no) => { let result = seq.receive_ack(reply_no); if let Ok(Exclamation) = result { // Our Hello World exchange ends right here. print!("\n"); } } - PacketType::Payload(seq_no, reply_no, payload) => { + PacketOwned::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) } diff --git a/src/seq_queue.rs b/src/seq_queue.rs index 6e3ecc2..4d5f58f 100644 --- a/src/seq_queue.rs +++ b/src/seq_queue.rs @@ -113,7 +113,7 @@ pub enum Packet<'a, SendData> { }, Ack { reply_no: SeqNo, - } + }, } /// An iterator over all packets in the send window. It will iterate over all packets currently @@ -247,7 +247,11 @@ impl SeqEx { debug_assert!(slot.is_none()); let entry = slot.insert(SendEntry { seq_no, reply_no: None, next_resend_time, data: packet_data }); - Ok(Packet::Payload { seq_no: entry.seq_no, reply_no: entry.reply_no, data: &entry.data }) + Ok(Packet::Payload { + seq_no: entry.seq_no, + reply_no: entry.reply_no, + data: &entry.data, + }) } /// If this returns `Ok` then `try_send` might succeed on next call. @@ -378,7 +382,11 @@ impl SeqEx { data: packet_data, }); - Some(Packet::Payload { seq_no: entry.seq_no, reply_no: entry.reply_no, data: &entry.data }) + Some(Packet::Payload { + seq_no: entry.seq_no, + reply_no: entry.reply_no, + data: &entry.data, + }) } else { None } @@ -391,19 +399,21 @@ impl SeqEx { } } - pub fn service<'a>(&'a mut self, current_time: i64, iter: &mut Option) -> Option> { + 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, - }); + 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(Packet::Payload { seq_no: entry.seq_no, reply_no: entry.reply_no, data: &entry.data }); + return Some(Packet::Payload { + seq_no: entry.seq_no, + reply_no: entry.reply_no, + data: &entry.data, + }); } else { iter.next_time = iter.next_time.min(entry.next_resend_time); } diff --git a/src/single_thread.rs b/src/single_thread.rs index 9514063..0681679 100644 --- a/src/single_thread.rs +++ b/src/single_thread.rs @@ -1,8 +1,8 @@ -use crate::{Error, SeqEx, SeqNo, TransportLayer, DEFAULT_WINDOW_CAP, Payload}; +use crate::{DirectError, Error, Packet, SeqEx, SeqNo, TransportLayer, DEFAULT_WINDOW_CAP}; pub struct ReplyGuard<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP>( &'a mut SeqEx, - TL, + Option, SeqNo, ); impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> ReplyGuard<'a, TL, SendData, RecvData, CAP> { @@ -12,14 +12,18 @@ 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(mut self, packet_data: SendData) { - self.0.reply_raw(self.1, self.2, packet_data); + let mut app = None; + core::mem::swap(&mut app, &mut self.1); + self.0.reply_raw(app.unwrap(), self.2, packet_data); 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) { - if let Some(p) = self.0.ack_direct(self.2) { - self.1.send(p) + if let Some(app) = &mut self.1 { + if let Some(p) = self.0.ack_direct(self.2) { + app.send(p) + } } } } @@ -31,8 +35,14 @@ pub struct RecvSuccess<'a, TL: TransportLayer, P: Into, Send } impl SeqEx { - pub fn try_send_raw(&mut self, packet_data: SendData, current_time: i64) -> Result, SendData> { - + pub fn try_send(&mut self, mut app: impl TransportLayer, packet_data: SendData) -> Result<(), SendData> { + match self.try_send_direct(packet_data, app.time()) { + Ok(p) => { + app.send(p); + Ok(()) + } + Err(e) => Err(e), + } } /// If this returns `Ok` then `try_send` might succeed on next call. pub fn receive_raw>( @@ -42,18 +52,34 @@ impl SeqEx { reply_no: Option, packet: P, ) -> Result<(SeqNo, P, Option), Error> { - let ret = self.receive_direct(seq_no, reply_no, packet); + match self.receive_direct(seq_no, reply_no, 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) => Err(Error::WindowIsFull), + } } pub fn reply_raw(&mut self, mut app: impl TransportLayer, reply_no: SeqNo, packet_data: SendData) { - if let Some(Payload { seq_no, reply_no, data }) = self.reply_direct(reply_no, packet_data, app.time()) { - app.send(seq_no, reply_no, data) + if let Some(p) = self.reply_direct(reply_no, packet_data, app.time()) { + app.send(p) } } pub fn ack_raw(&mut self, mut app: impl TransportLayer, reply_no: SeqNo) { - if self.ack_direct(reply_no) { - app.send_ack(reply_no) + if let Some(p) = self.ack_direct(reply_no) { + 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, @@ -61,11 +87,18 @@ impl SeqEx { reply_no: Option, packet: P, ) -> Result, Error> { - self.receive_raw(app, seq_no, reply_no, packet) - .map(|(reply_no, packet, send_data)| RecvSuccess { guard: ReplyGuard(self, app, reply_no), packet, send_data }) + self.receive_raw(app.clone(), seq_no, reply_no, packet) + .map(|(reply_no, packet, send_data)| RecvSuccess { + guard: ReplyGuard(self, Some(app), reply_no), + packet, + send_data, + }) } 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 }) + self.pump_raw().map(|(reply_no, packet, send_data)| RecvSuccess { + guard: ReplyGuard(self, Some(app), reply_no), + packet, + send_data, + }) } } diff --git a/src/sync.rs b/src/sync.rs index 33c23ef..8830801 100644 --- a/src/sync.rs +++ b/src/sync.rs @@ -7,7 +7,7 @@ use std::{ time::Instant, }; -use crate::{Error, SeqEx, SeqNo, TransportLayer, DEFAULT_INITIAL_SEQ_NO, DEFAULT_RESEND_INTERVAL_MS, DEFAULT_WINDOW_CAP}; +use crate::{Error, Packet, SeqEx, SeqNo, TransportLayer, DEFAULT_INITIAL_SEQ_NO, DEFAULT_RESEND_INTERVAL_MS, DEFAULT_WINDOW_CAP}; pub struct SeqExSync { seq_ex: Mutex<(SeqEx, usize)>, @@ -189,25 +189,33 @@ impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> Ite #[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] #[derive(Clone)] -pub enum PacketType { +pub enum PacketOwned { Payload(SeqNo, Option, Payload), Ack(SeqNo), } +impl<'a, Payload: Clone> From> for PacketOwned { + fn from(value: Packet<'a, Payload>) -> Self { + match value { + Packet::Payload { seq_no, reply_no, data } => PacketOwned::Payload(seq_no, reply_no, data.clone()), + Packet::Ack { reply_no } => PacketOwned::Ack(reply_no), + } + } +} #[derive(Clone)] 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 +224,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(PacketOwned::from(packet)); } } 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();