From d2205fd8c11b5c59582997629ba41142d8de84b7 Mon Sep 17 00:00:00 2001 From: Monica Moniot Date: Mon, 21 Aug 2023 11:10:42 -0400 Subject: [PATCH] finished tokio support --- Cargo.lock | 296 ++++++++++++++++ Cargo.toml | 6 +- async_file_share.rs | 250 ------------- examples/calculator.rs | 18 +- examples/file_download.rs | 25 +- examples/hello_world.rs | 18 +- examples/hello_world_tokio.rs | 118 +++++++ src/seq_queue.rs | 381 +++++++++++--------- src/single_thread.rs | 146 ++++---- src/sync.rs | 261 +++++++------- src/tokio.rs | 642 ++++++++++++++++++++-------------- 11 files changed, 1246 insertions(+), 915 deletions(-) delete mode 100644 async_file_share.rs create mode 100644 examples/hello_world_tokio.rs diff --git a/Cargo.lock b/Cargo.lock index 887d4c1..b981aac 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,6 +2,63 @@ # It is not intended for manual editing. version = 3 +[[package]] +name = "addr2line" +version = "0.20.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f4fa78e18c64fce05e902adecd7a5eed15a5e0a3439f7b0e169f0252214865e3" +dependencies = [ + "gimli", +] + +[[package]] +name = "adler" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f26201604c87b1e01bd3d98f8d5d9a8fcbb815e8cedb41ffccbeb4bf593a35fe" + +[[package]] +name = "autocfg" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d468802bab17cbc0cc575e9b053f41e72aa36bfa6b7f55e3529ffa43161b97fa" + +[[package]] +name = "backtrace" +version = "0.3.68" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4319208da049c43661739c5fade2ba182f09d1dc2299b32298d3a31692b17e12" +dependencies = [ + "addr2line", + "cc", + "cfg-if", + "libc", + "miniz_oxide", + "object", + "rustc-demangle", +] + +[[package]] +name = "bitflags" +version = "1.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bef38d45163c2f1dde094a7dfd33ccf595c92905c8f8f4fdc18d06fb1037718a" + +[[package]] +name = "bytes" +version = "1.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "89b2fd2a0dcf38d7971e2194b6b6eebab45ae01067456a7fd93d5547a61b70be" + +[[package]] +name = "cc" +version = "1.0.82" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "305fe645edc1442a0fa8b6726ba61d422798d37a52e12eaecf4b022ebbb88f01" +dependencies = [ + "libc", +] + [[package]] name = "cfg-if" version = "1.0.0" @@ -19,18 +76,114 @@ dependencies = [ "wasi", ] +[[package]] +name = "gimli" +version = "0.27.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6c80984affa11d98d1b88b66ac8853f143217b399d3c74116778ff8fdb4ed2e" + [[package]] name = "half" version = "1.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "eabb4a44450da02c90444cf74558da904edde8fb4e9035a9a6a4e15445af0bd7" +[[package]] +name = "hermit-abi" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "443144c8cdadd93ebf52ddb4056d257f5b52c04d3c804e657d19eb73fc33668b" + [[package]] name = "libc" version = "0.2.147" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b4668fb0ea861c1df094127ac5f1da3409a82116a4ba74fca2e58ef927159bb3" +[[package]] +name = "lock_api" +version = "0.4.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c1cc9717a20b1bb222f333e6a92fd32f7d8a18ddc5a3191a11af45dcbf4dcd16" +dependencies = [ + "autocfg", + "scopeguard", +] + +[[package]] +name = "memchr" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2dffe52ecf27772e601905b7522cb4ef790d2cc203488bbd0e2fe85fcb74566d" + +[[package]] +name = "miniz_oxide" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e7810e0be55b428ada41041c41f32c9f1a42817901b4ccf45fa3d4b6561e74c7" +dependencies = [ + "adler", +] + +[[package]] +name = "mio" +version = "0.8.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "927a765cd3fc26206e66b296465fa9d3e5ab003e651c1b3c060e7956d96b19d2" +dependencies = [ + "libc", + "wasi", + "windows-sys", +] + +[[package]] +name = "num_cpus" +version = "1.16.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4161fcb6d602d4d2081af7c3a45852d875a03dd337a6bfdd6e06407b61342a43" +dependencies = [ + "hermit-abi", + "libc", +] + +[[package]] +name = "object" +version = "0.31.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8bda667d9f2b5051b8833f59f3bf748b28ef54f850f4fcb389a252aa383866d1" +dependencies = [ + "memchr", +] + +[[package]] +name = "parking_lot" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3742b2c103b9f06bc9fff0a37ff4912935851bee6d36f3c02bcc755bcfec228f" +dependencies = [ + "lock_api", + "parking_lot_core", +] + +[[package]] +name = "parking_lot_core" +version = "0.9.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93f00c865fe7cabf650081affecd3871070f26767e7b2070a3ffae14c654b447" +dependencies = [ + "cfg-if", + "libc", + "redox_syscall", + "smallvec", + "windows-targets", +] + +[[package]] +name = "pin-project-lite" +version = "0.2.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "12cc1b0bf1727a77a54b6654e7b5f1af8604923edc8b81885f8ec92f9e3f0a05" + [[package]] name = "proc-macro2" version = "1.0.66" @@ -58,6 +211,27 @@ dependencies = [ "getrandom", ] +[[package]] +name = "redox_syscall" +version = "0.3.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "567664f262709473930a4bf9e51bf2ebf3348f2e748ccc50dea20646858f8f29" +dependencies = [ + "bitflags", +] + +[[package]] +name = "rustc-demangle" +version = "0.1.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d626bb9dae77e28219937af045c257c28bfd3f69333c512553507f5f9798cb76" + +[[package]] +name = "scopeguard" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" + [[package]] name = "seq_ex" version = "0.1.0" @@ -65,6 +239,7 @@ dependencies = [ "rand_core", "serde", "serde_cbor", + "tokio", ] [[package]] @@ -97,6 +272,31 @@ dependencies = [ "syn", ] +[[package]] +name = "signal-hook-registry" +version = "1.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d8229b473baa5980ac72ef434c4415e70c4b5e71b423043adb4ba059f89c99a1" +dependencies = [ + "libc", +] + +[[package]] +name = "smallvec" +version = "1.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "62bb4feee49fdd9f707ef802e22365a35de4b7b299de4763d44bfea899442ff9" + +[[package]] +name = "socket2" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2538b18701741680e0322a2302176d3253a35388e2e62f172f64f4f16605f877" +dependencies = [ + "libc", + "windows-sys", +] + [[package]] name = "syn" version = "2.0.28" @@ -108,6 +308,36 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "tokio" +version = "1.32.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "17ed6077ed6cd6c74735e21f37eb16dc3935f96878b1fe961074089cc80893f9" +dependencies = [ + "backtrace", + "bytes", + "libc", + "mio", + "num_cpus", + "parking_lot", + "pin-project-lite", + "signal-hook-registry", + "socket2", + "tokio-macros", + "windows-sys", +] + +[[package]] +name = "tokio-macros" +version = "2.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "630bdcf245f78637c13ec01ffae6187cca34625e8c63150d424b59e55af2675e" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "unicode-ident" version = "1.0.11" @@ -119,3 +349,69 @@ name = "wasi" version = "0.11.0+wasi-snapshot-preview1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9c8d87e72b64a3b4db28d11ce29237c246188f4f51057d65a7eab63b7987e423" + +[[package]] +name = "windows-sys" +version = "0.48.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "677d2418bec65e3338edb076e806bc1ec15693c5d0104683f2efe857f61056a9" +dependencies = [ + "windows-targets", +] + +[[package]] +name = "windows-targets" +version = "0.48.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d1eeca1c172a285ee6c2c84c341ccea837e7c01b12fbb2d0fe3c9e550ce49ec8" +dependencies = [ + "windows_aarch64_gnullvm", + "windows_aarch64_msvc", + "windows_i686_gnu", + "windows_i686_msvc", + "windows_x86_64_gnu", + "windows_x86_64_gnullvm", + "windows_x86_64_msvc", +] + +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.48.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b10d0c968ba7f6166195e13d593af609ec2e3d24f916f081690695cf5eaffb2f" + +[[package]] +name = "windows_aarch64_msvc" +version = "0.48.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "571d8d4e62f26d4932099a9efe89660e8bd5087775a2ab5cdd8b747b811f1058" + +[[package]] +name = "windows_i686_gnu" +version = "0.48.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2229ad223e178db5fbbc8bd8d3835e51e566b8474bfca58d2e6150c48bb723cd" + +[[package]] +name = "windows_i686_msvc" +version = "0.48.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "600956e2d840c194eedfc5d18f8242bc2e17c7775b6684488af3a9fff6fe3287" + +[[package]] +name = "windows_x86_64_gnu" +version = "0.48.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ea99ff3f8b49fb7a8e0d305e5aec485bd068c2ba691b6e277d29eaeac945868a" + +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.48.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f1a05a1ece9a7a0d5a7ccf30ba2c33e3a61a30e042ffd247567d1de1d94120d" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.48.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d419259aba16b663966e29e6d7c6ecfa0bb8425818bb96f6f1f3c3eb71a6e7b9" diff --git a/Cargo.toml b/Cargo.toml index 44a7e00..a8a5c42 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -10,13 +10,15 @@ path = "src/lib.rs" doc = true [features] -default = ["std", "serde"] +default = ["serde", "tokio"] +tokio = ["std", "dep:tokio"] std = [] [dependencies] serde = { version = "1.0.183", default-features = false, features = ["derive"], optional = true } -#tokio = { version = "1.31.0", default-features = false, features = ["sync", "rt", "time", "macros"], optional = true } +tokio = { version = "1.31.0", default-features = false, features = ["sync", "time"], optional = true } [dev-dependencies] rand_core = { version = "0.6.4", features = ["getrandom"]} serde_cbor = { version = "0.11.2" } +tokio = { version = "1.31.0", default-features = false, features = ["full"] } diff --git a/async_file_share.rs b/async_file_share.rs deleted file mode 100644 index f9ea0ae..0000000 --- a/async_file_share.rs +++ /dev/null @@ -1,250 +0,0 @@ -use std::{ - collections::HashMap, - fs::{read_dir, File}, - io::{Read, Write}, - ops::Deref, - path::PathBuf, - sync::{ - mpsc::{channel, Receiver, Sender}, - Arc, RwLock, - }, - thread, - time::{Duration, Instant}, -}; - -use rand_core::{OsRng, RngCore}; -use seq_ex::{SeqNo, TransportLayer, Packet}; -use serde::{Deserialize, Serialize}; - -/// serde_cbor minimal format is both smaller and faster than default format. -/// The overhead is ~33% faster. -fn to_writer_minimal(value: &impl serde::Serialize, w: &mut impl serde_cbor::ser::Write) -> serde_cbor::Result<()> { - value.serialize(&mut serde_cbor::Serializer::new(w).packed_format().legacy_enums())?; - Ok(()) -} - -fn drop_packet() -> bool { - OsRng.next_u32() >= (u32::MAX / 4 * 3) -} - -const FILE_CHUNK_SIZE: usize = 1000; -const DOWNLOAD_LIMIT: u64 = 1000000; - -#[derive(Debug)] -enum SendData { - RequestFile { filename: String }, - ConfirmFileSize { file: File }, - ConfirmDownload, - FileDownload, - ReadDir { download_missing: bool }, - DirContents, -} - -#[derive(Serialize, Deserialize)] -enum Payload<'a> { - RequestFile { filename: &'a str }, - ConfirmFileSize { filesize: u64 }, - ConfirmDownload { fileid: u64 }, - FileDownload { fileid: u64, file_chunk: &'a [u8] }, - FileDownloadComplete { fileid: u64 }, - ReadDir, - DirContents { filenames: Vec<&'a str> }, -} - -const PACKET_TYPE_PAYLOAD: u8 = 0; -const PACKET_TYPE_LOCK_PAYLOAD: u8 = 1; -const PACKET_TYPE_REPLY: u8 = 2; -const PACKET_TYPE_LOCK_REPLY: u8 = 3; -const PACKET_TYPE_ACK: u8 = 4; - -fn create_payload(seq_no: SeqNo, packet: &Payload<'_>) -> Vec { - let mut p = vec![PACKET_TYPE_PAYLOAD]; - p.extend(&seq_no.to_be_bytes()); - to_writer_minimal(packet, &mut p); - p -} -fn create_reply(seq_no: SeqNo, reply_no: SeqNo, packet: &Payload<'_>) -> Vec { - let mut p = vec![PACKET_TYPE_REPLY]; - p.extend(&seq_no.to_be_bytes()); - p.extend(&reply_no.to_be_bytes()); - to_writer_minimal(packet, &mut p); - p -} -macro_rules! reply { - ($guard:expr, $send_data:expr, $packet: expr) => { - $guard.reply_with(|seq_no, reply_no| ($send_data, create_reply(seq_no, reply_no, &$packet))) - }; -} -macro_rules! send { - ($peer:expr, $trans:expr, $send_data:expr, $packet: expr) => { - $peer.seqex.send_with($trans, |seq_no| ($send_data, create_payload(seq_no, &$packet))) - }; -} - -type SeqEx = seq_ex::sync::SeqExSync<(SendData, Vec), Vec>; -type ReplyGuard<'a> = seq_ex::sync::ReplyGuard<'a, &'a Transport, (SendData, Vec), Vec>; - -#[derive(Clone)] -struct Transport { - sender: Sender>, - time: Instant, -} - -struct Peer { - downloads_in_progress: HashMap, - home_dir: PathBuf, - seqex: SeqEx, -} - -impl TransportLayer<(SendData, Vec)> for &Transport { - fn time(&mut self) -> i64 { - self.time.elapsed().as_millis() as i64 - } - - fn send(&mut self, packet: Packet<&(SendData, Vec)>) { - let p = match packet.consume() { - Ok((_, packet)) => packet.clone(), - Err(reply_no) => { - let mut p = Vec::with_capacity(5); - p.push(PACKET_TYPE_ACK); - p.extend(&reply_no.to_be_bytes()); - p - } - }; - let _ = self.sender.send(p); - } -} - -fn process(peer: &Arc, transport: &Transport, guard: ReplyGuard<'_>, payload: Payload<'_>, data: Option) -> Option<()> { - match (payload, data) { - (Payload::RequestFile { filename }, None) => { - let file = File::open(peer.home_dir.join(filename)).ok()?; - let metadata = file.metadata().ok()?; - let filesize = metadata.len(); - reply!( - guard, - SendData::ConfirmFileSize { file }, - Payload::ConfirmFileSize { filesize } - ); - } - (Payload::ConfirmFileSize { filesize }, Some(SendData::RequestFile { filename })) => { - let path = peer.home_dir.join(filename); - if filesize > DOWNLOAD_LIMIT { - return None; - } - let file = File::create(&path).ok()?; - let fileid = OsRng.next_u64(); - peer.downloads_in_progress.insert(fileid, file); - reply!(guard, SendData::ConfirmDownload, Payload::ConfirmDownload { fileid }); - } - (Payload::ConfirmDownload { fileid }, Some(SendData::ConfirmFileSize { mut file })) => { - let peer = peer.clone(); - let transport = transport.clone(); - thread::spawn(move || { - const BUFFERED_CHUNKS: usize = 10; - let mut buffer = [0u8; BUFFERED_CHUNKS * FILE_CHUNK_SIZE]; - while let Ok(n) = file.read(&mut buffer) { - if n == 0 { - break; - } - let mut i = 0; - while i < n { - let j = n.min(i + FILE_CHUNK_SIZE); - send!( - peer, - &transport, - SendData::FileDownload, - Payload::FileDownload { fileid, file_chunk: &buffer[i..j] } - ); - i = j; - } - } - send!( - peer, - &transport, - SendData::FileDownload, - Payload::FileDownloadComplete { fileid } - ); - }); - } - (Payload::FileDownload { fileid, file_chunk }, None) => { - let mut file = peer.downloads_in_progress.get(&fileid)?; - let result = file.write_all(file_chunk); - } - (Payload::FileDownloadComplete { fileid }, None) => { - - } - (Payload::ReadDir, None) => { - let dir = read_dir(&peer.home_dir).ok()?; - let mut filenames = Vec::new(); - for entry in dir { - if let Ok(entry) = entry { - if entry.file_type().map_or(false, |f| f.is_file()) { - if let Ok(filename) = entry.file_name().into_string() { - filenames.push(filename) - } - } - } else { - return None; - } - } - let filenames: Vec<&str> = filenames.iter().map(|f| f.as_str()).collect(); - reply!(guard, SendData::DirContents, Payload::DirContents { filenames }); - } - (Payload::DirContents { filenames }, Some(SendData::ReadDir { download_missing: true })) => { - drop(guard); - for filename in filenames {} - } - _ => {} - } - Some(()) -} - -fn receive(peer: &Arc, transport: &Transport, receiver: &Receiver>) -> Option<()> { - while let Ok(packet) = receiver.try_recv() { - if drop_packet() { - continue; - } - let pt = *packet.get(0)?; - let (mut parsed_packet, payload_offset) = match pt { - PACKET_TYPE_ACK => { - let reply_no = SeqNo::from_be_bytes(packet.get(1..5)?.try_into().ok()?); - (Packet::Ack(reply_no), 5) - } - PACKET_TYPE_PAYLOAD | PACKET_TYPE_LOCK_PAYLOAD => { - let seq_no = SeqNo::from_be_bytes(packet.get(1..5)?.try_into().ok()?); - (Packet::Payload(seq_no, packet), 5) - } - PACKET_TYPE_REPLY | PACKET_TYPE_LOCK_REPLY => { - let seq_no = SeqNo::from_be_bytes(packet.get(1..5)?.try_into().ok()?); - let reply_no = SeqNo::from_be_bytes(packet.get(5..9)?.try_into().ok()?); - (Packet::Reply(seq_no, reply_no, packet), 9) - } - _ => return None, - }; - if pt == PACKET_TYPE_LOCK_PAYLOAD || pt == PACKET_TYPE_LOCK_REPLY { - parsed_packet.set_locking(true); - } - for recv_data in peer.seqex.receive_all(transport, parsed_packet) { - let recv_data = recv_data.map(|packet| serde_cbor::from_slice::(&packet[payload_offset..])); - if let Ok(parsed_payload) = serde_cbor::from_slice::(packet.get(payload_offset..)?) { - process(peer, transport, guard, parsed_packet, send_data.map(|d| d.0)); - } - } - } - Some(()) -} - -fn main() { - let alice_root = std::path::Path::new("examples").join("alice_home"); - let bob_root = std::path::Path::new("examples").join("bob_home"); - let (s, r) = channel::(); - thread::spawn(move || { - s.send(0); - }); -} - -#[test] -fn test() { - main() -} diff --git a/examples/calculator.rs b/examples/calculator.rs index 91e7377..bce0892 100644 --- a/examples/calculator.rs +++ b/examples/calculator.rs @@ -1,7 +1,7 @@ use std::{sync::mpsc::Receiver, thread, time::Duration}; use seq_ex::{ - sync::{MpscSeqEx, MpscTransport}, + sync::{MpscTransport, SeqExSync}, Packet, }; @@ -19,7 +19,7 @@ fn drop_packet() -> bool { rand_core::OsRng.next_u32() & 1 > 0 } -fn receive(recv: &Receiver>, seq: &MpscSeqEx, transport: &MpscTransport, value: &mut f32) { +fn receive(recv: &Receiver>, seq: &SeqExSync, transport: &MpscTransport, value: &mut f32) { while let Ok(packet) = recv.try_recv() { if drop_packet() { continue; @@ -42,20 +42,20 @@ fn receive(recv: &Receiver>, seq: &MpscSeqEx, transport fn main() { let (transport1, recv2) = MpscTransport::new(); let (transport2, recv1) = MpscTransport::new(); - let seq1 = MpscSeqEx::new(5, 1); - let seq2 = MpscSeqEx::new(5, 1); + let seq1 = SeqExSync::new(5, 1); + let seq2 = SeqExSync::new(5, 1); let mut value = 0.0; let mut remote_value = value; - seq1.send(&transport1, Payload::Add(1.0)); + seq1.send(&transport1, true, Payload::Add(1.0)); value += 1.0; - seq1.send(&transport1, Payload::Sub(2.0)); + seq1.send(&transport1, true, Payload::Sub(2.0)); value -= 2.0; - seq1.send(&transport1, Payload::Mul(3.0)); + seq1.send(&transport1, true, Payload::Mul(3.0)); value *= 3.0; - seq1.send(&transport1, Payload::Div(4.0)); + seq1.send(&transport1, true, Payload::Div(4.0)); value /= 4.0; - seq1.send(&transport1, Payload::Mod(5.0)); + seq1.send(&transport1, true, Payload::Mod(5.0)); value %= 5.0; for _ in 0..16 { diff --git a/examples/file_download.rs b/examples/file_download.rs index c39bb96..7a31388 100644 --- a/examples/file_download.rs +++ b/examples/file_download.rs @@ -60,7 +60,7 @@ fn process(peer: &Peer, recv_data: RecvOk<'_, &Transport, Payload, Payload, Payl let transport = peer.transport.clone(); let seqex = peer.seqex.clone(); if let Some(file) = filesystem.read().unwrap().get(&filename) { - guard.reply(ConfirmRequestFile { filesize: file.len() as u64 }); + guard.reply(true, ConfirmRequestFile { filesize: file.len() as u64 }); } thread::spawn(move || { let filesystem = filesystem.read().unwrap(); @@ -70,18 +70,22 @@ fn process(peer: &Peer, recv_data: RecvOk<'_, &Transport, Payload, Payload, Payl let mut i = 0; while i < file.len() { let j = file.len().min(i + FILE_CHUNK_SIZE); - seqex.send_locked(&transport, FileDownload { filename: filename.clone(), file_chunk: file[i..j].to_vec() }); + seqex.send( + &transport, + true, + FileDownload { filename: filename.clone(), file_chunk: file[i..j].to_vec() }, + ); i = j; } } }); } - (Some((_, ConfirmRequestFile { filesize })), Some(RequestFile { filename })) => { + (Some((_g, ConfirmRequestFile { filesize })), Some(RequestFile { filename })) => { let mut filesystem = peer.filesystem.write().unwrap(); let file = Vec::with_capacity(filesize as usize); filesystem.insert(filename, file); } - (Some((_, FileDownload { filename, file_chunk })), None) => { + (Some((_g, 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() { @@ -102,7 +106,7 @@ fn receive(peer: &Peer) { continue; } if let Ok(parsed_packet) = serde_cbor::from_slice::>(&packet) { - for recv_data in peer.seqex.receive_all(&peer.transport, parsed_packet) { + for recv_data in peer.seqex.try_receive_all(&peer.transport, parsed_packet) { process(peer, recv_data); } } @@ -137,14 +141,15 @@ fn main() { receiver: recv2, }; - 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() }); + let tl = &peer1.transport; + peer1.seqex.send(tl, false, Payload::RequestFile { filename: "File1".to_string() }); + peer1.seqex.send(tl, false, Payload::RequestFile { filename: "File3".to_string() }); + peer1.seqex.send(tl, false, Payload::RequestFile { filename: "File2".to_string() }); - for _ in 0..500 { + for _ in 0..400 { receive(&peer1); receive(&peer2); - thread::sleep(Duration::from_millis(1)); + thread::sleep(Duration::from_millis(2)); peer1.seqex.service(&peer1.transport); peer2.seqex.service(&peer2.transport); } diff --git a/examples/hello_world.rs b/examples/hello_world.rs index 3e284ca..5b249e4 100644 --- a/examples/hello_world.rs +++ b/examples/hello_world.rs @@ -1,7 +1,7 @@ use std::sync::mpsc::Receiver; use seq_ex::{ - sync::{MpscSeqEx, MpscTransport, RecvOk}, + sync::{MpscTransport, RecvOk, SeqExSync}, Packet, }; @@ -14,23 +14,23 @@ enum Payload { } use Payload::*; -fn receive(recv: &Receiver>, seq: &MpscSeqEx, transport: &MpscTransport) { +fn receive(recv: &Receiver>, seq: &SeqExSync, 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); + guard.reply(false, Space); } (Some((guard, Space)), Some(Hello)) => { print!(" "); - guard.reply(World); + guard.reply(false, World); } (Some((guard, World)), Some(Space)) => { print!("World"); - guard.reply(Exclamation); + guard.reply(false, Exclamation); } - (Some((_, Exclamation)), Some(World)) => { + (Some((_g, Exclamation)), Some(World)) => { print!("!"); } (None, Some(Exclamation)) => { @@ -48,11 +48,11 @@ fn receive(recv: &Receiver>, seq: &MpscSeqEx, transport fn main() { let (transport1, recv2) = MpscTransport::new(); let (transport2, recv1) = MpscTransport::new(); - let seq1 = MpscSeqEx::default(); - let seq2 = MpscSeqEx::default(); + let seq1 = SeqExSync::default(); + let seq2 = SeqExSync::default(); // We begin a "Hello World" exchange right here. - seq1.send(&transport1, Payload::Hello); + seq1.send(&transport1, false, Payload::Hello); receive(&recv2, &seq2, &transport2); receive(&recv1, &seq1, &transport1); diff --git a/examples/hello_world_tokio.rs b/examples/hello_world_tokio.rs new file mode 100644 index 0000000..e720bc0 --- /dev/null +++ b/examples/hello_world_tokio.rs @@ -0,0 +1,118 @@ +use std::sync::Arc; + +use tokio::{sync::mpsc, task}; + +use seq_ex::{ + tokio::{AsyncRecvError, MpscTransport, ReplyGuard, SeqExTokio}, + Packet, +}; + +#[derive(Clone, Debug, PartialEq, Eq)] +enum Payload { + Hello, + Space, + World, + Exclamation, + NewLine, +} +use Payload::*; + +async fn receive(reply_guard: ReplyGuard<'_, &MpscTransport, Payload>, payload: Payload) -> Option<()> { + match payload { + Hello => { + print!("Hello"); + + let (reply_guard, payload) = reply_guard.reply(false, Space).await.ok()?; + if payload != World { + return None; + } + print!("World"); + + let (reply_guard, payload) = reply_guard.reply(false, Exclamation).await.ok()?; + if payload != NewLine { + return None; + } + print!("\n"); + drop(reply_guard); + } + _ => { + assert!(false, "Unsolicited payload received: {:?}", payload); + } + } + Some(()) +} + +async fn say_hello(seq: &SeqExTokio, transport: &MpscTransport) -> Option<()> { + let (reply_guard, payload) = seq.send(transport, false, Hello).await.ok()?; + if payload != Space { + return None; + } + print!(" "); + + let (reply_guard, payload) = reply_guard.reply(false, World).await.ok()?; + if payload != Exclamation { + return None; + } + print!("!"); + + let _ = reply_guard.reply(false, NewLine).await; + Some(()) +} + +async fn pump_all(peer: Arc>, transport: MpscTransport) { + if let Some((g, payload)) = peer.pump(&transport).await { + spawn_pump(peer.clone(), transport.clone()); + receive(g, payload).await; + } +} +fn spawn_pump(peer: Arc>, transport: MpscTransport) { + task::spawn(async move { + pump_all(peer, transport).await; + }); +} + +fn peer_main(transport: MpscTransport, mut recv: mpsc::Receiver>) -> Arc> { + let (seq, mut service) = SeqExTokio::::new_default(); + let peer = Arc::new(seq); + let peer_weak = Arc::downgrade(&peer); + let tl = transport.clone(); + task::spawn(async move { + while let Some(peer) = peer_weak.upgrade() { + peer.service_task(&tl, &mut service).await; + } + }); + let peer_weak = Arc::downgrade(&peer); + task::spawn(async move { + while let Some(packet) = recv.recv().await { + if let Some(peer) = peer_weak.upgrade() { + let tl = transport.clone(); + task::spawn(async move { + let result = peer.receive(&tl, packet); + if matches!(&result, Ok(_) | Err(AsyncRecvError::AsyncReply)) { + spawn_pump(peer.clone(), tl.clone()); + } + if let Ok((g, payload)) = result { + receive(g, payload).await; + } + }); + } else { + break; + } + } + }); + peer +} + +#[tokio::main] +async fn main() { + let (transport2, recv1) = MpscTransport::new(32); + let (transport1, recv2) = MpscTransport::new(32); + let peer1 = peer_main(transport1.clone(), recv1); + let _peer2 = peer_main(transport2, recv2); + + say_hello(&peer1, &transport1).await; +} +#[test] +fn test() { + main() +} diff --git a/src/seq_queue.rs b/src/seq_queue.rs index a876696..1d4bdf2 100644 --- a/src/seq_queue.rs +++ b/src/seq_queue.rs @@ -42,10 +42,11 @@ pub type SeqNo = u32; /// The resend interval for a default instance of SeqEx. pub const DEFAULT_RESEND_INTERVAL_MS: i64 = 200; /// The initial sequence number for a default instance of SeqEx. -pub const DEFAULT_INITIAL_SEQ_NO: SeqNo = 1; +pub const DEFAULT_INITIAL_SEQ_NO: SeqNo = 0; pub const DEFAULT_WINDOW_CAP: usize = 64; +#[derive(Debug)] pub struct SeqEx { /// The interval at which packets will be resent if they have not yet been acknowledged by the /// remote peer. @@ -56,7 +57,7 @@ pub struct SeqEx { pre_recv_seq_no: SeqNo, /// This could be made more efficient by changing to SoA format. send_window: [Option>; CAP], - recv_window: [Option>; CAP], + recv_window: [RecvEntry; CAP], /// The size of this array determines the maximum number of received packets that the application /// may attempt to process concurrently before new received packets start being dropped. concurrent_replies: [SeqNo; CAP], @@ -68,63 +69,89 @@ pub struct SeqEx { is_locked: bool, } -struct RecvEntry { - seq_no: SeqNo, - reply_no: Option, - locked: bool, - data: RecvData, +#[derive(Debug)] +enum RecvEntry { + Occupied { + seq_no: SeqNo, + reply_no: Option, + seq_cst: bool, + data: RecvData, + }, + Unlocked { + seq_no: SeqNo, + }, + Empty, } +#[derive(Debug)] struct SendEntry { seq_no: SeqNo, reply_no: Option, + seq_cst: bool, next_resend_time: i64, data: SendData, } #[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), +pub enum DirectRecvError { + DroppedTooEarly, + DroppedDuplicate, + DroppedDuplicateResendAck(SeqNo), + WaitingForRecv, + WaitingForReply, } +#[cfg(feature = "std")] +impl std::fmt::Display for DirectRecvError { + 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"), + } + } +} +#[cfg(feature = "std")] +impl std::error::Error for DirectRecvError {} #[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, + WaitingForRecv, + WaitingForReply, } +#[cfg(feature = "std")] +impl std::fmt::Display for PumpError { + 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"), + } + } +} +#[cfg(feature = "std")] +impl std::error::Error for PumpError {} #[derive(Clone, Copy, Debug, PartialEq, Eq)] #[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] pub enum Packet { Payload(SeqNo, RecvData), - LockPayload(SeqNo, RecvData), + SeqCstPayload(SeqNo, RecvData), Reply(SeqNo, SeqNo, RecvData), - LockReply(SeqNo, SeqNo, RecvData), + SeqCstReply(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_with_data(seq_no: SeqNo, reply_no: Option, seq_cst: bool, data: RecvData) -> Self { + Self::new(Some(seq_no), reply_no, seq_cst, 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) { + pub fn new(seq_no: Option, reply_no: Option, seq_cst: bool, data: Option) -> Option { + match (seq_no, reply_no, seq_cst, data) { (Some(s), None, false, Some(d)) => Some(Payload(s, d)), - (Some(s), None, true, Some(d)) => Some(LockPayload(s, d)), + (Some(s), None, true, Some(d)) => Some(SeqCstPayload(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)), + (Some(s), Some(r), true, Some(d)) => Some(SeqCstReply(s, r, d)), (None, Some(r), false, None) => Some(Ack(r)), _ => None, } @@ -132,18 +159,18 @@ impl Packet { 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), + SeqCstPayload(seq_no, data) => SeqCstPayload(*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), + SeqCstReply(seq_no, reply_no, data) => SeqCstReply(*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)), + SeqCstPayload(seq_no, data) => SeqCstPayload(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)), + SeqCstReply(seq_no, reply_no, data) => SeqCstReply(seq_no, reply_no, f(data)), Ack(reply_no) => Ack(reply_no), } } @@ -152,27 +179,27 @@ impl Packet { } pub fn consume(self) -> Result { match self { - Payload(_, data) | LockPayload(_, data) | Reply(_, _, data) | LockReply(_, _, data) => Ok(data), + Payload(_, data) | SeqCstPayload(_, data) | Reply(_, _, data) | SeqCstReply(_, _, data) => Ok(data), Ack(r) => Err(r), } } - pub fn is_locking(&self) -> bool { - matches!(self, LockPayload(..) | LockReply(..)) + pub fn is_seq_cst(&self) -> bool { + matches!(self, SeqCstPayload(..) | SeqCstReply(..)) } - pub fn set_locking(&mut self, locking: bool) { + pub fn set_seq_cst(&mut self, seq_cst: 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) + Payload(seq_no, data) | SeqCstPayload(seq_no, data) => { + *self = if seq_cst { + SeqCstPayload(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) + Reply(seq_no, reply_no, data) | SeqCstReply(seq_no, reply_no, data) => { + *self = if seq_cst { + SeqCstReply(seq_no, reply_no, data) } else { Reply(seq_no, reply_no, data) } @@ -187,15 +214,16 @@ impl Packet<&RecvData> { } } +#[derive(Clone, Debug)] pub enum RecvOkRaw { Payload { reply_no: SeqNo, - locked: bool, + seq_cst: bool, recv_data: RecvData, }, Reply { reply_no: SeqNo, - locked: bool, + seq_cst: bool, recv_data: RecvData, send_data: SendData, }, @@ -241,31 +269,30 @@ impl SeqEx { next_service_timestamp: i64::MAX, next_send_seq_no: initial_seq_no, pre_recv_seq_no: initial_seq_no.wrapping_sub(1), - recv_window: core::array::from_fn(|_| None), + recv_window: core::array::from_fn(|_| RecvEntry::Empty), send_window: core::array::from_fn(|_| None), concurrent_replies: core::array::from_fn(|_| 0), concurrent_replies_total: 0, is_locked: false, } } + #[inline] fn send_window_slot_mut(&mut self, seq_no: SeqNo) -> &mut Option> { &mut self.send_window[seq_no as usize % self.send_window.len()] } - 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() { + #[inline] + fn is_full_inner(&self, is_for_reply: bool, reply_no: Option) -> bool { + let reply_idx = reply_no.map_or(self.send_window.len(), |r| r as usize % self.send_window.len()); + for i in 0..self.concurrent_replies_total as u32 + 1 + is_for_reply 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 true; } } false } + + #[inline] 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) { @@ -274,7 +301,6 @@ impl SeqEx { 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 { @@ -292,7 +318,7 @@ impl SeqEx { pub fn is_full(&self) -> bool { // We claim that the window is full one entry before it is actually full for the sake of // making it always possible for both peers to process at least one reply at all times. - self.is_full_inner(true) + self.is_full_inner(true, None) } /// Returns the next sequence number to be attached to the next sent packet. /// This should be called before `SeqEx::send`, and the return value should be @@ -320,9 +346,23 @@ 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. - fn try_send_direct_inner(&mut self, current_time: i64) -> Option<(&mut Option>, SeqNo, i64)> { + /// + /// Can mutate `next_service_timestamp`. + pub fn try_send_direct(&mut self, current_time: i64, seq_cst: bool, packet_data: SendData) -> Result, SendData> { + let mut tmp = Some(packet_data); + self.try_send_direct_with(current_time, seq_cst, |_| tmp.take().unwrap()) + .map_err(|_| ()) + .map_err(|_| tmp.unwrap()) + } + /// Can mutate `next_service_timestamp`. + pub fn try_send_direct_with SendData>( + &mut self, + current_time: i64, + seq_cst: bool, + packet_data: F, + ) -> Result, F> { if self.is_full() { - return None; + return Err(packet_data); } let seq_no = self.next_send_seq_no; self.next_send_seq_no = self.next_send_seq_no.wrapping_add(1); @@ -333,41 +373,30 @@ impl SeqEx { } let slot = self.send_window_slot_mut(seq_no); debug_assert!(slot.is_none()); - 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(()) - } + let entry = slot.insert(SendEntry { + seq_no, + reply_no: None, + seq_cst, + next_resend_time, + data: packet_data(seq_no), + }); + + let mut p = Payload(entry.seq_no, &entry.data); + p.set_seq_cst(seq_cst); + Ok(p) } /// If this returns `Ok` then `try_send` might succeed on next call. - pub fn receive_raw_and_direct>(&mut self, packet: Packet

) -> Result, DirectError

