1 use super::*; 2 use std::sync::atomic::{AtomicU32, Ordering}; 3 use std::sync::Arc; 4 use tokio::io::AsyncReadExt; 5 use tokio::io::AsyncWriteExt; 6 7 #[test] 8 fn test_stream_buffered_amount() -> Result<()> { 9 let s = Stream::default(); 10 11 assert_eq!(0, s.buffered_amount()); 12 assert_eq!(0, s.buffered_amount_low_threshold()); 13 14 s.buffered_amount.store(8192, Ordering::SeqCst); 15 s.set_buffered_amount_low_threshold(2048); 16 assert_eq!(8192, s.buffered_amount(), "unexpected bufferedAmount"); 17 assert_eq!( 18 2048, 19 s.buffered_amount_low_threshold(), 20 "unexpected threshold" 21 ); 22 23 Ok(()) 24 } 25 26 #[tokio::test] 27 async fn test_stream_amount_on_buffered_amount_low() -> Result<()> { 28 let s = Stream::default(); 29 30 s.buffered_amount.store(4096, Ordering::SeqCst); 31 s.set_buffered_amount_low_threshold(2048); 32 33 let n_cbs = Arc::new(AtomicU32::new(0)); 34 let n_cbs2 = n_cbs.clone(); 35 36 s.on_buffered_amount_low(Box::new(move || { 37 n_cbs2.fetch_add(1, Ordering::SeqCst); 38 Box::pin(async {}) 39 })); 40 41 // Negative value should be ignored (by design) 42 s.on_buffer_released(-32).await; // bufferedAmount = 3072 43 assert_eq!(4096, s.buffered_amount(), "unexpected bufferedAmount"); 44 assert_eq!(0, n_cbs.load(Ordering::SeqCst), "callback count mismatch"); 45 46 // Above to above, no callback 47 s.on_buffer_released(1024).await; // bufferedAmount = 3072 48 assert_eq!(3072, s.buffered_amount(), "unexpected bufferedAmount"); 49 assert_eq!(0, n_cbs.load(Ordering::SeqCst), "callback count mismatch"); 50 51 // Above to equal, callback should be made 52 s.on_buffer_released(1024).await; // bufferedAmount = 2048 53 assert_eq!(2048, s.buffered_amount(), "unexpected bufferedAmount"); 54 assert_eq!(1, n_cbs.load(Ordering::SeqCst), "callback count mismatch"); 55 56 // Eaual to below, no callback 57 s.on_buffer_released(1024).await; // bufferedAmount = 1024 58 assert_eq!(1024, s.buffered_amount(), "unexpected bufferedAmount"); 59 assert_eq!(1, n_cbs.load(Ordering::SeqCst), "callback count mismatch"); 60 61 // Blow to below, no callback 62 s.on_buffer_released(1024).await; // bufferedAmount = 0 63 assert_eq!(0, s.buffered_amount(), "unexpected bufferedAmount"); 64 assert_eq!(1, n_cbs.load(Ordering::SeqCst), "callback count mismatch"); 65 66 // Capped at 0, no callback 67 s.on_buffer_released(1024).await; // bufferedAmount = 0 68 assert_eq!(0, s.buffered_amount(), "unexpected bufferedAmount"); 69 assert_eq!(1, n_cbs.load(Ordering::SeqCst), "callback count mismatch"); 70 71 Ok(()) 72 } 73 74 #[tokio::test] 75 async fn test_stream() -> std::result::Result<(), io::Error> { 76 let s = Stream::new( 77 "test_poll_stream".to_owned(), 78 0, 79 4096, 80 Arc::new(AtomicU32::new(4096)), 81 Arc::new(AtomicU8::new(AssociationState::Established as u8)), 82 None, 83 Arc::new(PendingQueue::new()), 84 ); 85 86 // getters 87 assert_eq!(0, s.stream_identifier()); 88 assert_eq!(0, s.buffered_amount()); 89 assert_eq!(0, s.buffered_amount_low_threshold()); 90 assert_eq!(0, s.get_num_bytes_in_reassembly_queue().await); 91 92 // setters 93 s.set_default_payload_type(PayloadProtocolIdentifier::Binary); 94 s.set_reliability_params(true, ReliabilityType::Reliable, 0); 95 96 // write 97 let n = s.write(&Bytes::from("Hello "))?; 98 assert_eq!(6, n); 99 assert_eq!(6, s.buffered_amount()); 100 let n = s.write_sctp(&Bytes::from("world"), PayloadProtocolIdentifier::Binary)?; 101 assert_eq!(5, n); 102 assert_eq!(11, s.buffered_amount()); 103 104 // async read 105 // 1. pretend that we've received a chunk 106 s.handle_data(ChunkPayloadData { 107 unordered: true, 108 beginning_fragment: true, 109 ending_fragment: true, 110 user_data: Bytes::from_static(&[0, 1, 2, 3, 4]), 111 payload_type: PayloadProtocolIdentifier::Binary, 112 ..Default::default() 113 }) 114 .await; 115 // 2. read it 116 let mut buf = [0; 5]; 117 s.read(&mut buf).await?; 118 assert_eq!(buf, [0, 1, 2, 3, 4]); 119 120 // shutdown write 121 s.shutdown(Shutdown::Write).await?; 122 // write must fail 123 assert!(s.write(&Bytes::from("error")).is_err()); 124 // read should continue working 125 s.handle_data(ChunkPayloadData { 126 unordered: true, 127 beginning_fragment: true, 128 ending_fragment: true, 129 user_data: Bytes::from_static(&[5, 6, 7, 8, 9]), 130 payload_type: PayloadProtocolIdentifier::Binary, 131 ..Default::default() 132 }) 133 .await; 134 let mut buf = [0; 5]; 135 s.read(&mut buf).await?; 136 assert_eq!(buf, [5, 6, 7, 8, 9]); 137 138 // shutdown read 139 s.shutdown(Shutdown::Read).await?; 140 // read must return 0 141 assert_eq!(Ok(0), s.read(&mut buf).await); 142 143 Ok(()) 144 } 145 146 #[tokio::test] 147 async fn test_poll_stream() -> std::result::Result<(), io::Error> { 148 let s = Arc::new(Stream::new( 149 "test_poll_stream".to_owned(), 150 0, 151 4096, 152 Arc::new(AtomicU32::new(4096)), 153 Arc::new(AtomicU8::new(AssociationState::Established as u8)), 154 None, 155 Arc::new(PendingQueue::new()), 156 )); 157 let mut poll_stream = PollStream::new(s.clone()); 158 159 // getters 160 assert_eq!(0, poll_stream.stream_identifier()); 161 assert_eq!(0, poll_stream.buffered_amount()); 162 assert_eq!(0, poll_stream.buffered_amount_low_threshold()); 163 assert_eq!(0, poll_stream.get_num_bytes_in_reassembly_queue().await); 164 165 // async write 166 let n = poll_stream.write(&[1, 2, 3]).await?; 167 assert_eq!(3, n); 168 poll_stream.flush().await?; 169 assert_eq!(3, poll_stream.buffered_amount()); 170 171 // async read 172 // 1. pretend that we've received a chunk 173 let sc = s.clone(); 174 sc.handle_data(ChunkPayloadData { 175 unordered: true, 176 beginning_fragment: true, 177 ending_fragment: true, 178 user_data: Bytes::from_static(&[0, 1, 2, 3, 4]), 179 payload_type: PayloadProtocolIdentifier::Binary, 180 ..Default::default() 181 }) 182 .await; 183 // 2. read it 184 let mut buf = [0; 5]; 185 poll_stream.read(&mut buf).await?; 186 assert_eq!(buf, [0, 1, 2, 3, 4]); 187 188 // shutdown write 189 poll_stream.shutdown().await?; 190 // write must fail 191 assert!(poll_stream.write(&[1, 2, 3]).await.is_err()); 192 // read should continue working 193 sc.handle_data(ChunkPayloadData { 194 unordered: true, 195 beginning_fragment: true, 196 ending_fragment: true, 197 user_data: Bytes::from_static(&[5, 6, 7, 8, 9]), 198 payload_type: PayloadProtocolIdentifier::Binary, 199 ..Default::default() 200 }) 201 .await; 202 let mut buf = [0; 5]; 203 poll_stream.read(&mut buf).await?; 204 assert_eq!(buf, [5, 6, 7, 8, 9]); 205 206 // misc. 207 let clone = poll_stream.clone(); 208 assert_eq!(clone.stream_identifier(), poll_stream.stream_identifier()); 209 210 Ok(()) 211 } 212