From fd85f8e697349a067846ae677e904e50be65d405 Mon Sep 17 00:00:00 2001 From: Monica Moniot Date: Fri, 18 Aug 2023 01:10:47 -0400 Subject: [PATCH] massive improvement --- examples/calculator.rs | 51 +++++------ examples/file_download.rs | 57 +++++------- examples/hello_world.rs | 69 +++++++-------- src/seq_queue.rs | 112 +++++++++++++++--------- src/single_thread.rs | 114 ++++++++++++++++-------- src/sync.rs | 178 +++++++++++++++++++++++--------------- src/transport_layer.rs | 2 +- 7 files changed, 338 insertions(+), 245 deletions(-) diff --git a/examples/calculator.rs b/examples/calculator.rs index fd4df9a..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, PacketOwned, 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 { - PacketOwned::Ack(reply_no) => { - let _ = seq.receive_ack(reply_no); - } - 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); - } + 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 3f0bd5c..0319dfe 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::{PacketOwned, RecvSuccess, ReplyGuard, SeqExSync}, - 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,18 +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, packet: seq_ex::Packet<'_, Packet>) { - let p = PacketOwned::from(packet); - 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); } } @@ -53,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(); @@ -70,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(&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() { @@ -92,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)); } + _ => {} } } @@ -103,17 +101,10 @@ fn receive(peer: &Peer) { if drop_packet() { continue; } - let parsed_packet = serde_json::from_slice::>(&packet); - match parsed_packet { - Ok(PacketOwned::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(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); - } - } - _ => {} } } } @@ -146,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 b2bc0d0..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, PacketOwned, 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() { - PacketOwned::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"); } - } - 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) + (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 689f994..1992e78 100644 --- a/src/seq_queue.rs +++ b/src/seq_queue.rs @@ -105,14 +105,47 @@ pub enum Error { WindowIsFull, } -pub enum Packet<'a, SendData> { +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] +pub enum Packet { + Payload(SeqNo, RecvData), + Reply(SeqNo, SeqNo, RecvData), + Ack(SeqNo), +} +impl Packet { + pub fn as_ref(&self) -> Packet<&RecvData> { + match self { + Packet::Payload(seq_no, data) => Packet::Payload(*seq_no, data), + Packet::Reply(seq_no, reply_no, data) => Packet::Reply(*seq_no, *reply_no, data), + Packet::Ack(reply_no) => Packet::Ack(*reply_no), + } + } + pub fn map(self, f: impl FnOnce(RecvData) -> SendData) -> Packet { + match self { + Packet::Payload(seq_no, data) => Packet::Payload(seq_no, f(data)), + Packet::Reply(seq_no, reply_no, data) => Packet::Reply(seq_no, reply_no, f(data)), + Packet::Ack(reply_no) => Packet::Ack(reply_no), + } + } +} +impl Packet<&RecvData> { + pub fn cloned(&self) -> Packet { + self.map(|d| d.clone()) + } +} + +pub enum RecvOkRaw { Payload { - seq_no: SeqNo, - reply_no: Option, - data: &'a SendData, + reply_no: SeqNo, + recv_data: RecvData, + }, + Reply { + reply_no: SeqNo, + recv_data: RecvData, + send_data: SendData, }, Ack { - reply_no: SeqNo, + send_data: SendData, }, } @@ -232,7 +265,7 @@ impl SeqEx { /// 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_direct(&mut self, packet_data: SendData, current_time: i64) -> Result, SendData> { + pub fn try_send_direct(&mut self, packet_data: SendData, current_time: i64) -> Result, SendData> { if self.is_full() { return Err(packet_data); } @@ -247,20 +280,21 @@ 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(entry.seq_no, &entry.data)) } /// If this returns `Ok` then `try_send` might succeed on next call. - pub fn receive_raw_and_direct>( - &mut self, - seq_no: SeqNo, - reply_no: Option, - packet: P, - ) -> Result<(SeqNo, P, Option), DirectError> { + pub fn receive_raw_and_direct>(&mut self, packet: Packet

) -> Result, DirectError> { + let (seq_no, reply_no, recv_data) = match packet { + Packet::Payload(seq_no, recv_data) => (seq_no, None, recv_data), + Packet::Reply(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 @@ -320,10 +354,13 @@ impl SeqEx { 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)) + return 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 } + } else { + RecvOkRaw::Payload { reply_no: seq_no, recv_data } + }); } else { - self.recv_window[i] = Some(RecvEntry { seq_no, reply_no, data: packet.into() }); + self.recv_window[i] = Some(RecvEntry { seq_no, reply_no, data: recv_data.into() }); if is_full { Err(DirectError::WindowIsFull) } else { @@ -332,11 +369,7 @@ impl SeqEx { } } /// 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) - } - /// 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, Error> { let next_seq_no = self.pre_recv_seq_no.wrapping_add(1); let i = next_seq_no as usize % self.recv_window.len(); @@ -349,8 +382,11 @@ impl SeqEx { 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)) + return Ok(if let Some(send_data) = entry.reply_no.and_then(|r| self.take_send(r)) { + RecvOkRaw::Reply { reply_no: entry.seq_no, recv_data: entry.data, send_data } + } else { + RecvOkRaw::Payload { reply_no: entry.seq_no, recv_data: entry.data } + }); } else { Err(Error::OutOfSequence) } @@ -364,7 +400,7 @@ impl SeqEx { /// and since each fragment will be received in order it will be trivial for them to reconstruct /// the original file. #[must_use] - pub fn reply_raw_and_direct(&mut self, reply_no: SeqNo, packet_data: SendData, current_time: i64) -> Option> { + pub fn reply_raw_and_direct(&mut self, reply_no: SeqNo, 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); @@ -382,24 +418,20 @@ impl SeqEx { data: packet_data, }); - Some(Packet::Payload { - seq_no: entry.seq_no, - reply_no: entry.reply_no, - data: &entry.data, - }) + Some(Packet::Reply(entry.seq_no, reply_no, &entry.data)) } else { None } } - pub fn ack_raw_and_direct(&mut self, reply_no: SeqNo) -> Option> { + pub fn ack_raw_and_direct(&mut self, reply_no: SeqNo) -> Option> { if self.remove_reservation(reply_no) { - Some(Packet::Ack { reply_no }) + Some(Packet::Ack(reply_no)) } else { None } } - pub fn service_direct<'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 }); while let Some(entry) = self.send_window.get(iter.idx) { @@ -409,10 +441,10 @@ impl SeqEx { 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(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); diff --git a/src/single_thread.rs b/src/single_thread.rs index 15ba365..d9924ac 100644 --- a/src/single_thread.rs +++ b/src/single_thread.rs @@ -1,10 +1,10 @@ -use crate::{DirectError, Error, Packet, SeqEx, SeqNo, TransportLayer, DEFAULT_WINDOW_CAP}; +use crate::{DirectError, Error, Packet, RecvOkRaw, SeqEx, SeqNo, TransportLayer, DEFAULT_WINDOW_CAP}; -pub struct ReplyGuard<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP>( - &'a mut SeqEx, - Option, - SeqNo, -); +pub struct ReplyGuard<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { + seq: &'a mut SeqEx, + app: Option, + reply_no: SeqNo, +} 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. @@ -13,25 +13,80 @@ impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> Rep /// the original file. pub fn reply(mut self, packet_data: SendData) { let mut app = None; - core::mem::swap(&mut app, &mut self.1); - self.0.reply_raw(app.unwrap(), self.2, packet_data); + core::mem::swap(&mut app, &mut self.app); + self.seq.reply_raw(app.unwrap(), self.reply_no, 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(app) = &mut self.1 { - if let Some(p) = self.0.ack_raw_and_direct(self.2) { + if let Some(app) = &mut self.app { + if let Some(p) = self.seq.ack_raw_and_direct(self.reply_no) { 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)] +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, + }, +} +impl<'a, TL: TransportLayer, P, SendData, RecvData, const CAP: usize> RecvOk<'a, TL, P, SendData, RecvData, CAP> { + pub fn from_raw(seq: &'a mut SeqEx, app: TL, value: RecvOkRaw) -> Self { + match value { + RecvOkRaw::Payload { reply_no, recv_data } => RecvOk::Payload { + reply_guard: ReplyGuard { seq, app: Some(app), reply_no }, + recv_data, + }, + RecvOkRaw::Reply { reply_no, recv_data, send_data } => RecvOk::Reply { + reply_guard: ReplyGuard { seq, app: Some(app), reply_no }, + recv_data, + send_data, + }, + RecvOkRaw::Ack { send_data } => RecvOk::Ack { send_data }, + } + } + pub fn consume(self) -> (Option<(ReplyGuard<'a, TL, SendData, RecvData, CAP>, P)>, Option) { + match self { + RecvOk::Payload { reply_guard, recv_data } => (Some((reply_guard, recv_data)), None), + RecvOk::Reply { reply_guard, recv_data, send_data } => (Some((reply_guard, recv_data)), Some(send_data)), + RecvOk::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(RecvOk::Payload { reply_guard, recv_data }), + (Some((reply_guard, recv_data)), Some(send_data)) => Some(RecvOk::Reply { reply_guard, recv_data, send_data }), + (None, Some(send_data)) => Some(RecvOk::Ack { send_data }), + (None, None) => None, + } + } +} +impl<'a, TL: TransportLayer, P: Into, SendData, RecvData, const CAP: usize> RecvOk<'a, TL, P, SendData, RecvData, CAP> { + pub fn into(self) -> RecvOk<'a, TL, RecvData, SendData, RecvData, CAP> { + match self { + RecvOk::Payload { reply_guard, recv_data } => RecvOk::Payload { reply_guard, recv_data: recv_data.into() }, + RecvOk::Reply { reply_guard, recv_data, send_data } => RecvOk::Reply { reply_guard, recv_data: recv_data.into(), send_data }, + RecvOk::Ack { send_data } => RecvOk::Ack { send_data }, + } + } } impl SeqEx { @@ -48,14 +103,12 @@ impl SeqEx { pub fn receive_raw>( &mut self, mut app: impl TransportLayer, - seq_no: SeqNo, - reply_no: Option, - packet: P, - ) -> Result<(SeqNo, P, Option), Error> { - match self.receive_raw_and_direct(seq_no, reply_no, packet) { + 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 }); + app.send(Packet::Ack(reply_no)); Err(Error::OutOfSequence) } Err(DirectError::OutOfSequence) => Err(Error::OutOfSequence), @@ -83,22 +136,11 @@ impl SeqEx { 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, Some(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, Some(app), reply_no), - packet, - send_data, - }) + pub fn pump>(&mut self, app: TL) -> Result, Error> { + self.pump_raw().map(|r| RecvOk::from_raw(self, app, r)) } } diff --git a/src/sync.rs b/src/sync.rs index 8830801..69ee58e 100644 --- a/src/sync.rs +++ b/src/sync.rs @@ -7,7 +7,7 @@ use std::{ time::Instant, }; -use crate::{Error, Packet, SeqEx, SeqNo, TransportLayer, DEFAULT_INITIAL_SEQ_NO, DEFAULT_RESEND_INTERVAL_MS, DEFAULT_WINDOW_CAP}; +use crate::{Error, Packet, RecvOkRaw, SeqEx, SeqNo, TransportLayer, DEFAULT_INITIAL_SEQ_NO, DEFAULT_RESEND_INTERVAL_MS, DEFAULT_WINDOW_CAP}; pub struct SeqExSync { seq_ex: Mutex<(SeqEx, usize)>, @@ -15,7 +15,7 @@ pub struct SeqExSync } pub struct ReplyGuard<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { - origin: &'a SeqExSync, + seq: &'a SeqExSync, app: TL, reply_no: SeqNo, } @@ -26,12 +26,12 @@ 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(); + let mut seq = self.seq.lock(); seq.reply_raw(self.app.clone(), self.reply_no, packet_data); core::mem::forget(self); } pub fn reply_with(self, packet_data: impl FnOnce(SeqNo, SeqNo) -> SendData) { - let mut seq = self.origin.lock(); + 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)); core::mem::forget(self); @@ -39,21 +39,92 @@ 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) { - let mut seq = self.origin.lock(); + let mut seq = self.seq.lock(); seq.ack_raw(self.app.clone(), self.reply_no); } } - -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, +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 ReplyIter<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { +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, + }, +} +impl<'a, TL: TransportLayer, P: std::fmt::Debug, SendData: std::fmt::Debug, RecvData, const CAP: usize> std::fmt::Debug + for RecvOk<'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> RecvOk<'a, TL, P, SendData, RecvData, CAP> { + pub fn from_raw(seq: &'a SeqExSync, app: TL, value: RecvOkRaw) -> Self { + match value { + RecvOkRaw::Payload { reply_no, recv_data } => RecvOk::Payload { reply_guard: ReplyGuard { seq, app, reply_no }, recv_data }, + RecvOkRaw::Reply { reply_no, recv_data, send_data } => RecvOk::Reply { + reply_guard: ReplyGuard { seq, app, reply_no }, + recv_data, + send_data, + }, + RecvOkRaw::Ack { send_data } => RecvOk::Ack { send_data }, + } + } + pub fn consume(self) -> (Option<(ReplyGuard<'a, TL, SendData, RecvData, CAP>, P)>, Option) { + match self { + RecvOk::Payload { reply_guard, recv_data } => (Some((reply_guard, recv_data)), None), + RecvOk::Reply { reply_guard, recv_data, send_data } => (Some((reply_guard, recv_data)), Some(send_data)), + RecvOk::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(RecvOk::Payload { reply_guard, recv_data }), + (Some((reply_guard, recv_data)), Some(send_data)) => Some(RecvOk::Reply { reply_guard, recv_data, send_data }), + (None, Some(send_data)) => Some(RecvOk::Ack { send_data }), + (None, None) => None, + } + } +} +impl<'a, TL: TransportLayer, P: Into, SendData, RecvData, const CAP: usize> RecvOk<'a, TL, P, SendData, RecvData, CAP> { + pub fn into(self) -> RecvOk<'a, TL, RecvData, SendData, RecvData, CAP> { + match self { + RecvOk::Payload { reply_guard, recv_data } => RecvOk::Payload { reply_guard, recv_data: recv_data.into() }, + RecvOk::Reply { reply_guard, recv_data, send_data } => RecvOk::Reply { reply_guard, recv_data: recv_data.into(), send_data }, + RecvOk::Ack { send_data } => RecvOk::Ack { send_data }, + } + } +} + +pub struct ReplyIter<'a, TL: TransportLayer, P: Into, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { origin: Option<&'a SeqExSync>, app: TL, - first: Option>, + first: Option>, } pub struct SeqExGuard<'a, SendData, RecvData, const CAP: usize>(MutexGuard<'a, (SeqEx, usize)>); @@ -81,42 +152,30 @@ impl SeqExSync { pub fn receive, P: Into>( &self, app: TL, - seq_no: SeqNo, - reply_no: Option, - packet: P, - ) -> Result, Error> { + 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); + let ret = seq.0.receive_raw(app.clone(), packet); if seq.1 > 0 && ret.is_ok() { self.send_block.notify_one(); } - ret.map(|(reply_no, packet, send_data)| RecvSuccess { - guard: ReplyGuard { origin: self, app, reply_no }, - packet, - send_data, - }) + ret.map(|r| RecvOk::from_raw(self, app, r)) } - 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(); } - ret.map(|(reply_no, packet, send_data)| RecvSuccess { - guard: ReplyGuard { origin: self, app, reply_no }, - packet, - send_data, - }) + ret.map(|r| RecvOk::from_raw(self, app, r)) } - 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 { origin: Some(self), app, first: Some(r) } } else { ReplyIter { origin: None, app, first: None } } @@ -151,15 +210,6 @@ impl SeqExSync { 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 - } pub fn service>(&self, app: TL) -> i64 { self.lock().service(app) } @@ -174,11 +224,18 @@ 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) + Some(g.into()) } else if let Some(origin) = self.origin { origin.pump(self.app.clone()).ok() } else { @@ -187,35 +244,20 @@ impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> Ite } } -#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] -#[derive(Clone)] -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)] +#[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() } } } @@ -224,7 +266,7 @@ impl TransportLayer for &MpscTransport { self.time.elapsed().as_millis() as i64 } - fn send(&mut self, packet: Packet<'_, Payload>) { - let _ = self.channel.send(PacketOwned::from(packet)); + fn send(&mut self, packet: Packet<&Payload>) { + let _ = self.channel.send(packet.cloned()); } } diff --git a/src/transport_layer.rs b/src/transport_layer.rs index 62c19dc..ef67169 100644 --- a/src/transport_layer.rs +++ b/src/transport_layer.rs @@ -9,5 +9,5 @@ use crate::Packet; pub trait TransportLayer: Clone { fn time(&mut self) -> i64; - fn send(&mut self, packet: Packet<'_, SendData>); + fn send(&mut self, packet: Packet<&SendData>); }