> { - let locked = packet.is_locking(); + pub fn receive_raw_and_direct>(&mut self, packet: Packet

) -> Result, DirectRecvError> { + let seq_cst = packet.is_seq_cst(); 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) => { + Payload(seq_no, recv_data) | SeqCstPayload(seq_no, recv_data) => (seq_no, None, recv_data), + Reply(seq_no, reply_no, recv_data) | SeqCstReply(seq_no, reply_no, recv_data) => (seq_no, Some(reply_no), recv_data), + Ack(reply_no) => { return self .take_send(reply_no) .map(|send_data| RecvOkRaw::Ack { send_data }) - .ok_or(DirectError::OutOfSequence) + .ok_or(DirectRecvError::DroppedDuplicate) } }; // We only want to accept packets with sequence numbers in the range: @@ -390,104 +419,110 @@ impl SeqEx { // resending the packet. for entry in self.send_window.iter().flatten() { if entry.reply_no == Some(seq_no) { - return Err(DirectError::OutOfSequence); + return Err(DirectRecvError::DroppedDuplicate); } } for i in 0..self.concurrent_replies_total { if self.concurrent_replies[i] == seq_no { - return Err(DirectError::OutOfSequence); + return Err(DirectRecvError::DroppedDuplicate); } } - return Err(DirectError::ResendAck(seq_no)); + return Err(DirectRecvError::DroppedDuplicateResendAck(seq_no)); } else if is_above_range { - return Err(DirectError::OutOfSequence); + return Err(DirectRecvError::DroppedTooEarly); } - // 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(); - let is_in_window = if let Some(pre) = self.recv_window[i].as_mut() { - if seq_no == pre.seq_no { - 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(DirectError::OutOfSequence); - } + let is_duplicate = if let RecvEntry::Occupied { seq_no: pre_seq_no, .. } | RecvEntry::Unlocked { seq_no: pre_seq_no } = &self.recv_window[i] { + // Due to the range check these should always be equal. + // `self.pre_recv_seq_no < seq_no <= self.pre_recv_seq_no + self.recv_window.len()`. + debug_assert_eq!(seq_no, *pre_seq_no); + true } else { false }; - - if !is_next { - if !is_in_window { - self.recv_window[i] = Some(RecvEntry { seq_no, reply_no, locked, data: recv_data.into() }) + // 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_for_recv = self.is_full_inner(false, reply_no) || (seq_cst && !is_next); + let wait_for_reply = self.concurrent_replies_total >= self.concurrent_replies.len() || (seq_cst && self.is_locked); + if wait_for_recv || wait_for_reply { + if !is_duplicate { + self.recv_window[i] = RecvEntry::Occupied { seq_no, reply_no, seq_cst, 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; - } - self.pre_recv_seq_no = seq_no; - self.concurrent_replies[self.concurrent_replies_total] = seq_no; - self.concurrent_replies_total += 1; - 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 } + return if wait_for_recv { + Err(DirectRecvError::WaitingForRecv) } else { - RecvOkRaw::Payload { reply_no: seq_no, recv_data, locked } - }) + Err(DirectRecvError::WaitingForReply) + }; } + + if is_next { + self.recv_window[i] = RecvEntry::Empty; + self.pre_recv_seq_no = seq_no; + } else { + debug_assert!(!seq_cst); + self.recv_window[i] = RecvEntry::Unlocked { seq_no } + } + self.concurrent_replies[self.concurrent_replies_total] = seq_no; + self.concurrent_replies_total += 1; + if seq_cst { + 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, seq_cst } + } else { + RecvOkRaw::Payload { reply_no: seq_no, recv_data, seq_cst } + }) } /// If this returns `Ok` then `try_send` might succeed on next call. - 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(); + pub fn try_pump_raw(&mut self) -> Result, PumpError> { + let mut next_seq_no; + let mut i; + loop { + next_seq_no = self.pre_recv_seq_no.wrapping_add(1); + i = next_seq_no as usize % self.recv_window.len(); + if let RecvEntry::Unlocked { seq_no } = &self.recv_window[i] { + debug_assert_eq!(*seq_no, next_seq_no); + self.recv_window[i] = RecvEntry::Empty; + self.pre_recv_seq_no = next_seq_no; + } else { + break; + } + } - 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 let RecvEntry::Occupied { seq_no, reply_no, seq_cst, .. } = &self.recv_window[i] { + debug_assert_eq!(*seq_no, next_seq_no); + if self.is_full_inner(false, *reply_no) { + return Err(PumpError::WaitingForRecv); + } - if !entry.locked || !self.is_locked { - let entry = self.recv_window[i].take().unwrap(); + if (!*seq_cst || !self.is_locked) && self.concurrent_replies_total < self.concurrent_replies.len() { + let mut entry = RecvEntry::Empty; + core::mem::swap(&mut entry, &mut self.recv_window[i]); + if let RecvEntry::Occupied { seq_no, reply_no, seq_cst, data } = entry { self.pre_recv_seq_no = next_seq_no; - self.concurrent_replies[self.concurrent_replies_total] = entry.seq_no; + self.concurrent_replies[self.concurrent_replies_total] = seq_no; self.concurrent_replies_total += 1; - if entry.locked { + if seq_cst { 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, - } + return Ok(if let Some(send_data) = reply_no.and_then(|r| self.take_send(r)) { + RecvOkRaw::Reply { reply_no: seq_no, seq_cst, recv_data: data, send_data } } else { - RecvOkRaw::Payload { - reply_no: entry.seq_no, - locked: entry.locked, - recv_data: entry.data, - } + RecvOkRaw::Payload { reply_no: seq_no, seq_cst, recv_data: data } }); } else { - return Err(PumpError::WindowIsLocked); + unreachable!(); } + } else { + return Err(PumpError::WaitingForReply); } } - Err(PumpError::OutOfSequence) + Err(PumpError::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. @@ -497,8 +532,18 @@ 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. + /// + /// Can mutate `next_service_timestamp`. + /// If `unlock` is true and the return value is `Some` pump may return new values. #[must_use] - pub fn reply_raw_and_direct(&mut self, reply_no: SeqNo, unlock: bool, packet_data: SendData, current_time: i64) -> Option> { + pub fn reply_raw_and_direct( + &mut self, + current_time: i64, + reply_no: SeqNo, + unlock: bool, + seq_cst: bool, + packet_data: SendData, + ) -> 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); @@ -516,15 +561,19 @@ impl SeqEx { let entry = slot.insert(SendEntry { seq_no, reply_no: Some(reply_no), + seq_cst, next_resend_time, data: packet_data, }); - Some(Packet::Reply(entry.seq_no, reply_no, &entry.data)) + let mut p = Reply(entry.seq_no, reply_no, &entry.data); + p.set_seq_cst(seq_cst); + Some(p) } else { None } } + /// If `unlock` is true and the return value is `Some` pump may return new values. #[must_use] pub fn ack_raw_and_direct(&mut self, reply_no: SeqNo, unlock: bool) -> Option> { if self.remove_reservation(reply_no) { @@ -532,13 +581,14 @@ impl SeqEx { 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)) + Some(Ack(reply_no)) } else { None } } - pub fn service_direct<'a>(&'a mut self, current_time: i64, iter: &mut Option) -> Option> { + /// Can mutate `next_service_timestamp`. + pub fn service_direct(&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) { @@ -548,11 +598,14 @@ 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(if let Some(reply_no) = entry.reply_no { - Packet::Reply(entry.seq_no, reply_no, &entry.data) + + let mut p = if let Some(reply_no) = entry.reply_no { + Reply(entry.seq_no, reply_no, &entry.data) } else { - Packet::Payload(entry.seq_no, &entry.data) - }); + Payload(entry.seq_no, &entry.data) + }; + p.set_seq_cst(entry.seq_cst); + return Some(p); } 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 34d30e9..402a016 100644 --- a/src/single_thread.rs +++ b/src/single_thread.rs @@ -1,10 +1,10 @@ -use crate::{DirectError, Packet, PumpError, RecvOkRaw, SeqEx, SeqNo, TransportLayer, DEFAULT_WINDOW_CAP}; +use crate::{DirectRecvError, Packet, PumpError, 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, app: Option, reply_no: SeqNo, - locked: bool, + is_holding_lock: 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 @@ -12,31 +12,25 @@ impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> Rep /// 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(self, packet_data: SendData) { - self.reply_inner(false, |_, _| packet_data) + pub fn reply(self, seq_cst: bool, packet_data: SendData) { + self.reply_with(seq_cst, |_, _| 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); + fn reply_with(mut self, seq_cst: bool, packet_data: impl FnOnce(SeqNo, SeqNo) -> SendData) { + let app = self.app.take(); 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); + self.seq.reply_raw( + app.unwrap(), + self.reply_no, + self.is_holding_lock, + seq_cst, + packet_data(seq_no, self.reply_no), + ); } } 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.app { - if let Some(p) = self.seq.ack_raw_and_direct(self.reply_no, self.locked) { + if let Some(p) = self.seq.ack_raw_and_direct(self.reply_no, self.is_holding_lock) { app.send(p) } } @@ -44,20 +38,33 @@ impl<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize> Dro } 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() + f.debug_struct("ReplyGuard") + .field("reply_no", &self.reply_no) + .field("is_holding_lock", &self.is_holding_lock) + .finish() } } #[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 RecvError { + DroppedTooEarly, + DroppedDuplicate, + WaitingForRecv, + WaitingForReply, } +#[cfg(feature = "std")] +impl std::fmt::Display for RecvError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + RecvError::DroppedTooEarly => write!(f, "packet arrived too early"), + RecvError::DroppedDuplicate => write!(f, "packet was a duplicate"), + RecvError::WaitingForRecv => write!(f, "can't process until another packet is received"), + RecvError::WaitingForReply => write!(f, "can't process until a reply is finished"), + } + } +} +#[cfg(feature = "std")] +impl std::error::Error for RecvError {} pub enum RecvOk<'a, TL: TransportLayer, P, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { Payload { @@ -97,14 +104,14 @@ macro_rules! impl_recvok { } } 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 { + 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 }, + RecvOkRaw::Payload { reply_no, seq_cst, recv_data } => Self::Payload { + reply_guard: ReplyGuard { seq, app: Some(app), reply_no, is_holding_lock: seq_cst }, recv_data, }, - RecvOkRaw::Reply { reply_no, locked, recv_data, send_data } => Self::Reply { - reply_guard: ReplyGuard { seq, app: Some(app), reply_no, locked }, + RecvOkRaw::Reply { reply_no, seq_cst, recv_data, send_data } => Self::Reply { + reply_guard: ReplyGuard { seq, app: Some(app), reply_no, is_holding_lock: seq_cst }, recv_data, send_data, }, @@ -126,13 +133,6 @@ macro_rules! impl_recvok { (None, None) => None, } } - pub fn map(self, f: impl FnOnce(P) -> R) -> $recv<'a, TL, R, SendData, RecvData, CAP> { - match self { - Self::Payload { reply_guard, recv_data } => $recv::Payload { reply_guard, recv_data: f(recv_data) }, - Self::Reply { reply_guard, recv_data, send_data } => $recv::Reply { reply_guard, recv_data: f(recv_data), send_data }, - Self::Ack { send_data } => $recv::Ack { send_data }, - } - } } 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> { @@ -149,25 +149,25 @@ 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); + /// Can mutate `next_service_timestamp`. + pub fn try_send(&mut self, mut app: impl TransportLayer, seq_cst: bool, packet_data: SendData) -> Result<(), SendData> { + match self.try_send_direct(app.time(), seq_cst, packet_data) { + Ok(p) => { app.send(p); Ok(()) } Err(e) => Err(e), } } - pub fn try_send_with( + /// Can mutate `next_service_timestamp`. + pub fn try_send_with SendData>( &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); + seq_cst: bool, + packet_data: F, + ) -> Result<(), F> { + match self.try_send_direct_with(app.time(), seq_cst, packet_data) { + Ok(p) => { app.send(p); Ok(()) } @@ -179,29 +179,41 @@ impl SeqEx { &mut self, mut app: impl TransportLayer, packet: Packet

, - ) -> Result, Error

