1 #[cfg(test)] 2 mod sctp_transport_test; 3 4 pub mod sctp_transport_capabilities; 5 pub mod sctp_transport_state; 6 7 use sctp_transport_state::RTCSctpTransportState; 8 use std::collections::HashSet; 9 10 use crate::api::setting_engine::SettingEngine; 11 use crate::data_channel::data_channel_state::RTCDataChannelState; 12 use crate::data_channel::RTCDataChannel; 13 use crate::dtls_transport::dtls_role::DTLSRole; 14 use crate::dtls_transport::*; 15 use crate::error::*; 16 use crate::sctp_transport::sctp_transport_capabilities::SCTPTransportCapabilities; 17 use crate::stats::stats_collector::StatsCollector; 18 use crate::stats::StatsReportType::{PeerConnection, SCTPTransport}; 19 use crate::stats::{ICETransportStats, PeerConnectionStats}; 20 21 use data::message::message_channel_open::ChannelType; 22 use sctp::association::Association; 23 24 use crate::data_channel::data_channel_parameters::DataChannelParameters; 25 26 use arc_swap::ArcSwapOption; 27 use data::data_channel::DataChannel; 28 use std::collections::HashMap; 29 use std::future::Future; 30 use std::pin::Pin; 31 use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU8, Ordering}; 32 use std::sync::Arc; 33 use tokio::sync::{Mutex, Notify}; 34 use util::Conn; 35 36 const SCTP_MAX_CHANNELS: u16 = u16::MAX; 37 38 pub type OnDataChannelHdlrFn = Box< 39 dyn (FnMut(Arc<RTCDataChannel>) -> Pin<Box<dyn Future<Output = ()> + Send + 'static>>) 40 + Send 41 + Sync, 42 >; 43 44 pub type OnDataChannelOpenedHdlrFn = Box< 45 dyn (FnMut(Arc<RTCDataChannel>) -> Pin<Box<dyn Future<Output = ()> + Send + 'static>>) 46 + Send 47 + Sync, 48 >; 49 50 struct AcceptDataChannelParams { 51 notify_rx: Arc<Notify>, 52 sctp_association: Arc<Association>, 53 data_channels: Arc<Mutex<Vec<Arc<RTCDataChannel>>>>, 54 on_error_handler: Arc<ArcSwapOption<Mutex<OnErrorHdlrFn>>>, 55 on_data_channel_handler: Arc<ArcSwapOption<Mutex<OnDataChannelHdlrFn>>>, 56 on_data_channel_opened_handler: Arc<ArcSwapOption<Mutex<OnDataChannelOpenedHdlrFn>>>, 57 data_channels_opened: Arc<AtomicU32>, 58 data_channels_accepted: Arc<AtomicU32>, 59 setting_engine: Arc<SettingEngine>, 60 } 61 62 /// SCTPTransport provides details about the SCTP transport. 63 #[derive(Default)] 64 pub struct RTCSctpTransport { 65 pub(crate) dtls_transport: Arc<RTCDtlsTransport>, 66 67 // State represents the current state of the SCTP transport. 68 state: AtomicU8, // RTCSctpTransportState 69 70 // SCTPTransportState doesn't have an enum to distinguish between New/Connecting 71 // so we need a dedicated field 72 is_started: AtomicBool, 73 74 // max_message_size represents the maximum size of data that can be passed to 75 // DataChannel's send() method. 76 max_message_size: usize, 77 78 // max_channels represents the maximum amount of DataChannel's that can 79 // be used simultaneously. 80 max_channels: u16, 81 82 sctp_association: Mutex<Option<Arc<Association>>>, 83 84 on_error_handler: Arc<ArcSwapOption<Mutex<OnErrorHdlrFn>>>, 85 on_data_channel_handler: Arc<ArcSwapOption<Mutex<OnDataChannelHdlrFn>>>, 86 on_data_channel_opened_handler: Arc<ArcSwapOption<Mutex<OnDataChannelOpenedHdlrFn>>>, 87 88 // DataChannels 89 pub(crate) data_channels: Arc<Mutex<Vec<Arc<RTCDataChannel>>>>, 90 pub(crate) data_channels_opened: Arc<AtomicU32>, 91 pub(crate) data_channels_requested: Arc<AtomicU32>, 92 data_channels_accepted: Arc<AtomicU32>, 93 94 notify_tx: Arc<Notify>, 95 96 setting_engine: Arc<SettingEngine>, 97 } 98 99 impl RTCSctpTransport { new( dtls_transport: Arc<RTCDtlsTransport>, setting_engine: Arc<SettingEngine>, ) -> Self100 pub(crate) fn new( 101 dtls_transport: Arc<RTCDtlsTransport>, 102 setting_engine: Arc<SettingEngine>, 103 ) -> Self { 104 RTCSctpTransport { 105 dtls_transport, 106 state: AtomicU8::new(RTCSctpTransportState::Connecting as u8), 107 is_started: AtomicBool::new(false), 108 max_message_size: RTCSctpTransport::calc_message_size(65536, 65536), 109 max_channels: SCTP_MAX_CHANNELS, 110 sctp_association: Mutex::new(None), 111 on_error_handler: Arc::new(ArcSwapOption::empty()), 112 on_data_channel_handler: Arc::new(ArcSwapOption::empty()), 113 on_data_channel_opened_handler: Arc::new(ArcSwapOption::empty()), 114 115 data_channels: Arc::new(Mutex::new(vec![])), 116 data_channels_opened: Arc::new(AtomicU32::new(0)), 117 data_channels_requested: Arc::new(AtomicU32::new(0)), 118 data_channels_accepted: Arc::new(AtomicU32::new(0)), 119 120 notify_tx: Arc::new(Notify::new()), 121 122 setting_engine, 123 } 124 } 125 126 /// transport returns the DTLSTransport instance the SCTPTransport is sending over. transport(&self) -> Arc<RTCDtlsTransport>127 pub fn transport(&self) -> Arc<RTCDtlsTransport> { 128 Arc::clone(&self.dtls_transport) 129 } 130 131 /// get_capabilities returns the SCTPCapabilities of the SCTPTransport. get_capabilities(&self) -> SCTPTransportCapabilities132 pub fn get_capabilities(&self) -> SCTPTransportCapabilities { 133 SCTPTransportCapabilities { 134 max_message_size: 0, 135 } 136 } 137 138 /// Start the SCTPTransport. Since both local and remote parties must mutually 139 /// create an SCTPTransport, SCTP SO (Simultaneous Open) is used to establish 140 /// a connection over SCTP. start(&self, _remote_caps: SCTPTransportCapabilities) -> Result<()>141 pub async fn start(&self, _remote_caps: SCTPTransportCapabilities) -> Result<()> { 142 if self.is_started.load(Ordering::SeqCst) { 143 return Ok(()); 144 } 145 self.is_started.store(true, Ordering::SeqCst); 146 147 let dtls_transport = self.transport(); 148 if let Some(net_conn) = &dtls_transport.conn().await { 149 let sctp_association = loop { 150 tokio::select! { 151 _ = self.notify_tx.notified() => { 152 // It seems like notify_tx is only notified on Stop so perhaps this check 153 // is redundant. 154 // TODO: Consider renaming notify_tx to shutdown_tx. 155 if self.state.load(Ordering::SeqCst) == RTCSctpTransportState::Closed as u8 { 156 return Err(Error::ErrSCTPTransportDTLS); 157 } 158 }, 159 association = sctp::association::Association::client(sctp::association::Config { 160 net_conn: Arc::clone(net_conn) as Arc<dyn Conn + Send + Sync>, 161 max_receive_buffer_size: 0, 162 max_message_size: 0, 163 name: String::new(), 164 }) => { 165 break Arc::new(association?); 166 } 167 }; 168 }; 169 170 { 171 let mut sa = self.sctp_association.lock().await; 172 *sa = Some(Arc::clone(&sctp_association)); 173 } 174 self.state 175 .store(RTCSctpTransportState::Connected as u8, Ordering::SeqCst); 176 177 let param = AcceptDataChannelParams { 178 notify_rx: self.notify_tx.clone(), 179 sctp_association, 180 data_channels: Arc::clone(&self.data_channels), 181 on_error_handler: Arc::clone(&self.on_error_handler), 182 on_data_channel_handler: Arc::clone(&self.on_data_channel_handler), 183 on_data_channel_opened_handler: Arc::clone(&self.on_data_channel_opened_handler), 184 data_channels_opened: Arc::clone(&self.data_channels_opened), 185 data_channels_accepted: Arc::clone(&self.data_channels_accepted), 186 setting_engine: Arc::clone(&self.setting_engine), 187 }; 188 tokio::spawn(async move { 189 RTCSctpTransport::accept_data_channels(param).await; 190 }); 191 192 Ok(()) 193 } else { 194 Err(Error::ErrSCTPTransportDTLS) 195 } 196 } 197 198 /// Stop stops the SCTPTransport stop(&self) -> Result<()>199 pub async fn stop(&self) -> Result<()> { 200 { 201 let mut sctp_association = self.sctp_association.lock().await; 202 if let Some(sa) = sctp_association.take() { 203 sa.close().await?; 204 } 205 } 206 207 self.state 208 .store(RTCSctpTransportState::Closed as u8, Ordering::SeqCst); 209 210 self.notify_tx.notify_waiters(); 211 212 Ok(()) 213 } 214 accept_data_channels(param: AcceptDataChannelParams)215 async fn accept_data_channels(param: AcceptDataChannelParams) { 216 let dcs = param.data_channels.lock().await; 217 let mut existing_data_channels = Vec::new(); 218 for dc in dcs.iter() { 219 if let Some(dc) = dc.data_channel.lock().await.clone() { 220 existing_data_channels.push(dc); 221 } 222 } 223 drop(dcs); 224 225 loop { 226 let dc = tokio::select! { 227 _ = param.notify_rx.notified() => break, 228 result = DataChannel::accept( 229 ¶m.sctp_association, 230 data::data_channel::Config::default(), 231 &existing_data_channels, 232 ) => { 233 match result { 234 Ok(dc) => dc, 235 Err(err) => { 236 if data::Error::ErrStreamClosed == err { 237 log::error!("Failed to accept data channel: {}", err); 238 if let Some(handler) = &*param.on_error_handler.load() { 239 let mut f = handler.lock().await; 240 f(err.into()).await; 241 } 242 } 243 break; 244 } 245 } 246 } 247 }; 248 249 let mut max_retransmits = 0; 250 let mut max_packet_lifetime = 0; 251 let val = dc.config.reliability_parameter as u16; 252 let ordered; 253 254 match dc.config.channel_type { 255 ChannelType::Reliable => { 256 ordered = true; 257 } 258 ChannelType::ReliableUnordered => { 259 ordered = false; 260 } 261 ChannelType::PartialReliableRexmit => { 262 ordered = true; 263 max_retransmits = val; 264 } 265 ChannelType::PartialReliableRexmitUnordered => { 266 ordered = false; 267 max_retransmits = val; 268 } 269 ChannelType::PartialReliableTimed => { 270 ordered = true; 271 max_packet_lifetime = val; 272 } 273 ChannelType::PartialReliableTimedUnordered => { 274 ordered = false; 275 max_packet_lifetime = val; 276 } 277 }; 278 279 let negotiated = if dc.config.negotiated { 280 Some(dc.stream_identifier()) 281 } else { 282 None 283 }; 284 let rtc_dc = Arc::new(RTCDataChannel::new( 285 DataChannelParameters { 286 label: dc.config.label.clone(), 287 protocol: dc.config.protocol.clone(), 288 negotiated, 289 ordered, 290 max_packet_life_time: max_packet_lifetime, 291 max_retransmits, 292 }, 293 Arc::clone(¶m.setting_engine), 294 )); 295 296 if let Some(handler) = &*param.on_data_channel_handler.load() { 297 let mut f = handler.lock().await; 298 f(Arc::clone(&rtc_dc)).await; 299 300 param.data_channels_accepted.fetch_add(1, Ordering::SeqCst); 301 302 let mut dcs = param.data_channels.lock().await; 303 dcs.push(Arc::clone(&rtc_dc)); 304 } 305 306 rtc_dc.handle_open(Arc::new(dc)).await; 307 308 if let Some(handler) = &*param.on_data_channel_opened_handler.load() { 309 let mut f = handler.lock().await; 310 f(rtc_dc).await; 311 param.data_channels_opened.fetch_add(1, Ordering::SeqCst); 312 } 313 } 314 } 315 316 /// on_error sets an event handler which is invoked when 317 /// the SCTP connection error occurs. on_error(&self, f: OnErrorHdlrFn)318 pub fn on_error(&self, f: OnErrorHdlrFn) { 319 self.on_error_handler.store(Some(Arc::new(Mutex::new(f)))); 320 } 321 322 /// on_data_channel sets an event handler which is invoked when a data 323 /// channel message arrives from a remote peer. on_data_channel(&self, f: OnDataChannelHdlrFn)324 pub fn on_data_channel(&self, f: OnDataChannelHdlrFn) { 325 self.on_data_channel_handler 326 .store(Some(Arc::new(Mutex::new(f)))); 327 } 328 329 /// on_data_channel_opened sets an event handler which is invoked when a data 330 /// channel is opened on_data_channel_opened(&self, f: OnDataChannelOpenedHdlrFn)331 pub fn on_data_channel_opened(&self, f: OnDataChannelOpenedHdlrFn) { 332 self.on_data_channel_opened_handler 333 .store(Some(Arc::new(Mutex::new(f)))); 334 } 335 calc_message_size(remote_max_message_size: usize, can_send_size: usize) -> usize336 fn calc_message_size(remote_max_message_size: usize, can_send_size: usize) -> usize { 337 if remote_max_message_size == 0 && can_send_size == 0 { 338 usize::MAX 339 } else if remote_max_message_size == 0 { 340 can_send_size 341 } else if can_send_size == 0 || can_send_size > remote_max_message_size { 342 remote_max_message_size 343 } else { 344 can_send_size 345 } 346 } 347 348 /// max_channels is the maximum number of RTCDataChannels that can be open simultaneously. max_channels(&self) -> u16349 pub fn max_channels(&self) -> u16 { 350 if self.max_channels == 0 { 351 SCTP_MAX_CHANNELS 352 } else { 353 self.max_channels 354 } 355 } 356 357 /// state returns the current state of the SCTPTransport state(&self) -> RTCSctpTransportState358 pub fn state(&self) -> RTCSctpTransportState { 359 self.state.load(Ordering::SeqCst).into() 360 } 361 collect_stats( &self, collector: &StatsCollector, peer_connection_id: String, )362 pub(crate) async fn collect_stats( 363 &self, 364 collector: &StatsCollector, 365 peer_connection_id: String, 366 ) { 367 let dtls_transport = self.transport(); 368 369 // TODO: should this be collected? 370 dtls_transport.collect_stats(collector).await; 371 372 // data channels 373 let mut data_channels_closed = 0; 374 let data_channels = self.data_channels.lock().await; 375 for data_channel in &*data_channels { 376 match data_channel.ready_state() { 377 RTCDataChannelState::Connecting => (), 378 RTCDataChannelState::Open => (), 379 _ => data_channels_closed += 1, 380 } 381 data_channel.collect_stats(collector).await; 382 } 383 384 let mut reports = HashMap::new(); 385 let peer_connection_stats = 386 PeerConnectionStats::new(self, peer_connection_id.clone(), data_channels_closed); 387 reports.insert(peer_connection_id, PeerConnection(peer_connection_stats)); 388 389 // conn 390 if let Some(agent) = dtls_transport.ice_transport.gatherer.get_agent().await { 391 let stats = ICETransportStats::new("sctp_transport".to_owned(), agent); 392 reports.insert(stats.id.clone(), SCTPTransport(stats)); 393 } 394 395 collector.merge(reports); 396 } 397 generate_and_set_data_channel_id( &self, dtls_role: DTLSRole, ) -> Result<u16>398 pub(crate) async fn generate_and_set_data_channel_id( 399 &self, 400 dtls_role: DTLSRole, 401 ) -> Result<u16> { 402 let mut id = 0u16; 403 if dtls_role != DTLSRole::Client { 404 id += 1; 405 } 406 407 // Create map of ids so we can compare without double-looping each time. 408 let mut ids_map = HashSet::new(); 409 { 410 let data_channels = self.data_channels.lock().await; 411 for dc in &*data_channels { 412 ids_map.insert(dc.id()); 413 } 414 } 415 416 let max = self.max_channels(); 417 while id < max - 1 { 418 if ids_map.contains(&id) { 419 id += 2; 420 } else { 421 return Ok(id); 422 } 423 } 424 425 Err(Error::ErrMaxDataChannelID) 426 } 427 association(&self) -> Option<Arc<Association>>428 pub(crate) async fn association(&self) -> Option<Arc<Association>> { 429 let sctp_association = self.sctp_association.lock().await; 430 sctp_association.clone() 431 } 432 data_channels_accepted(&self) -> u32433 pub(crate) fn data_channels_accepted(&self) -> u32 { 434 self.data_channels_accepted.load(Ordering::SeqCst) 435 } 436 data_channels_opened(&self) -> u32437 pub(crate) fn data_channels_opened(&self) -> u32 { 438 self.data_channels_opened.load(Ordering::SeqCst) 439 } 440 data_channels_requested(&self) -> u32441 pub(crate) fn data_channels_requested(&self) -> u32 { 442 self.data_channels_requested.load(Ordering::SeqCst) 443 } 444 } 445