File
Blob: firmware/vendor/str0m/tests/handshake-direct.rs
| 1 | use std::net::{Ipv4Addr, SocketAddr}; |
| 2 | use std::sync::Arc; |
| 3 | use std::sync::atomic::{AtomicUsize, Ordering}; |
| 4 | use std::sync::mpsc::{self, Receiver, Sender}; |
| 5 | use std::thread; |
| 6 | use std::time::{Duration, Instant}; |
| 7 | |
| 8 | use str0m::channel::{ChannelConfig, ChannelId, Reliability}; |
| 9 | use str0m::config::{DtlsVersion, Fingerprint}; |
| 10 | use str0m::crypto::dtls::ProtocolVersion; |
| 11 | use str0m::ice::IceCreds; |
| 12 | use str0m::net::{Protocol, Receive}; |
| 13 | use str0m::{Candidate, Event, IceConnectionState, Input, Output, Rtc, RtcConfig, RtcError}; |
| 14 | use tracing::{Span, info_span}; |
| 15 | |
| 16 | mod common; |
| 17 | use common::{Peer, init_crypto_default, init_log, snap_init_data}; |
| 18 | |
| 19 | /// Maximum total packets expected when using SNAP. |
| 20 | /// |
| 21 | /// SNAP skips the 4-way SCTP handshake (INIT, INIT-ACK, COOKIE-ECHO, |
| 22 | /// COOKIE-ACK), so the total packet count should be well under this limit. |
| 23 | /// The exact number depends on ICE/DTLS setup specifics. |
| 24 | const MAX_SNAP_PACKETS: usize = 20; |
| 25 | |
| 26 | /// Pre-negotiated data channel SCTP stream ID |
| 27 | const DATA_CHANNEL_ID: u16 = 0; |
| 28 | |
| 29 | /// Set to `true` to save packet captures to `target/pcap/` for Wireshark analysis. |
| 30 | const SAVE_PCAP: bool = false; |
| 31 | |
| 32 | #[test] |
| 33 | pub fn handshake_dtls_auto_to_12() -> Result<(), RtcError> { |
| 34 | run_handshake_test(DtlsVersion::Auto, DtlsVersion::Dtls12) |
| 35 | } |
| 36 | |
| 37 | #[test] |
| 38 | pub fn handshake_dtls_auto_to_13() -> Result<(), RtcError> { |
| 39 | run_handshake_test(DtlsVersion::Auto, DtlsVersion::Dtls13) |
| 40 | } |
| 41 | |
| 42 | #[test] |
| 43 | pub fn handshake_dtls_auto_to_auto() -> Result<(), RtcError> { |
| 44 | run_handshake_test(DtlsVersion::Auto, DtlsVersion::Auto) |
| 45 | } |
| 46 | |
| 47 | #[test] |
| 48 | pub fn handshake_dtls_12_to_auto() -> Result<(), RtcError> { |
| 49 | run_handshake_test(DtlsVersion::Dtls12, DtlsVersion::Auto) |
| 50 | } |
| 51 | |
| 52 | #[test] |
| 53 | pub fn handshake_dtls_13_to_auto() -> Result<(), RtcError> { |
| 54 | run_handshake_test(DtlsVersion::Dtls13, DtlsVersion::Auto) |
| 55 | } |
| 56 | |
| 57 | #[test] |
| 58 | pub fn handshake_dtls_12_to_12() -> Result<(), RtcError> { |
| 59 | run_handshake_test(DtlsVersion::Dtls12, DtlsVersion::Dtls12) |
| 60 | } |
| 61 | |
| 62 | #[test] |
| 63 | pub fn handshake_dtls_13_to_13() -> Result<(), RtcError> { |
| 64 | run_handshake_test(DtlsVersion::Dtls13, DtlsVersion::Dtls13) |
| 65 | } |
| 66 | |
| 67 | /// Standard direct API handshake (no SNAP). |
| 68 | #[test] |
| 69 | pub fn handshake_direct_api() -> Result<(), RtcError> { |
| 70 | run_direct_handshake(DtlsVersion::Auto, DtlsVersion::Auto, false) |
| 71 | } |
| 72 | |
| 73 | /// Direct API handshake with SNAP (out-of-band SCTP INIT exchange, skips 4-way handshake). |
| 74 | #[test] |
| 75 | pub fn handshake_direct_api_snap() -> Result<(), RtcError> { |
| 76 | run_direct_handshake(DtlsVersion::Auto, DtlsVersion::Auto, true) |
| 77 | } |
| 78 | |
| 79 | /// Returns the name of the default crypto provider based on compile-time feature flags. |
| 80 | /// Mirrors the priority order in `str0m::crypto::from_feature_flags()`. |
| 81 | #[allow(unreachable_code)] |
| 82 | fn default_crypto_name() -> &'static str { |
| 83 | #[cfg(feature = "aws-lc-rs")] |
| 84 | return "aws-lc-rs"; |
| 85 | #[cfg(feature = "rust-crypto")] |
| 86 | return "rust-crypto"; |
| 87 | #[cfg(feature = "openssl-dimpl")] |
| 88 | return "openssl-dimpl"; |
| 89 | #[cfg(feature = "openssl")] |
| 90 | return "openssl"; |
| 91 | #[cfg(feature = "wincrypto-dimpl")] |
| 92 | return "wincrypto-dimpl"; |
| 93 | #[cfg(all(feature = "wincrypto", target_os = "windows"))] |
| 94 | return "wincrypto"; |
| 95 | #[cfg(all(feature = "apple-crypto", target_vendor = "apple"))] |
| 96 | return "apple-crypto"; |
| 97 | "unknown" |
| 98 | } |
| 99 | |
| 100 | fn run_handshake_test(client_dtls: DtlsVersion, server_dtls: DtlsVersion) -> Result<(), RtcError> { |
| 101 | run_direct_handshake(client_dtls, server_dtls, false) |
| 102 | } |
| 103 | |
| 104 | /// Shared implementation for both standard and SNAP direct API handshake tests. |
| 105 | fn run_direct_handshake( |
| 106 | client_dtls: DtlsVersion, |
| 107 | server_dtls: DtlsVersion, |
| 108 | use_snap: bool, |
| 109 | ) -> Result<(), RtcError> { |
| 110 | init_log(); |
| 111 | init_crypto_default(); |
| 112 | |
| 113 | let test_start = Instant::now(); |
| 114 | |
| 115 | let client_crypto_name = |
| 116 | std::env::var("L_CRYPTO").unwrap_or_else(|_| default_crypto_name().into()); |
| 117 | let server_crypto_name = |
| 118 | std::env::var("R_CRYPTO").unwrap_or_else(|_| default_crypto_name().into()); |
| 119 | |
| 120 | // wincrypto and openssl only support DTLS 1.2 - skip tests requiring 1.3/Auto. |
| 121 | // Also skip Auto client -> 1.2-only server: dimpl advertises X25519 in the hybrid |
| 122 | // ClientHello but its DTLS 1.2 engine can't process X25519 in ServerKeyExchange. |
| 123 | let dtls12_only = |name: &str| matches!(name, "wincrypto" | "openssl"); |
| 124 | let needs_13 = |v: DtlsVersion| matches!(v, DtlsVersion::Auto | DtlsVersion::Dtls13); |
| 125 | |
| 126 | if (dtls12_only(&client_crypto_name) && needs_13(client_dtls)) |
| 127 | || (dtls12_only(&server_crypto_name) && needs_13(server_dtls)) |
| 128 | || (matches!(client_dtls, DtlsVersion::Auto) && dtls12_only(&server_crypto_name)) |
| 129 | { |
| 130 | println!( |
| 131 | "\n=== SKIPPED: client={} ({}), server={} ({}) - DTLS 1.3/Auto not supported ===", |
| 132 | client_dtls, client_crypto_name, server_dtls, server_crypto_name |
| 133 | ); |
| 134 | return Ok(()); |
| 135 | } |
| 136 | |
| 137 | println!( |
| 138 | "\n=== Test: client={} ({}), server={} ({}) ===", |
| 139 | client_dtls, client_crypto_name, server_dtls, server_crypto_name |
| 140 | ); |
| 141 | |
| 142 | // Channels for communication between threads |
| 143 | let (client_tx, server_rx) = mpsc::channel::<Message>(); |
| 144 | let (server_tx, client_rx) = mpsc::channel::<Message>(); |
| 145 | |
| 146 | let client_packets_sent = Arc::new(AtomicUsize::new(0)); |
| 147 | let server_packets_sent = Arc::new(AtomicUsize::new(0)); |
| 148 | let client_packets_sent_clone = client_packets_sent.clone(); |
| 149 | let server_packets_sent_clone = server_packets_sent.clone(); |
| 150 | |
| 151 | let client_addr: SocketAddr = (Ipv4Addr::new(192, 168, 1, 1), 5000).into(); |
| 152 | let server_addr: SocketAddr = (Ipv4Addr::new(192, 168, 1, 2), 5001).into(); |
| 153 | |
| 154 | // Test name for pcap files |
| 155 | let test_name = format!( |
| 156 | "handshake_dtls_{}_to_{}", |
| 157 | dtls_version_short(client_dtls), |
| 158 | dtls_version_short(server_dtls) |
| 159 | ); |
| 160 | |
| 161 | // Spawn server thread |
| 162 | // Returns (packets, Result) so pcap is available even on failure. |
| 163 | let server_handle = thread::spawn( |
| 164 | move || -> (Vec<PcapPacket>, Result<TimingReport, RtcError>) { |
| 165 | let span = info_span!("SERVER"); |
| 166 | let _guard = span.enter(); |
| 167 | let mut timing = TimingReport::new(); |
| 168 | let mut packets = Vec::new(); |
| 169 | |
| 170 | let result = (|| -> Result<TimingReport, RtcError> { |
| 171 | // Initialize server (is_client = false) |
| 172 | let (mut rtc, local_creds, local_fingerprint) = |
| 173 | init_rtc(false, server_addr, server_dtls, Peer::Right, &mut timing)?; |
| 174 | |
| 175 | // If SNAP, generate local SCTP INIT for out-of-band exchange |
| 176 | let snap = snap_init_data(use_snap); |
| 177 | let local_sctp_init = snap.as_ref().map(|(init, _)| init.clone()); |
| 178 | |
| 179 | // Send server's credentials to client |
| 180 | server_tx |
| 181 | .send(Message::Credentials { |
| 182 | ice_ufrag: local_creds.ufrag.clone(), |
| 183 | ice_pwd: local_creds.pass.clone(), |
| 184 | dtls_fingerprint: local_fingerprint, |
| 185 | sctp_init: local_sctp_init, |
| 186 | }) |
| 187 | .expect("Failed to send server credentials"); |
| 188 | |
| 189 | // Wait for client's credentials |
| 190 | let (remote_ice_ufrag, remote_ice_pwd, remote_fingerprint, remote_sctp_init) = |
| 191 | match server_rx.recv_timeout(Duration::from_secs(5)) { |
| 192 | Ok(Message::Credentials { |
| 193 | ice_ufrag, |
| 194 | ice_pwd, |
| 195 | dtls_fingerprint, |
| 196 | sctp_init, |
| 197 | }) => { |
| 198 | timing.got_offer = Some(Instant::now()); |
| 199 | (ice_ufrag, ice_pwd, dtls_fingerprint, sctp_init) |
| 200 | } |
| 201 | Ok(_) => panic!("Server expected Credentials, got something else"), |
| 202 | Err(e) => panic!("Server failed to receive credentials: {:?}", e), |
| 203 | }; |
| 204 | |
| 205 | // Configure with remote credentials (is_client = false) |
| 206 | configure_rtc( |
| 207 | &mut rtc, |
| 208 | false, |
| 209 | client_addr, |
| 210 | remote_ice_ufrag, |
| 211 | remote_ice_pwd, |
| 212 | remote_fingerprint, |
| 213 | snap.map(|(_, d)| d), |
| 214 | remote_sctp_init, |
| 215 | )?; |
| 216 | timing.sent_answer = Some(Instant::now()); |
| 217 | |
| 218 | // Run the event loop with message exchange |
| 219 | run_rtc_loop_with_exchange( |
| 220 | &mut rtc, |
| 221 | &span, |
| 222 | &server_rx, |
| 223 | &server_tx, |
| 224 | &mut timing, |
| 225 | false, |
| 226 | &mut packets, |
| 227 | &server_packets_sent_clone, |
| 228 | )?; |
| 229 | |
| 230 | timing.dtls_protocol_version = rtc.direct_api().dtls_protocol_version(); |
| 231 | |
| 232 | Ok(timing) |
| 233 | })(); |
| 234 | |
| 235 | (packets, result) |
| 236 | }, |
| 237 | ); |
| 238 | |
| 239 | // Spawn client thread |
| 240 | // Returns (packets, Result) so pcap is available even on failure. |
| 241 | let client_handle = thread::spawn( |
| 242 | move || -> (Vec<PcapPacket>, Result<TimingReport, RtcError>) { |
| 243 | let span = info_span!("CLIENT"); |
| 244 | let _guard = span.enter(); |
| 245 | let mut timing = TimingReport::new(); |
| 246 | let mut packets = Vec::new(); |
| 247 | |
| 248 | let result = (|| -> Result<TimingReport, RtcError> { |
| 249 | // Initialize client (is_client = true) |
| 250 | let (mut rtc, local_creds, local_fingerprint) = |
| 251 | init_rtc(true, client_addr, client_dtls, Peer::Left, &mut timing)?; |
| 252 | |
| 253 | // If SNAP, generate local SCTP INIT for out-of-band exchange |
| 254 | let snap = snap_init_data(use_snap); |
| 255 | let local_sctp_init = snap.as_ref().map(|(init, _)| init.clone()); |
| 256 | |
| 257 | // Wait for server's credentials first |
| 258 | let (remote_ice_ufrag, remote_ice_pwd, remote_fingerprint, remote_sctp_init) = |
| 259 | match client_rx.recv_timeout(Duration::from_secs(5)) { |
| 260 | Ok(Message::Credentials { |
| 261 | ice_ufrag, |
| 262 | ice_pwd, |
| 263 | dtls_fingerprint, |
| 264 | sctp_init, |
| 265 | }) => (ice_ufrag, ice_pwd, dtls_fingerprint, sctp_init), |
| 266 | Ok(_) => panic!("Client expected Credentials, got something else"), |
| 267 | Err(e) => panic!("Client failed to receive server credentials: {:?}", e), |
| 268 | }; |
| 269 | |
| 270 | // Send client's credentials to server |
| 271 | client_tx |
| 272 | .send(Message::Credentials { |
| 273 | ice_ufrag: local_creds.ufrag.clone(), |
| 274 | ice_pwd: local_creds.pass.clone(), |
| 275 | dtls_fingerprint: local_fingerprint, |
| 276 | sctp_init: local_sctp_init, |
| 277 | }) |
| 278 | .expect("Failed to send client credentials"); |
| 279 | timing.sent_offer = Some(Instant::now()); |
| 280 | |
| 281 | // Configure with remote credentials (is_client = true) |
| 282 | configure_rtc( |
| 283 | &mut rtc, |
| 284 | true, |
| 285 | server_addr, |
| 286 | remote_ice_ufrag, |
| 287 | remote_ice_pwd, |
| 288 | remote_fingerprint, |
| 289 | snap.map(|(_, d)| d), |
| 290 | remote_sctp_init, |
| 291 | )?; |
| 292 | timing.got_answer = Some(Instant::now()); |
| 293 | |
| 294 | // Run the event loop with message exchange |
| 295 | run_rtc_loop_with_exchange( |
| 296 | &mut rtc, |
| 297 | &span, |
| 298 | &client_rx, |
| 299 | &client_tx, |
| 300 | &mut timing, |
| 301 | true, |
| 302 | &mut packets, |
| 303 | &client_packets_sent_clone, |
| 304 | )?; |
| 305 | |
| 306 | timing.dtls_protocol_version = rtc.direct_api().dtls_protocol_version(); |
| 307 | |
| 308 | Ok(timing) |
| 309 | })(); |
| 310 | |
| 311 | (packets, result) |
| 312 | }, |
| 313 | ); |
| 314 | |
| 315 | // Wait for both threads to complete |
| 316 | let (server_packets, server_result) = server_handle.join().expect("Server thread panicked"); |
| 317 | let (client_packets, client_result) = client_handle.join().expect("Client thread panicked"); |
| 318 | |
| 319 | // Save pcap files BEFORE checking errors so we capture failing handshakes |
| 320 | if SAVE_PCAP { |
| 321 | let pcap_dir = std::path::Path::new("target/pcap"); |
| 322 | std::fs::create_dir_all(pcap_dir).expect("Failed to create target/pcap directory"); |
| 323 | |
| 324 | let client_path = pcap_dir.join(format!("{test_name}_client.pcap")); |
| 325 | let server_path = pcap_dir.join(format!("{test_name}_server.pcap")); |
| 326 | |
| 327 | write_pcap(&client_path, &client_packets).expect("Failed to write client pcap"); |
| 328 | write_pcap(&server_path, &server_packets).expect("Failed to write server pcap"); |
| 329 | |
| 330 | println!(" PCAP saved: {}", client_path.display()); |
| 331 | println!(" PCAP saved: {}", server_path.display()); |
| 332 | } |
| 333 | |
| 334 | let server_timing = server_result.expect("Server returned error"); |
| 335 | let client_timing = client_result.expect("Client returned error"); |
| 336 | |
| 337 | let total_time = test_start.elapsed(); |
| 338 | let variant = if use_snap { "SNAP" } else { "standard" }; |
| 339 | |
| 340 | client_timing.print(&format!("CLIENT ({})", variant)); |
| 341 | server_timing.print(&format!("SERVER ({})", variant)); |
| 342 | |
| 343 | println!( |
| 344 | "\n=== Total Test Time ({}): {:.3}ms ===", |
| 345 | variant, |
| 346 | total_time.as_secs_f64() * 1000.0 |
| 347 | ); |
| 348 | |
| 349 | let client_sent = client_packets_sent.load(Ordering::SeqCst); |
| 350 | let server_sent = server_packets_sent.load(Ordering::SeqCst); |
| 351 | let total_packets = client_sent + server_sent; |
| 352 | println!("\n=== Packet Counts ({}) ===", variant); |
| 353 | println!(" Client packets sent: {}", client_sent); |
| 354 | println!(" Server packets sent: {}", server_sent); |
| 355 | println!(" Total packets: {}", total_packets); |
| 356 | |
| 357 | if use_snap { |
| 358 | // SNAP skips the 4-way SCTP handshake, so it must use strictly fewer |
| 359 | // packets than a standard connection. |
| 360 | assert!( |
| 361 | total_packets < MAX_SNAP_PACKETS, |
| 362 | "SNAP should use fewer packets, got {total_packets}" |
| 363 | ); |
| 364 | } |
| 365 | |
| 366 | // Verify the exchange happened |
| 367 | assert!( |
| 368 | client_timing.sent_data.is_some(), |
| 369 | "Client should have sent data" |
| 370 | ); |
| 371 | assert!( |
| 372 | client_timing.received_data.is_some(), |
| 373 | "Client should have received reply" |
| 374 | ); |
| 375 | assert!( |
| 376 | server_timing.received_data.is_some(), |
| 377 | "Server should have received data" |
| 378 | ); |
| 379 | assert!( |
| 380 | server_timing.sent_data.is_some(), |
| 381 | "Server should have sent reply" |
| 382 | ); |
| 383 | |
| 384 | // Verify the negotiated DTLS protocol version matches the expected outcome |
| 385 | // for the requested (client, server) version combination. DTLS 1.3 and Auto |
| 386 | // variants are tested only with dimpl. |
| 387 | let expected = match (client_dtls, server_dtls) { |
| 388 | (DtlsVersion::Dtls12, _) | (_, DtlsVersion::Dtls12) => ProtocolVersion::DTLS1_2, |
| 389 | (DtlsVersion::Dtls13, _) | (_, DtlsVersion::Dtls13) => ProtocolVersion::DTLS1_3, |
| 390 | (DtlsVersion::Auto, DtlsVersion::Auto) => ProtocolVersion::DTLS1_3, |
| 391 | _ => unreachable!("unexpected DTLS version combo: {client_dtls:?}/{server_dtls:?}"), |
| 392 | }; |
| 393 | assert_eq!( |
| 394 | client_timing.dtls_protocol_version, |
| 395 | Some(expected), |
| 396 | "Client negotiated DTLS version mismatch" |
| 397 | ); |
| 398 | assert_eq!( |
| 399 | server_timing.dtls_protocol_version, |
| 400 | Some(expected), |
| 401 | "Server negotiated DTLS version mismatch" |
| 402 | ); |
| 403 | |
| 404 | Ok(()) |
| 405 | } |
| 406 | |
| 407 | /// Initialize an Rtc instance configured for client or server role. |
| 408 | /// |
| 409 | /// Returns the Rtc instance and the local ICE credentials/DTLS fingerprint for exchange. |
| 410 | fn init_rtc( |
| 411 | is_client: bool, |
| 412 | local_addr: SocketAddr, |
| 413 | dtls_version: DtlsVersion, |
| 414 | peer: Peer, |
| 415 | timing: &mut TimingReport, |
| 416 | ) -> Result<(Rtc, IceCreds, String), RtcError> { |
| 417 | let ice_creds = IceCreds::new(); |
| 418 | |
| 419 | let mut rtc_config = RtcConfig::new() |
| 420 | .set_local_ice_credentials(ice_creds.clone()) |
| 421 | .set_dtls_version(dtls_version); |
| 422 | if !is_client { |
| 423 | rtc_config = rtc_config.set_ice_lite(true); |
| 424 | } |
| 425 | if let Some(crypto) = peer.crypto_provider() { |
| 426 | rtc_config = rtc_config.set_crypto_provider(crypto); |
| 427 | } |
| 428 | let mut rtc = rtc_config.build(Instant::now()); |
| 429 | timing.rtc_built = Some(Instant::now()); |
| 430 | |
| 431 | let fingerprint = rtc.direct_api().local_dtls_fingerprint().to_string(); |
| 432 | |
| 433 | let local_candidate = Candidate::host(local_addr, "udp")?; |
| 434 | rtc.add_local_candidate(local_candidate); |
| 435 | |
| 436 | Ok((rtc, ice_creds, fingerprint)) |
| 437 | } |
| 438 | |
| 439 | /// Configure the Rtc instance with remote credentials and start DTLS/SCTP. |
| 440 | /// |
| 441 | /// If `local_init_data` and `remote_sctp_init` are both provided, SNAP is used |
| 442 | /// to skip the 4-way SCTP handshake. |
| 443 | /// |
| 444 | /// `local_init_data` is expected to already contain the local INIT chunk from |
| 445 | /// `local_init_chunk()`. This function adds the remote INIT chunk and then calls |
| 446 | /// `start_sctp_with_snap()`. If either side of that exchange is missing, it |
| 447 | /// falls back to the normal `start_sctp()` path. |
| 448 | fn configure_rtc( |
| 449 | rtc: &mut Rtc, |
| 450 | is_client: bool, |
| 451 | remote_addr: SocketAddr, |
| 452 | remote_ice_ufrag: String, |
| 453 | remote_ice_pwd: String, |
| 454 | remote_fingerprint: String, |
| 455 | local_init_data: Option<str0m::channel::SctpInitData>, |
| 456 | remote_sctp_init: Option<Vec<u8>>, |
| 457 | ) -> Result<(), RtcError> { |
| 458 | let remote_candidate = Candidate::host(remote_addr, "udp")?; |
| 459 | rtc.add_remote_candidate(remote_candidate); |
| 460 | |
| 461 | // Build SctpInitData with remote INIT if both sides provided SNAP data |
| 462 | let sctp_init_data = match (local_init_data, remote_sctp_init) { |
| 463 | (Some(mut data), Some(remote_init)) => { |
| 464 | data.set_remote_init_chunk(remote_init); |
| 465 | Some(data) |
| 466 | } |
| 467 | _ => None, |
| 468 | }; |
| 469 | |
| 470 | { |
| 471 | let mut direct_api = rtc.direct_api(); |
| 472 | |
| 473 | direct_api.set_ice_lite(!is_client); |
| 474 | direct_api.set_ice_controlling(is_client); |
| 475 | |
| 476 | direct_api.set_remote_ice_credentials(IceCreds { |
| 477 | ufrag: remote_ice_ufrag, |
| 478 | pass: remote_ice_pwd, |
| 479 | }); |
| 480 | |
| 481 | let fingerprint: Fingerprint = remote_fingerprint |
| 482 | .parse() |
| 483 | .expect("Failed to parse remote fingerprint"); |
| 484 | direct_api.set_remote_fingerprint(fingerprint); |
| 485 | |
| 486 | direct_api.start_dtls(is_client)?; |
| 487 | |
| 488 | if let Some(sctp_init_data) = sctp_init_data { |
| 489 | direct_api.start_sctp_with_snap(is_client, sctp_init_data)?; |
| 490 | } else { |
| 491 | direct_api.start_sctp(is_client); |
| 492 | } |
| 493 | |
| 494 | direct_api.create_data_channel(ChannelConfig { |
| 495 | label: "test-channel".into(), |
| 496 | negotiated: Some(DATA_CHANNEL_ID), |
| 497 | ordered: true, |
| 498 | reliability: Reliability::Reliable, |
| 499 | protocol: "".into(), |
| 500 | }); |
| 501 | } |
| 502 | |
| 503 | rtc.handle_input(Input::Timeout(Instant::now()))?; |
| 504 | |
| 505 | Ok(()) |
| 506 | } |
| 507 | |
| 508 | /// Messages exchanged between client and server threads. |
| 509 | #[derive(Debug)] |
| 510 | enum Message { |
| 511 | /// ICE, DTLS, and optionally SCTP credentials exchange |
| 512 | Credentials { |
| 513 | ice_ufrag: String, |
| 514 | ice_pwd: String, |
| 515 | dtls_fingerprint: String, |
| 516 | /// SCTP INIT chunk for SNAP (`None` when not using SNAP) |
| 517 | sctp_init: Option<Vec<u8>>, |
| 518 | }, |
| 519 | /// RTP/DTLS/SCTP packet |
| 520 | Packet { |
| 521 | proto: Protocol, |
| 522 | source: SocketAddr, |
| 523 | destination: SocketAddr, |
| 524 | contents: Vec<u8>, |
| 525 | }, |
| 526 | /// Signal to exit (sent by client to server) |
| 527 | Exit, |
| 528 | } |
| 529 | |
| 530 | /// Timing report for major events |
| 531 | #[derive(Debug, Default)] |
| 532 | struct TimingReport { |
| 533 | start: Option<Instant>, |
| 534 | rtc_built: Option<Instant>, |
| 535 | sent_offer: Option<Instant>, |
| 536 | got_offer: Option<Instant>, |
| 537 | sent_answer: Option<Instant>, |
| 538 | got_answer: Option<Instant>, |
| 539 | ice_checking: Option<Instant>, |
| 540 | ice_completed: Option<Instant>, |
| 541 | channel_open: Option<Instant>, |
| 542 | sent_data: Option<Instant>, |
| 543 | received_data: Option<Instant>, |
| 544 | dtls_protocol_version: Option<ProtocolVersion>, |
| 545 | } |
| 546 | |
| 547 | impl TimingReport { |
| 548 | fn new() -> Self { |
| 549 | Self { |
| 550 | start: Some(Instant::now()), |
| 551 | ..Default::default() |
| 552 | } |
| 553 | } |
| 554 | |
| 555 | fn print(&self, name: &str) { |
| 556 | let start = self.start.unwrap(); |
| 557 | println!("\n=== {} Timing Report ===", name); |
| 558 | if let Some(t) = self.rtc_built { |
| 559 | println!( |
| 560 | " Rtc built: {:>8.3}ms", |
| 561 | (t - start).as_secs_f64() * 1000.0 |
| 562 | ); |
| 563 | } |
| 564 | if let Some(t) = self.sent_offer { |
| 565 | println!( |
| 566 | " Sent offer: {:>8.3}ms", |
| 567 | (t - start).as_secs_f64() * 1000.0 |
| 568 | ); |
| 569 | } |
| 570 | if let Some(t) = self.got_offer { |
| 571 | println!( |
| 572 | " Got offer: {:>8.3}ms", |
| 573 | (t - start).as_secs_f64() * 1000.0 |
| 574 | ); |
| 575 | } |
| 576 | if let Some(t) = self.sent_answer { |
| 577 | println!( |
| 578 | " Sent answer: {:>8.3}ms", |
| 579 | (t - start).as_secs_f64() * 1000.0 |
| 580 | ); |
| 581 | } |
| 582 | if let Some(t) = self.got_answer { |
| 583 | println!( |
| 584 | " Got answer: {:>8.3}ms", |
| 585 | (t - start).as_secs_f64() * 1000.0 |
| 586 | ); |
| 587 | } |
| 588 | if let Some(t) = self.ice_checking { |
| 589 | println!( |
| 590 | " ICE Checking: {:>8.3}ms", |
| 591 | (t - start).as_secs_f64() * 1000.0 |
| 592 | ); |
| 593 | } |
| 594 | if let Some(t) = self.ice_completed { |
| 595 | println!( |
| 596 | " ICE Completed: {:>8.3}ms", |
| 597 | (t - start).as_secs_f64() * 1000.0 |
| 598 | ); |
| 599 | } |
| 600 | if let Some(t) = self.channel_open { |
| 601 | println!( |
| 602 | " Channel Open: {:>8.3}ms", |
| 603 | (t - start).as_secs_f64() * 1000.0 |
| 604 | ); |
| 605 | } |
| 606 | if let Some(t) = self.sent_data { |
| 607 | println!( |
| 608 | " Sent data: {:>8.3}ms", |
| 609 | (t - start).as_secs_f64() * 1000.0 |
| 610 | ); |
| 611 | } |
| 612 | if let Some(t) = self.received_data { |
| 613 | println!( |
| 614 | " Received data: {:>8.3}ms", |
| 615 | (t - start).as_secs_f64() * 1000.0 |
| 616 | ); |
| 617 | } |
| 618 | } |
| 619 | } |
| 620 | |
| 621 | /// State for managing message exchange |
| 622 | #[derive(Debug, PartialEq)] |
| 623 | enum DataExchangeState { |
| 624 | WaitingForChannelOpen, |
| 625 | ChannelOpen, |
| 626 | SentMessage, |
| 627 | Complete, |
| 628 | } |
| 629 | |
| 630 | /// Run the Rtc event loop with message exchange capability |
| 631 | fn run_rtc_loop_with_exchange( |
| 632 | rtc: &mut Rtc, |
| 633 | span: &Span, |
| 634 | incoming: &Receiver<Message>, |
| 635 | outgoing: &Sender<Message>, |
| 636 | timing: &mut TimingReport, |
| 637 | is_client: bool, |
| 638 | packets: &mut Vec<PcapPacket>, |
| 639 | packets_sent: &AtomicUsize, |
| 640 | ) -> Result<(), RtcError> { |
| 641 | let mut state = DataExchangeState::WaitingForChannelOpen; |
| 642 | let mut channel_id: Option<ChannelId> = None; |
| 643 | let role = if is_client { "CLIENT" } else { "SERVER" }; |
| 644 | |
| 645 | loop { |
| 646 | if state == DataExchangeState::Complete { |
| 647 | break; |
| 648 | } |
| 649 | |
| 650 | if timing.start.unwrap().elapsed() > Duration::from_secs(10) { |
| 651 | println!("[{}] Overall timeout reached", role); |
| 652 | break; |
| 653 | } |
| 654 | |
| 655 | let timeout = loop { |
| 656 | match span.in_scope(|| rtc.poll_output())? { |
| 657 | Output::Timeout(t) => break t, |
| 658 | Output::Transmit(t) => { |
| 659 | let data = t.contents.to_vec(); |
| 660 | packets_sent.fetch_add(1, Ordering::SeqCst); |
| 661 | if SAVE_PCAP { |
| 662 | packets.push(PcapPacket { |
| 663 | src: t.source, |
| 664 | dst: t.destination, |
| 665 | data: data.clone(), |
| 666 | }); |
| 667 | } |
| 668 | // Send packet to other peer |
| 669 | let _ = outgoing.send(Message::Packet { |
| 670 | proto: t.proto, |
| 671 | source: t.source, |
| 672 | destination: t.destination, |
| 673 | contents: data, |
| 674 | }); |
| 675 | } |
| 676 | Output::Event(e) => { |
| 677 | handle_event( |
| 678 | rtc, |
| 679 | &e, |
| 680 | timing, |
| 681 | is_client, |
| 682 | &mut state, |
| 683 | &mut channel_id, |
| 684 | outgoing, |
| 685 | ); |
| 686 | if state == DataExchangeState::Complete { |
| 687 | return Ok(()); |
| 688 | } |
| 689 | } |
| 690 | } |
| 691 | }; |
| 692 | |
| 693 | let now = Instant::now(); |
| 694 | let wait = timeout.saturating_duration_since(now); |
| 695 | println!("[{}] poll_output returned timeout in {:?}", role, wait); |
| 696 | |
| 697 | match incoming.recv_timeout(wait) { |
| 698 | Ok(Message::Packet { |
| 699 | proto, |
| 700 | source, |
| 701 | destination, |
| 702 | contents, |
| 703 | }) => { |
| 704 | println!("[{}] Received packet ({} bytes)", role, contents.len()); |
| 705 | if SAVE_PCAP { |
| 706 | packets.push(PcapPacket { |
| 707 | src: source, |
| 708 | dst: destination, |
| 709 | data: contents.clone(), |
| 710 | }); |
| 711 | } |
| 712 | let receive = Receive { |
| 713 | proto, |
| 714 | source, |
| 715 | destination, |
| 716 | contents: contents.as_slice().try_into()?, |
| 717 | }; |
| 718 | span.in_scope(|| rtc.handle_input(Input::Receive(Instant::now(), receive)))?; |
| 719 | } |
| 720 | Ok(Message::Exit) => { |
| 721 | println!("[{}] Received Exit signal", role); |
| 722 | state = DataExchangeState::Complete; |
| 723 | } |
| 724 | Ok(_) => { |
| 725 | unreachable!("Unexpected message type"); |
| 726 | } |
| 727 | Err(mpsc::RecvTimeoutError::Timeout) => { |
| 728 | println!("[{}] Timeout fired, calling handle_input(Timeout)", role); |
| 729 | span.in_scope(|| rtc.handle_input(Input::Timeout(Instant::now())))?; |
| 730 | } |
| 731 | Err(mpsc::RecvTimeoutError::Disconnected) => { |
| 732 | println!("[{}] Channel disconnected", role); |
| 733 | break; |
| 734 | } |
| 735 | } |
| 736 | } |
| 737 | |
| 738 | Ok(()) |
| 739 | } |
| 740 | |
| 741 | fn handle_event( |
| 742 | rtc: &mut Rtc, |
| 743 | event: &Event, |
| 744 | timing: &mut TimingReport, |
| 745 | is_client: bool, |
| 746 | state: &mut DataExchangeState, |
| 747 | channel_id: &mut Option<ChannelId>, |
| 748 | outgoing: &Sender<Message>, |
| 749 | ) { |
| 750 | match event { |
| 751 | Event::IceConnectionStateChange(ice_state) => match ice_state { |
| 752 | IceConnectionState::Checking => { |
| 753 | if timing.ice_checking.is_none() { |
| 754 | timing.ice_checking = Some(Instant::now()); |
| 755 | } |
| 756 | } |
| 757 | IceConnectionState::Completed => { |
| 758 | timing.ice_completed = Some(Instant::now()); |
| 759 | } |
| 760 | _ => {} |
| 761 | }, |
| 762 | Event::ChannelOpen(cid, label) => { |
| 763 | println!( |
| 764 | "[{}] Channel opened: {:?} - {}", |
| 765 | if is_client { "CLIENT" } else { "SERVER" }, |
| 766 | cid, |
| 767 | label |
| 768 | ); |
| 769 | timing.channel_open = Some(Instant::now()); |
| 770 | *channel_id = Some(*cid); |
| 771 | *state = DataExchangeState::ChannelOpen; |
| 772 | |
| 773 | // Client sends first message |
| 774 | if is_client { |
| 775 | if let Some(mut chan) = rtc.channel(*cid) { |
| 776 | chan.write(true, b"sixseven").expect("Failed to write"); |
| 777 | println!("[CLIENT] Sent 'sixseven'"); |
| 778 | timing.sent_data = Some(Instant::now()); |
| 779 | *state = DataExchangeState::SentMessage; |
| 780 | } |
| 781 | } |
| 782 | } |
| 783 | Event::ChannelData(data) => { |
| 784 | let msg = String::from_utf8_lossy(&data.data); |
| 785 | println!( |
| 786 | "[{}] Received data: '{}'", |
| 787 | if is_client { "CLIENT" } else { "SERVER" }, |
| 788 | msg |
| 789 | ); |
| 790 | if is_client { |
| 791 | if msg == "sevenofnine" { |
| 792 | println!("[CLIENT] Got reply 'sevenofnine' - sending Exit and completing"); |
| 793 | timing.received_data = Some(Instant::now()); |
| 794 | let _ = outgoing.send(Message::Exit); |
| 795 | *state = DataExchangeState::Complete; |
| 796 | } |
| 797 | } else if msg == "sixseven" { |
| 798 | timing.received_data = Some(Instant::now()); |
| 799 | let cid = data.id; |
| 800 | if let Some(mut chan) = rtc.channel(cid) { |
| 801 | chan.write(true, b"sevenofnine").expect("Failed to write"); |
| 802 | println!("[SERVER] Sent reply 'sevenofnine'"); |
| 803 | timing.sent_data = Some(Instant::now()); |
| 804 | *state = DataExchangeState::SentMessage; |
| 805 | } |
| 806 | } |
| 807 | } |
| 808 | _ => {} |
| 809 | } |
| 810 | } |
| 811 | |
| 812 | // --- PCAP support --- |
| 813 | |
| 814 | fn dtls_version_short(v: DtlsVersion) -> &'static str { |
| 815 | match v { |
| 816 | DtlsVersion::Auto => "auto", |
| 817 | DtlsVersion::Dtls12 => "12", |
| 818 | DtlsVersion::Dtls13 => "13", |
| 819 | _ => "unknown", |
| 820 | } |
| 821 | } |
| 822 | |
| 823 | /// A captured packet for pcap output. |
| 824 | struct PcapPacket { |
| 825 | src: SocketAddr, |
| 826 | dst: SocketAddr, |
| 827 | data: Vec<u8>, |
| 828 | } |
| 829 | |
| 830 | /// Write packets to a pcap file using the standard pcap format. |
| 831 | /// Uses raw IPv4 link type so Wireshark can dissect the UDP/DTLS layers. |
| 832 | fn write_pcap(path: &std::path::Path, packets: &[PcapPacket]) -> std::io::Result<()> { |
| 833 | use std::io::Write; |
| 834 | |
| 835 | let mut f = std::fs::File::create(path)?; |
| 836 | |
| 837 | // Global header (24 bytes) |
| 838 | // magic_number, version_major, version_minor, thiszone, sigfigs, snaplen, network |
| 839 | f.write_all(&0xa1b2c3d4u32.to_le_bytes())?; // magic |
| 840 | f.write_all(&2u16.to_le_bytes())?; // version major |
| 841 | f.write_all(&4u16.to_le_bytes())?; // version minor |
| 842 | f.write_all(&0i32.to_le_bytes())?; // thiszone |
| 843 | f.write_all(&0u32.to_le_bytes())?; // sigfigs |
| 844 | f.write_all(&65535u32.to_le_bytes())?; // snaplen |
| 845 | f.write_all(&228u32.to_le_bytes())?; // LINKTYPE_IPV4 (228 = raw IPv4) |
| 846 | |
| 847 | for (i, pkt) in packets.iter().enumerate() { |
| 848 | // Build a minimal IPv4 + UDP frame around the payload |
| 849 | let udp_len = 8 + pkt.data.len(); |
| 850 | let ip_total_len = 20 + udp_len; |
| 851 | |
| 852 | // IPv4 header (20 bytes, no options) |
| 853 | let mut ip_header = [0u8; 20]; |
| 854 | ip_header[0] = 0x45; // version=4, IHL=5 |
| 855 | ip_header[1] = 0; // DSCP/ECN |
| 856 | ip_header[2..4].copy_from_slice(&(ip_total_len as u16).to_be_bytes()); |
| 857 | ip_header[4..6].copy_from_slice(&(i as u16).to_be_bytes()); // identification |
| 858 | ip_header[8] = 64; // TTL |
| 859 | ip_header[9] = 17; // protocol = UDP |
| 860 | // checksum left as 0 (Wireshark will flag but still parse) |
| 861 | match pkt.src { |
| 862 | SocketAddr::V4(a) => ip_header[12..16].copy_from_slice(&a.ip().octets()), |
| 863 | _ => {} |
| 864 | } |
| 865 | match pkt.dst { |
| 866 | SocketAddr::V4(a) => ip_header[16..20].copy_from_slice(&a.ip().octets()), |
| 867 | _ => {} |
| 868 | } |
| 869 | |
| 870 | // UDP header (8 bytes) |
| 871 | let mut udp_header = [0u8; 8]; |
| 872 | udp_header[0..2].copy_from_slice(&pkt.src.port().to_be_bytes()); |
| 873 | udp_header[2..4].copy_from_slice(&pkt.dst.port().to_be_bytes()); |
| 874 | udp_header[4..6].copy_from_slice(&(udp_len as u16).to_be_bytes()); |
| 875 | // checksum left as 0 |
| 876 | |
| 877 | let frame_len = ip_total_len as u32; |
| 878 | |
| 879 | // Packet record header (16 bytes) |
| 880 | // Use packet index as fake timestamp (1ms apart) |
| 881 | let ts_sec = i as u32; |
| 882 | let ts_usec = 0u32; |
| 883 | f.write_all(&ts_sec.to_le_bytes())?; |
| 884 | f.write_all(&ts_usec.to_le_bytes())?; |
| 885 | f.write_all(&frame_len.to_le_bytes())?; // incl_len |
| 886 | f.write_all(&frame_len.to_le_bytes())?; // orig_len |
| 887 | |
| 888 | // Frame data |
| 889 | f.write_all(&ip_header)?; |
| 890 | f.write_all(&udp_header)?; |
| 891 | f.write_all(&pkt.data)?; |
| 892 | } |
| 893 | |
| 894 | Ok(()) |
| 895 | } |