lib

Core libraries for Radroots
git clone https://radroots.dev/git/lib.git
Log | Files | Refs | README

source_wire_tests.rs (10543B)


      1 use super::*;
      2 use crate::source::budget::{FetchBudget, MAX_FETCH_BYTES};
      3 use crate::source_ingress::{FetchRegistration, IngressRegistry};
      4 use futures::task::{ArcWake, waker};
      5 use std::sync::{
      6     Arc,
      7     atomic::{AtomicUsize, Ordering},
      8 };
      9 
     10 const UPGRADE: &[u8] = b"HTTP/1.1 101 Switching Protocols\r\n\r\n";
     11 const RELAY: &str = "ws://127.0.0.1:1234";
     12 
     13 fn upgraded() -> WireParser {
     14     let mut parser = WireParser::default();
     15     parser
     16         .consume(UPGRADE, |_| panic!("HTTP is not a frame"))
     17         .expect("upgrade");
     18     parser
     19 }
     20 
     21 fn header(length: u64) -> Vec<u8> {
     22     let mut bytes = vec![0x82, 127];
     23     bytes.extend(length.to_be_bytes());
     24     bytes
     25 }
     26 
     27 #[test]
     28 fn maximum_frame_length_is_admitted_and_next_byte_or_overflow_is_denied() {
     29     let maximum = crate::relay::MAX_WIRE_MESSAGE_BYTES as u64;
     30     let mut parser = upgraded();
     31     parser
     32         .consume(&header(maximum), |_| true)
     33         .expect("exact limit");
     34     assert_eq!(parser.payload_remaining, maximum as usize);
     35     for length in [maximum + 1, 1 << 63, u64::MAX] {
     36         assert!(upgraded().consume(&header(length), |_| true).is_err());
     37     }
     38 }
     39 
     40 #[test]
     41 fn headers_payloads_and_upgrade_can_arrive_one_byte_at_a_time() {
     42     let mut parser = WireParser::default();
     43     let mut bytes = UPGRADE.to_vec();
     44     bytes.extend([0x01, 126, 0, 126]);
     45     bytes.extend([b'x'; 126]);
     46     bytes.extend([0x89, 0, 0x80, 2, b'y', b'z', 0x82, 0]);
     47     let mut frames = 0;
     48     let mut data = 0;
     49     for byte in bytes {
     50         parser
     51             .consume(&[byte], |is_data| {
     52                 frames += 1;
     53                 data += usize::from(is_data);
     54                 true
     55             })
     56             .expect("chopped ingress");
     57     }
     58     assert_eq!((frames, data), (4, 2));
     59     parser.eof().expect("complete frames");
     60 }
     61 
     62 #[test]
     63 fn zero_length_continuations_and_controls_each_require_admission() {
     64     let mut parser = upgraded();
     65     let mut admissions = 0;
     66     let error = parser
     67         .consume(&[0x01, 0, 0x00, 0, 0x89, 0, 0x80, 0], |is_data| {
     68             admissions += 1;
     69             assert_eq!(is_data, admissions == 1);
     70             admissions <= 3
     71         })
     72         .expect_err("fourth frame is denied");
     73     assert_eq!(error.kind(), io::ErrorKind::ConnectionAborted);
     74     assert_eq!(admissions, 4);
     75 }
     76 
     77 #[test]
     78 fn invalid_flags_masks_lengths_and_controls_fail_closed() {
     79     for bytes in [
     80         vec![0xc1, 0],
     81         vec![0x83, 0],
     82         vec![0x09, 0],
     83         vec![0x81, 0x80],
     84         vec![0x82, 126, 0, 125],
     85         header(65535),
     86         vec![0x89, 126, 0, 126],
     87     ] {
     88         assert!(
     89             upgraded().consume(&bytes, |_| true).is_err(),
     90             "accepted {bytes:?}"
     91         );
     92     }
     93 }
     94 
     95 #[test]
     96 fn upgrade_prefix_has_an_exact_fixed_bound() {
     97     let mut exact = vec![b'x'; MAX_UPGRADE_BYTES - 4];
     98     exact.extend(b"\r\n\r\n");
     99     let mut parser = WireParser::default();
    100     parser.consume(&exact, |_| true).expect("exact maximum");
    101     parser.eof().expect("upgrade boundary");
    102     let mut excessive = vec![b'x'];
    103     excessive.extend(exact);
    104     assert!(WireParser::default().consume(&excessive, |_| true).is_err());
    105 }
    106 
    107 #[test]
    108 fn eof_rejects_partial_upgrade_header_and_payload() {
    109     assert!(WireParser::default().eof().is_err());
    110     for bytes in [vec![0x81], vec![0x82, 126, 0], vec![0x81, 2, b'x']] {
    111         let mut parser = upgraded();
    112         parser.consume(&bytes, |_| true).expect("partial input");
    113         assert_eq!(
    114             parser.eof().expect_err("truncation").kind(),
    115             io::ErrorKind::UnexpectedEof
    116         );
    117     }
    118     upgraded().eof().expect("frame boundary");
    119 }
    120 
    121 #[derive(Default)]
    122 struct Calls {
    123     reads: AtomicUsize,
    124     writes: AtomicUsize,
    125     wakes: AtomicUsize,
    126 }
    127 
    128 impl ArcWake for Calls {
    129     fn wake_by_ref(arc_self: &Arc<Self>) {
    130         arc_self.wakes.fetch_add(1, Ordering::SeqCst);
    131     }
    132 }
    133 
    134 struct TestIo {
    135     bytes: Vec<u8>,
    136     offset: usize,
    137     chunk: usize,
    138     pending: bool,
    139     calls: Arc<Calls>,
    140 }
    141 
    142 impl AsyncRead for TestIo {
    143     fn poll_read(
    144         mut self: Pin<&mut Self>,
    145         _: &mut Context<'_>,
    146         output: &mut ReadBuf<'_>,
    147     ) -> Poll<io::Result<()>> {
    148         self.calls.reads.fetch_add(1, Ordering::SeqCst);
    149         if self.pending {
    150             return Poll::Pending;
    151         }
    152         let count = (self.bytes.len() - self.offset)
    153             .min(output.remaining())
    154             .min(self.chunk);
    155         output.put_slice(&self.bytes[self.offset..self.offset + count]);
    156         self.offset += count;
    157         Poll::Ready(Ok(()))
    158     }
    159 }
    160 
    161 impl AsyncWrite for TestIo {
    162     fn poll_write(
    163         self: Pin<&mut Self>,
    164         _: &mut Context<'_>,
    165         bytes: &[u8],
    166     ) -> Poll<io::Result<usize>> {
    167         self.calls.writes.fetch_add(1, Ordering::SeqCst);
    168         Poll::Ready(Ok(bytes.len()))
    169     }
    170     fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<io::Result<()>> {
    171         self.calls.writes.fetch_add(1, Ordering::SeqCst);
    172         Poll::Ready(Ok(()))
    173     }
    174     fn poll_shutdown(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<io::Result<()>> {
    175         Poll::Ready(Ok(()))
    176     }
    177 }
    178 
    179 fn fixture(
    180     bytes: Vec<u8>,
    181 ) -> (
    182     MeteredIo<TestIo>,
    183     Arc<FetchBudget>,
    184     FetchRegistration,
    185     Arc<Calls>,
    186 ) {
    187     let registry = IngressRegistry::new([RELAY.to_owned()].into_iter());
    188     let budget = Arc::new(FetchBudget::default());
    189     let registration = registry
    190         .register([RELAY.to_owned()].into_iter(), budget.clone())
    191         .expect("registration");
    192     let connection = registry.connection(RELAY).expect("connection");
    193     let calls = Arc::new(Calls::default());
    194     let inner = TestIo {
    195         bytes,
    196         offset: 0,
    197         chunk: READ_CHUNK_BYTES,
    198         pending: false,
    199         calls: calls.clone(),
    200     };
    201     (
    202         MeteredIo::new(inner, connection),
    203         budget,
    204         registration,
    205         calls,
    206     )
    207 }
    208 
    209 fn read(meter: &mut MeteredIo<TestIo>, calls: &Arc<Calls>) -> (Poll<io::Result<()>>, usize) {
    210     let waker = waker(calls.clone());
    211     let mut context = Context::from_waker(&waker);
    212     let mut output = [0; READ_CHUNK_BYTES];
    213     let mut buffer = ReadBuf::new(&mut output);
    214     let result = Pin::new(meter).poll_read(&mut context, &mut buffer);
    215     (result, buffer.filled().len())
    216 }
    217 
    218 #[test]
    219 fn raw_upgrade_and_frame_overhead_count_before_decode() {
    220     let mut bytes = UPGRADE.to_vec();
    221     bytes.extend([0x81, 0]);
    222     let length = bytes.len();
    223     let (mut meter, budget, _registration, calls) = fixture(bytes);
    224     assert!(budget.wire(MAX_FETCH_BYTES - length, 0, 0));
    225     assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Ok(())), count) if count == length));
    226     assert!(!budget.wire(1, 0, 0));
    227 }
    228 
    229 #[test]
    230 fn excessive_read_is_not_exposed_and_subsequent_io_is_rejected() {
    231     let (mut meter, budget, _registration, calls) = fixture(UPGRADE.to_vec());
    232     assert!(budget.wire(MAX_FETCH_BYTES, 0, 0));
    233     assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Err(_)), 0)));
    234     assert!(budget.exhausted());
    235     let reads = calls.reads.load(Ordering::SeqCst);
    236     assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Err(_)), 0)));
    237     assert_eq!(calls.reads.load(Ordering::SeqCst), reads);
    238     let waker = waker(calls.clone());
    239     let mut context = Context::from_waker(&waker);
    240     assert!(matches!(
    241         Pin::new(&mut meter).poll_write(&mut context, b"REQ"),
    242         Poll::Ready(Err(_))
    243     ));
    244     assert!(matches!(
    245         Pin::new(&mut meter).poll_flush(&mut context),
    246         Poll::Ready(Err(_))
    247     ));
    248     assert_eq!(calls.writes.load(Ordering::SeqCst), 0);
    249 }
    250 
    251 #[test]
    252 fn cancellation_wakes_pending_read_and_revokes_both_io_directions() {
    253     let (mut meter, _budget, registration, calls) = fixture(Vec::new());
    254     meter.inner.pending = true;
    255     assert!(matches!(read(&mut meter, &calls), (Poll::Pending, 0)));
    256     let writer_calls = Arc::new(Calls::default());
    257     let writer_waker = waker(writer_calls.clone());
    258     let mut writer_context = Context::from_waker(&writer_waker);
    259     assert!(matches!(
    260         Pin::new(&mut meter).poll_write(&mut writer_context, b"REQ"),
    261         Poll::Ready(Ok(3))
    262     ));
    263     assert!(matches!(
    264         Pin::new(&mut meter).poll_flush(&mut writer_context),
    265         Poll::Ready(Ok(()))
    266     ));
    267     drop(registration);
    268     assert!(calls.wakes.load(Ordering::SeqCst) > 0);
    269     assert!(writer_calls.wakes.load(Ordering::SeqCst) > 0);
    270     assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Err(_)), 0)));
    271     let waker = waker(calls.clone());
    272     let mut context = Context::from_waker(&waker);
    273     assert!(matches!(
    274         Pin::new(&mut meter).poll_write(&mut context, b"REQ"),
    275         Poll::Ready(Err(_))
    276     ));
    277     assert!(matches!(
    278         Pin::new(&mut meter).poll_shutdown(&mut context),
    279         Poll::Ready(Ok(()))
    280     ));
    281     assert_eq!(calls.reads.load(Ordering::SeqCst), 1);
    282     assert_eq!(calls.writes.load(Ordering::SeqCst), 2);
    283 }
    284 
    285 #[test]
    286 fn continuation_flood_exhausts_frame_budget_without_assembled_message() {
    287     let mut bytes = UPGRADE.to_vec();
    288     bytes.extend([0x01, 0]);
    289     for _ in 0..8192 {
    290         bytes.extend([0x00, 0]);
    291     }
    292     let (mut meter, budget, _registration, calls) = fixture(bytes);
    293     loop {
    294         match read(&mut meter, &calls) {
    295             (Poll::Ready(Ok(())), count) => assert!(count > 0, "unexpected EOF"),
    296             (Poll::Ready(Err(_)), 0) => break,
    297             other => panic!("unexpected result {other:?}"),
    298         }
    299     }
    300     assert!(budget.exhausted());
    301 }
    302 
    303 #[test]
    304 fn empty_data_frames_exhaust_attempt_budget() {
    305     let mut bytes = UPGRADE.to_vec();
    306     for _ in 0..4097 {
    307         bytes.extend([0x81, 0]);
    308     }
    309     let (mut meter, budget, _registration, calls) = fixture(bytes);
    310     loop {
    311         match read(&mut meter, &calls) {
    312             (Poll::Ready(Ok(())), count) => assert!(count > 0, "unexpected EOF"),
    313             (Poll::Ready(Err(_)), 0) => break,
    314             other => panic!("unexpected result {other:?}"),
    315         }
    316     }
    317     assert!(budget.exhausted());
    318 }
    319 
    320 #[test]
    321 fn read_chunks_remain_bounded_and_truncation_is_terminal() {
    322     let mut bytes = UPGRADE.to_vec();
    323     bytes.extend([0x81, 126, 0x20, 0]);
    324     bytes.extend([b'x'; 5000]);
    325     let (mut meter, _budget, _registration, calls) = fixture(bytes);
    326     assert!(matches!(
    327         read(&mut meter, &calls),
    328         (Poll::Ready(Ok(())), READ_CHUNK_BYTES)
    329     ));
    330     assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Ok(())), count) if count > 0));
    331     assert!(
    332         matches!(read(&mut meter, &calls), (Poll::Ready(Err(error)), 0) if error.kind() == io::ErrorKind::UnexpectedEof)
    333     );
    334 }