> { + ) -> Result, RecvError> { match self.receive_raw_and_direct(packet) { Ok(a) => Ok(a), - Err(DirectError::ResendAck(reply_no)) => { + Err(DirectRecvError::DroppedDuplicateResendAck(reply_no)) => { app.send(Packet::Ack(reply_no)); - Err(Error::OutOfSequence) + Err(RecvError::DroppedDuplicate) } - Err(DirectError::OutOfSequence) => Err(Error::OutOfSequence), - Err(DirectError::WindowIsFull(p)) => Err(Error::WindowIsFull(p)), - Err(DirectError::WindowIsLocked(p)) => Err(Error::WindowIsLocked(p)), + Err(DirectRecvError::DroppedTooEarly) => Err(RecvError::DroppedTooEarly), + Err(DirectRecvError::DroppedDuplicate) => Err(RecvError::DroppedDuplicate), + Err(DirectRecvError::WaitingForRecv) => Err(RecvError::WaitingForRecv), + Err(DirectRecvError::WaitingForReply) => Err(RecvError::WaitingForReply), } } - 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) + /// Can mutate `next_service_timestamp`. + /// If `unlock` is true and the return value is true pump may return new values. + /// + /// Only returns false if the reply number was incorrect or used twice. + pub fn reply_raw(&mut self, mut app: impl TransportLayer, reply_no: SeqNo, unlock: bool, seq_cst: bool, packet_data: SendData) -> bool { + if let Some(p) = self.reply_raw_and_direct(app.time(), reply_no, unlock, seq_cst, packet_data) { + app.send(p); + true + } else { + false } } - pub fn ack_raw(&mut self, mut app: impl TransportLayer, reply_no: SeqNo, unlock: bool) { + /// If `unlock` is true and the return value is true pump may return new values. + pub fn ack_raw(&mut self, mut app: impl TransportLayer, reply_no: SeqNo, unlock: bool) -> bool { if let Some(p) = self.ack_raw_and_direct(reply_no, unlock) { - app.send(p) + app.send(p); + true + } else { + false } } + /// Can mutate `next_service_timestamp`. pub fn service(&mut self, mut app: impl TransportLayer) -> i64 { let current_time = app.time(); let mut iter = None; @@ -214,10 +226,10 @@ impl SeqEx { &mut self, app: TL, packet: Packet

