mod webrtc_data_channel;
use crate::nq_core::ScopedHeaders;
#[cfg(not(test))]
use crate::nq_core::{ConnectionType, Network, Time, TokioTime, client::Direction};
#[cfg(not(test))]
use crate::nq_load_generator::LoadedConnection;
use crate::nq_load_generator::{LoadConfig, LoadGenerator};
#[cfg(not(test))]
use crate::nq_tokio_network::TokioNetwork;
use serde::{Deserialize, Serialize};
use std::{cmp::min, collections::HashMap, fmt::Display, sync::Arc, time::Duration};
use tokio::sync::{RwLock, mpsc};
use tokio_util::sync::CancellationToken;
#[cfg(not(test))]
use tracing::Instrument;
use url::Url;
use webrtc_data_channel::{DataChannelEvent, WebRTCDataChannel};
#[derive(Clone, Debug)]
pub struct PacketLossConfig {
pub turn_server_uri: String,
pub turn_cred_request_url: Url,
pub num_packets: usize,
pub batch_size: usize,
pub batch_wait_time: Duration,
pub response_wait_time: Duration,
pub download_url: Url,
pub upload_url: Url,
pub scoped_headers: Option<ScopedHeaders>,
}
impl Default for PacketLossConfig {
fn default() -> Self {
Self {
turn_server_uri: "turn:turn.speed.cloudflare.com:50000?transport=udp".to_owned(),
turn_cred_request_url: "https://speed.cloudflare.com/turn-creds".parse().unwrap(),
num_packets: 1000,
batch_size: 10,
batch_wait_time: Duration::from_millis(10),
response_wait_time: Duration::from_millis(3000),
download_url: "https://h3.speed.cloudflare.com/__down?bytes=10000000000"
.parse()
.unwrap(),
upload_url: "https://h3.speed.cloudflare.com/__up".parse().unwrap(),
scoped_headers: None,
}
}
}
impl PacketLossConfig {
pub fn load_config(&self) -> LoadConfig {
LoadConfig {
headers: HashMap::default(),
scoped_headers: self.scoped_headers.clone(),
download_url: self.download_url.clone(),
upload_url: self.upload_url.clone(),
}
}
}
pub struct PacketLoss {
config: Arc<PacketLossConfig>,
load_generator: LoadGenerator,
message_tracker: Arc<RwLock<Vec<bool>>>,
}
impl PacketLoss {
pub fn new_with_config(config: PacketLossConfig) -> anyhow::Result<Self> {
let load_generator = LoadGenerator::new(config.load_config())?;
let message_vec = vec![false; config.num_packets];
let message_tracker: Arc<RwLock<Vec<bool>>> = Arc::new(RwLock::new(message_vec));
Ok(Self {
config: Arc::new(config),
load_generator,
message_tracker,
})
}
#[cfg(not(test))]
fn add_load_generators(
&self,
packet_event_tx: mpsc::Sender<PacketLossEvent>,
shutdown: CancellationToken,
) -> anyhow::Result<()> {
let time = Arc::new(TokioTime::new()) as Arc<dyn Time>;
let network =
Arc::new(TokioNetwork::new(Arc::clone(&time), shutdown.clone())) as Arc<dyn Network>;
self.new_load_generating_connection(
packet_event_tx.clone(),
Direction::Down,
network.clone(),
time.clone(),
shutdown.clone(),
)?;
self.new_load_generating_connection(
packet_event_tx.clone(),
Direction::Up(4_000_000_000),
network,
time,
shutdown.clone(),
)?;
Ok(())
}
#[cfg_attr(test, allow(unused_mut))]
pub async fn run_test(
mut self,
turn_server_creds: TurnServerCreds,
shutdown: CancellationToken,
) -> anyhow::Result<PacketLossResult> {
let (webrtc_event_tx, mut webrtc_event_rx) =
tokio::sync::mpsc::unbounded_channel::<DataChannelEvent>();
let (packet_event_tx, mut packet_event_rx) =
tokio::sync::mpsc::channel::<PacketLossEvent>(3);
#[cfg(not(test))]
self.add_load_generators(packet_event_tx.clone(), shutdown.clone())?;
let mut webrtc_data_channel = WebRTCDataChannel::create_with_config(
&self.config.turn_server_uri,
&turn_server_creds,
webrtc_event_tx.clone(),
)
.await?;
webrtc_data_channel.establish_data_channel().await?;
loop {
tokio::select! {
Some(event) = webrtc_event_rx.recv() => {
match event {
DataChannelEvent::OnOpenChannel => {
self.send_messages(webrtc_data_channel.clone(), packet_event_tx.clone(), shutdown.clone());
}
DataChannelEvent::OnReceivedMessage(message) => {
self.message_tracker
.write()
.await
.insert(message, true);
}
DataChannelEvent::ConnectionError(err) => {
tracing::warn!("Failed to complete packet loss test {}", err);
break;
}
}
}
Some(event) = packet_event_rx.recv() => {
match event {
PacketLossEvent::AllMessagesSent => {
break;
}
#[cfg(not(test))]
PacketLossEvent::NewLoadedConnection(connection) => {
self.load_generator.push(connection);
}
#[cfg(not(test))]
PacketLossEvent::Error(err) => {
tracing::warn!("Failed to complete packet loss test {}", err);
break;
}
}
}
_ = shutdown.cancelled() => {
tracing::debug!("Shutdown requested");
break;
}
}
}
let mut loads = self.load_generator.into_connections();
loads.iter_mut().for_each(|load| load.stop());
webrtc_data_channel.close_channel().await;
let num_messages = self
.message_tracker
.read()
.await
.iter()
.filter(|val| **val)
.count();
let loss_ratio = (((self.config.num_packets - num_messages) as f64
/ self.config.num_packets as f64)
* 10_000.0)
.trunc()
/ 100.0;
Ok(PacketLossResult {
num_messages,
loss_ratio,
})
}
fn send_messages(
&self,
mut webrtc_data_channel: WebRTCDataChannel,
send_message_tx: mpsc::Sender<PacketLossEvent>,
shutdown: CancellationToken,
) {
let config = self.config.clone();
tokio::spawn(async move {
let mut message_count = 0;
loop {
let batch_start = message_count;
let batch_count = if config.batch_size == 0 {
config.num_packets
} else {
min(message_count + config.batch_size, config.num_packets)
};
for i in batch_start..batch_count {
if let Err(err) = webrtc_data_channel.send_message(&i.to_be_bytes()).await {
tracing::warn!("Send message failed: {}", err);
}
message_count += 1;
}
if (message_count + 1) >= config.num_packets {
tokio::time::sleep(config.response_wait_time).await;
let _ = send_message_tx.send(PacketLossEvent::AllMessagesSent).await;
break;
}
tokio::time::sleep(config.batch_wait_time).await;
if shutdown.is_cancelled() {
tracing::debug!("Shutdown requested");
break;
}
}
});
}
#[cfg(not(test))]
#[tracing::instrument(skip_all)]
fn new_load_generating_connection(
&self,
event_tx: mpsc::Sender<PacketLossEvent>,
direction: Direction,
network: Arc<dyn Network>,
time: Arc<dyn Time>,
shutdown: CancellationToken,
) -> anyhow::Result<()> {
let oneshot_res = self.load_generator.new_loaded_connection(
direction,
ConnectionType::H2,
network,
time,
shutdown,
)?;
tokio::spawn(
async move {
let _ = match oneshot_res.await {
Ok(conn) => event_tx.send(PacketLossEvent::NewLoadedConnection(conn)),
Err(err) => event_tx.send(PacketLossEvent::Error(err)),
}
.await;
}
.in_current_span(),
);
Ok(())
}
}
#[derive(Debug, Serialize, Deserialize, PartialEq)]
pub struct PacketLossResult {
pub num_messages: usize,
pub loss_ratio: f64,
}
impl Display for PacketLossResult {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
writeln!(
f,
"messages: {} loss ratio: {:.2}",
self.num_messages, self.loss_ratio,
)
}
}
pub enum PacketLossEvent {
AllMessagesSent,
#[cfg(not(test))]
NewLoadedConnection(LoadedConnection),
#[cfg(not(test))]
Error(anyhow::Error),
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct TurnServerCreds {
pub username: String,
pub credential: String,
}
#[cfg(test)]
mod tests {
use crate::nq_packetloss::{
PacketLoss, PacketLossConfig, PacketLossResult, webrtc_data_channel::tests::TestTurnServer,
};
use std::time::Duration;
use tokio_util::sync::CancellationToken;
async fn run_test(
num_packets: usize,
batch_size: usize,
batch_wait_time: u64,
response_wait_time: u64,
) -> anyhow::Result<PacketLossResult> {
let shutdown = CancellationToken::new();
let server = TestTurnServer::start_turn_server().await?;
let config = PacketLossConfig {
turn_server_uri: format!("turn:127.0.0.1:{}?transport=udp", server.server_port),
turn_cred_request_url: "https://127.0.0.1/creds".parse().unwrap(),
num_packets,
batch_size,
batch_wait_time: Duration::from_millis(batch_wait_time),
response_wait_time: Duration::from_millis(response_wait_time),
download_url: "https://h3.speed.cloudflare.com/__down?bytes=10000000000"
.parse()
.unwrap(),
upload_url: "https://h3.speed.cloudflare.com/__up".parse().unwrap(),
scoped_headers: None,
};
let packet_loss = PacketLoss::new_with_config(config)?;
let packet_loss_result = packet_loss
.run_test(server.get_test_creds(), shutdown)
.await;
server.close().await?;
packet_loss_result
}
#[tokio::test]
async fn test_with_server() -> anyhow::Result<()> {
let packet_loss_result = run_test(1000, 10, 10, 100).await?;
assert_eq!(
PacketLossResult {
num_messages: 1000,
loss_ratio: 0.0,
},
packet_loss_result
);
println!("Results: {:?}", packet_loss_result);
Ok(())
}
#[tokio::test]
async fn test_no_batch_size() -> anyhow::Result<()> {
let packet_loss_result = run_test(1000, 0, 10, 100).await?;
assert_eq!(
PacketLossResult {
num_messages: 1000,
loss_ratio: 0.0,
},
packet_loss_result
);
println!("Results: {:?}", packet_loss_result);
Ok(())
}
#[tokio::test]
async fn test_too_large_batch_size() -> anyhow::Result<()> {
let packet_loss_result = run_test(100, 1000, 10, 100).await?;
assert_eq!(
PacketLossResult {
num_messages: 100,
loss_ratio: 0.0,
},
packet_loss_result
);
println!("Results: {:?}", packet_loss_result);
Ok(())
}
#[tokio::test]
async fn test_equal_batch_size() -> anyhow::Result<()> {
let packet_loss_result = run_test(100, 100, 10, 100).await?;
assert_eq!(
PacketLossResult {
num_messages: 100,
loss_ratio: 0.0,
},
packet_loss_result
);
println!("Results: {:?}", packet_loss_result);
Ok(())
}
#[tokio::test]
async fn test_zero_batch_wait() -> anyhow::Result<()> {
let packet_loss_result = run_test(100, 10, 0, 100).await?;
assert_eq!(
PacketLossResult {
num_messages: 100,
loss_ratio: 0.0,
},
packet_loss_result
);
println!("Results: {:?}", packet_loss_result);
Ok(())
}
#[tokio::test]
async fn test_cancel() -> anyhow::Result<()> {
let shutdown = CancellationToken::new();
let server = TestTurnServer::start_turn_server().await?;
let config = PacketLossConfig {
turn_server_uri: format!("turn:127.0.0.1:{}?transport=udp", server.server_port),
turn_cred_request_url: "https://127.0.0.1/creds".parse().unwrap(),
num_packets: 5000,
batch_size: 10,
batch_wait_time: Duration::from_millis(25),
response_wait_time: Duration::from_millis(100),
download_url: "https://h3.speed.cloudflare.com/__down?bytes=10000000000"
.parse()
.unwrap(),
upload_url: "https://h3.speed.cloudflare.com/__up".parse().unwrap(),
scoped_headers: None,
};
let packet_loss = PacketLoss::new_with_config(config)?;
let shutdown_clone = shutdown.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(5000)).await;
shutdown_clone.cancel();
});
let packet_loss_result = packet_loss
.run_test(server.get_test_creds(), shutdown)
.await;
server.close().await?;
let packet_loss_result = packet_loss_result?;
assert_ne!(
PacketLossResult {
num_messages: 5000,
loss_ratio: 0.0,
},
packet_loss_result
);
println!("Results: {:?}", packet_loss_result);
Ok(())
}
}