xref: /webrtc/webrtc/src/sctp_transport/mod.rs (revision 5b79f08a)
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                     &param.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(&param.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