, - ) -> Result, Error

> { + ) -> Result, RecvError> { self.receive_raw(app.clone(), packet).map(|r| RecvOk::from_raw(self, app, r)) } - pub fn pump>(&mut self, app: TL) -> Result, PumpError> { - self.pump_raw().map(|r| RecvOk::from_raw(self, app, r)) + pub fn try_pump>(&mut self, app: TL) -> Result, PumpError> { + self.try_pump_raw().map(|r| RecvOk::from_raw(self, app, r)) } } diff --git a/src/sync.rs b/src/sync.rs index ecdf918..7e61539 100644 --- a/src/sync.rs +++ b/src/sync.rs @@ -1,83 +1,71 @@ use std::{ - ops::{Deref, DerefMut}, sync::{ mpsc::{channel, Receiver, Sender}, - Condvar, Mutex, MutexGuard, + Condvar, Mutex, }, time::Instant, }; -use crate::{Packet, PumpError, RecvOkRaw, SeqEx, SeqNo, TransportLayer, DEFAULT_INITIAL_SEQ_NO, DEFAULT_RESEND_INTERVAL_MS, DEFAULT_WINDOW_CAP}; +use crate::{ + Packet, PumpError, RecvError, RecvOkRaw, SeqEx, SeqNo, TransportLayer, DEFAULT_INITIAL_SEQ_NO, DEFAULT_RESEND_INTERVAL_MS, DEFAULT_WINDOW_CAP, +}; pub struct SeqExSync { - seq_ex: Mutex<(SeqEx, usize)>, + seq_ex: Mutex<(SeqEx, usize, bool)>, send_block: Condvar, - recv_lock: Condvar, + reply_block: Condvar, } pub struct ReplyGuard<'a, TL: TransportLayer, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { seq: &'a SeqExSync, app: Option, reply_no: SeqNo, - locked: bool, + is_holding_lock: bool, } 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 + .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(); + } + } /// 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. /// 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(self, packet_data: SendData) { - 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 mut seq = self.seq.lock(); - let seq_no = seq.seq_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); + pub fn reply(self, seq_cst: bool, packet_data: SendData) { + self.reply_with(seq_cst, |_, _| packet_data) } } 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.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) + 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(); } } } } 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() + f.debug_struct("ReplyGuard") + .field("reply_no", &self.reply_no) + .field("is_holding_lock", &self.is_holding_lock) + .finish() } } -#[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 enum RecvOk<'a, TL: TransportLayer, P, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { Payload { reply_guard: ReplyGuard<'a, TL, SendData, RecvData, CAP>, @@ -94,135 +82,138 @@ pub enum RecvOk<'a, TL: TransportLayer, P, SendData, RecvData, const C } crate::impl_recvok!(RecvOk, &'a SeqExSync); -pub struct ReplyIter<'a, TL: TransportLayer, P: Into, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { +pub struct RecvIter<'a, TL: TransportLayer, P, SendData, RecvData, const CAP: usize = DEFAULT_WINDOW_CAP> { seq: Option<&'a SeqExSync>, app: TL, first: Option>, -} - -pub struct SeqExGuard<'a, SendData, RecvData, const CAP: usize>(MutexGuard<'a, (SeqEx, usize)>); -impl<'a, SendData, RecvData, const CAP: usize> Deref for SeqExGuard<'a, SendData, RecvData, CAP> { - type Target = SeqEx; - - fn deref(&self) -> &Self::Target { - &self.0 .0 - } -} -impl<'a, SendData, RecvData, const CAP: usize> DerefMut for SeqExGuard<'a, SendData, RecvData, CAP> { - fn deref_mut(&mut self) -> &mut Self::Target { - &mut self.0 .0 - } + 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)), + seq_ex: Mutex::new((SeqEx::new(retry_interval, initial_seq_no), 0, false)), send_block: Condvar::default(), - recv_lock: Condvar::default(), + reply_block: Condvar::default(), } } pub fn receive, P: Into>( &self, app: TL, - mut packet: Packet

, - ) -> Result, Error> { + 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)) + } + 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); + } + 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)) + } + 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 { + return None; + } loop { - match seq.0.receive_raw(app.clone(), packet) { + match seq.0.try_pump_raw() { Ok(r) => { if seq.1 > 0 { + seq.1 -= 1; + drop(seq); self.send_block.notify_one(); } - return Ok(RecvOk::from_raw(self, app, r)); + return Some(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; + Err(PumpError::WaitingForRecv) => return None, + Err(PumpError::WaitingForReply) => { + seq.2 = true; + seq = self.reply_block.wait(seq).unwrap(); } } } } - pub fn pump>(&self, app: TL) -> Result, Error> { - let mut seq = self.seq_ex.lock().unwrap(); - 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(); - } - } + + fn receive_all_inner, P: Into>( + &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 }, } } pub fn receive_all, P: Into>( &self, app: TL, 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 { seq: None, app, first: None } - } + ) -> RecvIter<'_, TL, P, SendData, RecvData, CAP> { + self.receive_all_inner(app, true, packet) } - pub fn try_send>(&self, app: TL, locked: bool, packet_data: SendData) -> Result<(), SendData> { - self.try_send_with(app, locked, |_| packet_data) + 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) } - fn send_with_inner>(&self, app: TL, locked: bool, mut packet_data: impl FnMut(SeqNo) -> SendData) { + + 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(); - 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; - } + seq.0.try_send_with(app, seq_cst, packet_data) } - fn send_inner>(&self, app: TL, locked: bool, mut packet_data: SendData) { + pub fn try_send>(&self, app: TL, seq_cst: bool, packet_data: SendData) -> Result<(), SendData> { let mut seq = self.seq_ex.lock().unwrap(); - while let Err(p) = seq.0.try_send(app.clone(), locked, packet_data) { + seq.0.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) { packet_data = p; seq.1 += 1; seq = self.send_block.wait(seq).unwrap(); - seq.1 -= 1; } } - pub fn send(&self, app: impl TransportLayer, packet_data: SendData) { - self.send_inner(app, false, packet_data) - } - 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, locked, packet_data(seq_no)) - } - pub fn service>(&self, app: TL) -> i64 { - self.lock().service(app) + 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) { + packet_data = p; + seq.1 += 1; + seq = self.send_block.wait(seq).unwrap(); + } } - pub fn lock(&self) -> SeqExGuard<'_, SendData, RecvData, CAP> { - SeqExGuard(self.seq_ex.lock().unwrap()) + pub fn service>(&self, app: TL) -> i64 { + self.seq_ex.lock().unwrap().0.service(app) } } impl Default for SeqExSync { @@ -231,20 +222,24 @@ impl Default for SeqExSync, P: Into, SendData, RecvData, const CAP: usize> ReplyIter<'a, TL, P, SendData, RecvData, CAP> { +impl<'a, TL: TransportLayer, 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 ReplyIter<'a, TL, P, SendData, RecvData, CAP> + for RecvIter<'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.into()) } else if let Some(origin) = self.seq { - origin.pump(self.app.clone()).ok() + if self.blocking { + origin.pump(self.app.clone()) + } else { + origin.try_pump(self.app.clone()).ok() + } } else { None } @@ -256,8 +251,6 @@ pub struct MpscTransport { 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>) { diff --git a/src/tokio.rs b/src/tokio.rs index c2ec721..f9eac5d 100644 --- a/src/tokio.rs +++ b/src/tokio.rs @@ -1,316 +1,418 @@ -use std::sync::{Mutex, MutexGuard}; +use std::sync::Mutex; use tokio::{ sync::{mpsc, oneshot, Notify}, - task, time, + time, }; -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}; -type SendData = (oneshot::Sender<(Packet, SeqNo)>, Packet); +type SendData = (oneshot::Sender>, Payload); -pub struct SeqExTokio { - seq_ex: Mutex<(SeqEx, Packet, CAP>, usize)>, +pub struct SeqExTokio { + seq_ex: Mutex<(SeqEx, Payload, CAP>, usize, bool)>, send_block: Notify, + reply_block: Notify, + update_queue: mpsc::Sender, } -pub struct ReplyGuard<'a, TL: TokioTransportLayer, Packet, const CAP: usize = DEFAULT_WINDOW_CAP> { - origin: &'a SeqExTokio, - app: Option<&'a TokioTransport>, +pub struct ReplyGuard<'a, TL: TokioLayer, Payload, const CAP: usize = DEFAULT_WINDOW_CAP> { + seq: &'a SeqExTokio, + app: Option, reply_no: SeqNo, + is_holding_lock: bool, } -impl<'a, TL: TokioTransportLayer, Packet, const CAP: usize> ReplyGuard<'a, TL, Packet, 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. - /// 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 async fn reply(self, packet: Packet) -> Option<(Packet, ReplyGuard<'a, TL, Packet, CAP>)> { - self.reply_with(|_, _| packet).await +impl<'a, TL: TokioLayer, Payload, const CAP: usize> ReplyGuard<'a, TL, Payload, CAP> { + fn new(seq: &'a SeqExTokio, app: TL, reply_no: SeqNo, seq_cst: bool) -> Self { + ReplyGuard { seq, app: Some(app), reply_no, is_holding_lock: seq_cst } } - pub async fn reply_with(mut self, packet: impl FnOnce(SeqNo, SeqNo) -> Packet) -> Option<(Packet, ReplyGuard<'a, TL, Packet, CAP>)> { - let (tx, rx) = oneshot::channel(); - let mut seq = self.origin.seq_ex.lock().unwrap(); + fn try_reply_with_inner(&mut self, app: TL, seq_cst: bool, packet_data: impl FnOnce(SeqNo, SeqNo) -> SendData) -> Option { + let mut seq = self.seq.seq_ex.lock().unwrap(); let seq_no = seq.0.seq_no(); - let app = self.app.take().unwrap(); let pre_ts = seq.0.next_service_timestamp; - seq.0.reply_raw(app, self.reply_no, (tx, packet(seq_no, self.reply_no))); - if seq.0.next_service_timestamp != pre_ts { - let _ = app.update_queue.send(seq.0.next_service_timestamp).await; + seq.0 + .reply_raw(app, self.reply_no, self.is_holding_lock, seq_cst, packet_data(seq_no, self.reply_no)); + let ret = (pre_ts != seq.0.next_service_timestamp).then_some(seq.0.next_service_timestamp); + if seq.2 { + seq.2 = false; + drop(seq); + self.seq.reply_block.notify_waiters(); } - - let (packet, reply_no) = rx.await.ok()?; - let g = ReplyGuard { origin: self.origin, app: Some(app), reply_no }; - Some((packet, g)) - } -} -impl<'a, TL: TokioTransportLayer, Packet, const CAP: usize> Drop for ReplyGuard<'a, TL, Packet, CAP> { - fn drop(&mut self) { - if let Some(app) = self.app.take() { - let mut seq = self.origin.seq_ex.lock().unwrap(); - seq.0.ack_raw(app, self.reply_no); - } - } -} - -pub struct ReplyIter<'a, TL: TokioTransportLayer, Packet, const CAP: usize = DEFAULT_WINDOW_CAP> { - origin: Option<&'a SeqExTokio>, - app: &'a TokioTransport, - first: Option<(Packet, ReplyGuard<'a, TL, Packet, CAP>)>, -} - -//pub struct SeqExGuard<'a, SendData, RecvData, const CAP: usize>(MutexGuard<'a, (SeqEx, usize)>); -//impl<'a, SendData, RecvData, const CAP: usize> Deref for SeqExGuard<'a, SendData, RecvData, CAP> { -// type Target = SeqEx; - -// fn deref(&self) -> &Self::Target { -// &self.0 .0 -// } -//} -//impl<'a, SendData, RecvData, const CAP: usize> DerefMut for SeqExGuard<'a, SendData, RecvData, CAP> { -// fn deref_mut(&mut self) -> &mut Self::Target { -// &mut self.0 .0 -// } -//} -#[derive(Clone)] -pub struct TokioTransport { - time: time::Instant, - update_queue: mpsc::Sender, - app: TL, -} -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 task_tl = ret.clone(); - task::spawn(async move { - let mut update_ts = i64::MAX; - loop { - if update_ts < i64::MAX { - let diff = update_ts - task_tl.time.elapsed().as_millis() as i64; - let mut do_update = diff <= 0; - if diff > 0 { - let sleep = time::sleep(time::Duration::from_millis(diff as u64)); - tokio::select! { - Some(up) = recv.recv() => { - update_ts = up; - } - _ = sleep => { - do_update = true; - } - }; - } - if do_update { - let mut seq = seq.as_ref().seq_ex.lock().unwrap(); - seq.0.service(&task_tl); - update_ts = seq.0.next_service_timestamp; - } - } else if let Some(up) = recv.recv().await { - update_ts = up; - } - } - }); ret } -} + ///// 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. + ///// 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 async fn reply(self, seq_cst: bool, packet_data: Payload) -> Result<(ReplyGuard<'a, TL, Payload, CAP>, Payload), AsyncError> { + self.reply_with(seq_cst, |_, _| packet_data).await + } + pub async fn reply_with( + mut self, + seq_cst: bool, + packet_data: impl FnOnce(SeqNo, SeqNo) -> Payload, + ) -> Result<(ReplyGuard<'a, TL, Payload, CAP>, Payload), AsyncError> { + let app = self.app.take().unwrap(); + let (tx, rx) = oneshot::channel(); -impl SeqExTokio { - pub fn new(retry_interval: i64, initial_seq_no: SeqNo) -> Self { - Self { - seq_ex: Mutex::new((SeqEx::new(retry_interval, initial_seq_no), 0)), - send_block: Notify::new(), + let update_ts = self.try_reply_with_inner(app.clone(), seq_cst, |s, r| (tx, packet_data(s, r))); + + if let Some(update_ts) = update_ts { + let _ = self.seq.update_queue.send(update_ts).await; + } + let (reply_no, seq_cst, recv_data) = rx.await.map_err(|_| AsyncError::SeqExClosed)?.ok_or(AsyncError::ReceivedAck)?; + Ok((Self::new(self.seq, app, reply_no, seq_cst), recv_data)) + } +} +impl<'a, TL: TokioLayer, Payload, const CAP: usize> Drop for ReplyGuard<'a, TL, Payload, 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_waiters(); + } } } +} +impl<'a, TL: TokioLayer, Payload, const CAP: usize> std::fmt::Debug for ReplyGuard<'a, TL, Payload, CAP> { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ReplyGuard") + .field("reply_no", &self.reply_no) + .field("is_holding_lock", &self.is_holding_lock) + .finish() + } +} - fn process<'a, TL: TokioTransportLayer>( +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum AsyncRecvError { + DroppedTooEarly, + DroppedDuplicate, + WaitingForRecv, + WaitingForReply, + AsyncReply, +} +impl std::fmt::Display for AsyncRecvError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + AsyncRecvError::DroppedTooEarly => write!(f, "packet arrived too early"), + AsyncRecvError::DroppedDuplicate => write!(f, "packet was a duplicate"), + AsyncRecvError::WaitingForRecv => write!(f, "can't process until another packet is received"), + AsyncRecvError::WaitingForReply => write!(f, "can't process until a reply is finished"), + AsyncRecvError::AsyncReply => write!(f, "packet was an async reply"), + } + } +} +impl std::error::Error for AsyncRecvError {} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum AsyncError { + ReceivedAck, + SeqExClosed, +} +impl std::fmt::Display for AsyncError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + AsyncError::ReceivedAck => write!(f, "peer replied with an ack"), + AsyncError::SeqExClosed => write!(f, "the window was closed before a reply could be received"), + } + } +} +impl std::error::Error for AsyncError {} + +pub struct AsyncRecvIter<'a, TL: TokioLayer, P: Into, Payload, const CAP: usize = DEFAULT_WINDOW_CAP> { + seq: Option<&'a SeqExTokio>, + app: TL, + first: Option<(ReplyGuard<'a, TL, Payload, CAP>, P)>, +} +pub struct RecvIter<'a, TL: TokioLayer, P: Into, Payload, const CAP: usize = DEFAULT_WINDOW_CAP> { + seq: Option<&'a SeqExTokio>, + app: TL, + first: Option<(ReplyGuard<'a, TL, Payload, CAP>, P)>, +} + +pub struct ServiceState { + next_service_timestamp: i64, + recv_service_update: mpsc::Receiver, +} + +impl SeqExTokio { + pub fn new(retry_interval: i64, initial_seq_no: SeqNo) -> (Self, ServiceState) { + let (update_queue, recv_service_update) = mpsc::channel(8); + ( + Self { + seq_ex: Mutex::new((SeqEx::new(retry_interval, initial_seq_no), 0, false)), + send_block: Notify::new(), + reply_block: Notify::new(), + update_queue, + }, + ServiceState { next_service_timestamp: i64::MAX, recv_service_update }, + ) + } + pub fn new_default() -> (Self, ServiceState) { + Self::new(DEFAULT_RESEND_INTERVAL_MS, DEFAULT_INITIAL_SEQ_NO) + } + + fn map_raw_or_reply<'a, TL: TokioLayer, P: Into>( &'a self, - app: &'a TokioTransport, - mut seq: MutexGuard<'_, (SeqEx, Packet, CAP>, usize)>, - 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 { - self.send_block.notify_one(); - } - if let Some((tx, _)) = send_data { - if let Err(_) = tx.send((packet, reply_no)) { - // Allow the drop code to be run - seq.0.ack_raw(app, reply_no); + app: TL, + seq: &mut SeqEx, Payload, CAP>, + ret: RecvOkRaw, P>, + ) -> Option<(ReplyGuard<'a, TL, Payload, CAP>, P)> { + match ret { + RecvOkRaw::Payload { reply_no, seq_cst, recv_data } => Some((ReplyGuard::new(self, app, reply_no, seq_cst), recv_data)), + RecvOkRaw::Reply { reply_no, seq_cst, recv_data, send_data: (tx, _) } => { + if tx.send(Some((reply_no, seq_cst, recv_data.into()))).is_err() { + // Send an ack if no one is receiving the reply on the other end. + // Could occur if the future holding the receiver is dropped. + seq.ack_raw(app, reply_no, seq_cst); } None - } else { - Some((packet, ReplyGuard { origin: self, app: Some(app), reply_no })) } - } else { - None - } - } - pub fn receive<'a, TL: TokioTransportLayer>( - &'a self, - app: &'a TokioTransport, - seq_no: SeqNo, - reply_no: Option, - packet: Packet, - ) -> Option<(Packet, ReplyGuard<'_, TL, Packet, CAP>)> { - let mut seq = self.seq_ex.lock().unwrap(); - 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>)> { - let mut seq = self.seq_ex.lock().unwrap(); - let result = seq.0.pump_raw(); - self.process(app, seq, result) - } - pub fn receive_all<'a, TL: TokioTransportLayer>( - &'a self, - app: &'a TokioTransport, - seq_no: SeqNo, - reply_no: Option, - packet: Packet, - ) -> ReplyIter<'_, TL, Packet, CAP> { - if let Some(g) = self.receive(app, seq_no, reply_no, packet) { - ReplyIter { origin: Some(self), app, first: Some(g) } - } else { - ReplyIter { origin: None, app, first: None } - } - } - pub fn receive_ack(&self, reply_no: SeqNo) { - let mut seq = self.seq_ex.lock().unwrap(); - // We drop the sender to notify the receiver that no packet was received. - if let Ok(_) = seq.0.receive_ack(reply_no) { - if seq.1 > 0 { - self.send_block.notify_one(); + RecvOkRaw::Ack { send_data: (tx, _) } => { + drop(tx); + None } } } - //pub fn try_send>(&self, app: TL, packet_data: SendData) -> Result<(), SendData> { - // let mut seq = self.lock(); - // seq.try_send(app, packet_data) - //} - async fn send_inner>( - &self, - mut seq: MutexGuard<'_, (SeqEx, Packet, CAP>, usize)>, - app: &TokioTransport, - mut tx: oneshot::Sender<(Packet, SeqNo)>, - mut packet: Packet, - ) { - let mut pre_ts = seq.0.next_service_timestamp; - while let Err(e) = seq.0.try_send(app, (tx, packet)) { - (tx, packet) = e; - seq.1 += 1; - drop(seq); - self.send_block.notified().await; - seq = self.seq_ex.lock().unwrap(); - pre_ts = seq.0.next_service_timestamp; - seq.1 -= 1; - } - if seq.0.next_service_timestamp != pre_ts { - let _ = app.update_queue.send(seq.0.next_service_timestamp).await; - } - } - /// 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>)> { - self.send_with(app, |_| packet).await - } - //pub fn try_send_with>(&self, app: TL, 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 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(); - self.send_inner(seq, app, tx, packet(seq_no)).await; - // This can only return an error if the sender was dropped. - let (packet, reply_no) = rx.await.ok()?; - Some((packet, ReplyGuard { origin: self, app: Some(app), reply_no })) - } - //pub fn lock(&self) -> SeqExGuard<'_, SendData, RecvData, CAP> { - // SeqExGuard(self.seq_ex.lock().unwrap()) - //} -} -//impl Default for SeqExTokio { -// fn default() -> Self { -// Self::new(DEFAULT_RESEND_INTERVAL_MS, DEFAULT_INITIAL_SEQ_NO) -// } -//} -impl<'a, TL: TokioTransportLayer, Packet, const CAP: usize> Iterator for ReplyIter<'a, TL, Packet, CAP> { - type Item = (Packet, ReplyGuard<'a, TL, Packet, CAP>); + pub fn receive, P: Into>( + &self, + app: TL, + packet: Packet

, + ) -> Result<(ReplyGuard<'_, TL, Payload, CAP>, P), AsyncRecvError> { + let mut seq = self.seq_ex.lock().unwrap(); + return match seq.0.receive_raw(app.clone(), packet) { + Ok(ret) => { + if seq.1 > 0 { + seq.1 -= 1; + self.send_block.notify_one(); + } + self.map_raw_or_reply(app, &mut seq.0, ret).ok_or(AsyncRecvError::AsyncReply) + } + Err(crate::RecvError::DroppedTooEarly) => Err(AsyncRecvError::DroppedTooEarly), + Err(crate::RecvError::DroppedDuplicate) => Err(AsyncRecvError::DroppedDuplicate), + Err(crate::RecvError::WaitingForRecv) => Err(AsyncRecvError::WaitingForRecv), + Err(crate::RecvError::WaitingForReply) => Err(AsyncRecvError::WaitingForReply), + }; + } + + fn try_pump_inner>( + &self, + app: TL, + blocking: bool, + ) -> Result<(ReplyGuard<'_, TL, Payload, CAP>, Payload), PumpError> { + let mut seq = self.seq_ex.lock().unwrap(); + // Enforce that only one thread may pump at a time. + loop { + match seq.0.try_pump_raw() { + Ok(ret) => { + if seq.1 > 0 { + seq.1 -= 1; + self.send_block.notify_one(); + } + if let Some(ret) = self.map_raw_or_reply(app.clone(), &mut seq.0, ret) { + return Ok(ret); + } + } + Err(PumpError::WaitingForRecv) => return Err(PumpError::WaitingForRecv), + Err(PumpError::WaitingForReply) => { + seq.2 |= blocking; + return Err(PumpError::WaitingForReply); + } + } + } + } + pub fn try_pump>(&self, app: TL) -> Result<(ReplyGuard<'_, TL, Payload, CAP>, Payload), PumpError> { + 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) => { + self.reply_block.notified().await; + } + } + } + } + pub fn try_receive_all, P: Into>( + &self, + app: TL, + packet: Packet

, + ) -> RecvIter<'_, TL, P, Payload, CAP> { + match self.receive(app.clone(), packet) { + Ok(ret) => RecvIter { seq: Some(self), app, first: Some(ret) }, + Err(AsyncRecvError::AsyncReply) => RecvIter { seq: Some(self), app, first: None }, + Err(_) => RecvIter { seq: None, app, first: None }, + } + } + pub fn receive_all, P: Into>( + &self, + app: TL, + packet: Packet

, + ) -> AsyncRecvIter<'_, TL, P, Payload, CAP> { + match self.receive(app.clone(), packet) { + Ok(ret) => AsyncRecvIter { seq: Some(self), app, first: Some(ret) }, + Err(AsyncRecvError::AsyncReply) => AsyncRecvIter { seq: Some(self), app, first: None }, + Err(_) => AsyncRecvIter { seq: None, app, first: None }, + } + } + + fn try_send_with_inner, F: FnOnce(SeqNo) -> SendData>( + &self, + app: TL, + blocking: bool, + seq_cst: bool, + packet_data: F, + ) -> Result, F> { + let mut seq = self.seq_ex.lock().unwrap(); + let pre_ts = seq.0.next_service_timestamp; + let result = seq.0.try_send_with(app, seq_cst, packet_data); + if let Err(e) = result { + seq.1 += blocking as usize; + Err(e) + } else { + Ok((pre_ts != seq.0.next_service_timestamp).then_some(seq.0.next_service_timestamp)) + } + } + + pub async fn send>( + &self, + app: TL, + seq_cst: bool, + packet_data: Payload, + ) -> Result<(ReplyGuard<'_, TL, Payload, CAP>, Payload), AsyncError> { + self.send_with(app, seq_cst, |_| packet_data).await + } + pub async fn send_with>( + &self, + app: TL, + seq_cst: bool, + packet_data: impl FnOnce(SeqNo) -> Payload, + ) -> Result<(ReplyGuard<'_, TL, Payload, CAP>, Payload), AsyncError> { + let (rx, tx) = oneshot::channel(); + let mut pf = |s| (rx, packet_data(s)); + loop { + let ret = self.try_send_with_inner(app.clone(), true, seq_cst, pf); + match ret { + Ok(update) => { + if let Some(update) = update { + let _ = self.update_queue.send(update).await; + } + let (reply_no, locked, recv_data) = tx.await.map_err(|_| AsyncError::SeqExClosed)?.ok_or(AsyncError::ReceivedAck)?; + return Ok((ReplyGuard::new(self, app, reply_no, locked), recv_data)); + } + Err(p) => { + pf = p; + self.send_block.notified().await; + } + } + } + } + + /// This function must be called with the same ServiceState instance returned upon creation of + /// the given SeqExTokio instance. + pub async fn service_task>(&self, mut app: TL, state: &mut ServiceState) { + let mut result = None; + if state.next_service_timestamp < i64::MAX { + let diff = state.next_service_timestamp - app.time(); + if diff > 0 { + if let Ok(up) = time::timeout(time::Duration::from_millis(diff as u64), state.recv_service_update.recv()).await { + result = up + } + } + } else { + result = state.recv_service_update.recv().await + }; + + if let Some(up) = result { + state.next_service_timestamp = state.next_service_timestamp.min(up); + } else { + let mut seq = self.seq_ex.lock().unwrap(); + seq.0.service(app.clone()); + state.next_service_timestamp = seq.0.next_service_timestamp; + } + } +} + +impl<'a, TL: TokioLayer, P: Into, Payload, const CAP: usize> RecvIter<'a, TL, P, Payload, CAP> { + pub fn take_first(&mut self) -> Option<(ReplyGuard<'a, TL, Payload, CAP>, P)> { + self.first.take() + } +} +impl<'a, TL: TokioLayer, P: Into, Payload, const CAP: usize> Iterator for RecvIter<'a, TL, P, Payload, CAP> { + type Item = (ReplyGuard<'a, TL, Payload, CAP>, Payload); + fn next(&mut self) -> Option { if let Some(g) = self.first.take() { - Some(g) - } else if let Some(origin) = self.origin { - origin.pump(self.app) + Some((g.0, g.1.into())) + } else if let Some(seq) = self.seq { + seq.try_pump(self.app.clone()).ok() + } else { + None + } + } +} +impl<'a, TL: TokioLayer, P: Into, Payload, const CAP: usize> AsyncRecvIter<'a, TL, P, Payload, CAP> { + pub fn take_first(&mut self) -> Option<(ReplyGuard<'a, TL, Payload, CAP>, P)> { + self.first.take() + } + + pub async fn next(&mut self) -> Option<(ReplyGuard<'a, TL, Payload, CAP>, Payload)> { + if let Some(g) = self.first.take() { + Some((g.0, g.1.into())) + } else if let Some(seq) = self.seq { + seq.pump(self.app.clone()).await } else { None } } } -pub trait TokioTransportLayer: Clone + Send + 'static { - type Packet; +pub trait TokioLayer: Clone { + type Payload; - fn send(&self, seq_no: SeqNo, reply_no: Option, payload: &Self::Packet); - fn send_ack(&self, reply_no: SeqNo); + fn time(&mut self) -> i64; + + fn send(&mut self, packet: Packet<&Self::Payload>); } -impl TransportLayer> for &TokioTransport { +impl TransportLayer> for TL { + fn time(&mut self) -> i64 { + self.time() + } + fn send(&mut self, packet: Packet<&SendData>) { + self.send(packet.map(|p| &p.1)) + } +} + +#[derive(Clone, Debug)] +pub struct MpscTransport { + pub channel: mpsc::Sender>, + pub time: time::Instant, +} + +impl MpscTransport { + pub fn new(buffer: usize) -> (Self, mpsc::Receiver>) { + let (send, recv) = mpsc::channel(buffer); + (Self { channel: send, time: time::Instant::now() }, recv) + } + pub fn from_sender(send: mpsc::Sender>) -> Self { + Self { channel: send, time: time::Instant::now() } + } +} +impl TokioLayer for &MpscTransport { + type Payload = Payload; + fn time(&mut self) -> i64 { self.time.elapsed().as_millis() as i64 } - fn send(&mut self, seq_no: SeqNo, reply_no: Option, (_, payload): &SendData) { - self.app.send(seq_no, reply_no, payload); - } - fn send_ack(&mut self, reply_no: SeqNo) { - self.app.send_ack(reply_no) + fn send(&mut self, packet: Packet<&Payload>) { + let _ = self.channel.try_send(packet.cloned()); } } -//#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] -//#[derive(Clone)] -//pub enum PacketType { -// Payload(SeqNo, Option, Payload), -// Ack(SeqNo), -//} - -//#[derive(Clone)] -//pub struct MpscTransport { -// pub channel: Sender>, -// pub time: Instant, -//} -//pub type MpscGuard<'a, Packet> = ReplyGuard<'a, &'a MpscTransport, Packet, Packet>; -//pub type MpscSeqEx = SeqExTokio; - -//impl MpscTransport { -// 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 { -// Self { channel: send, time: std::time::Instant::now() } -// } -//} -//impl TransportLayer for &MpscTransport { -// fn time(&mut self) -> i64 { -// 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)); -// } -//}