lib

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

source_wire.rs (7993B)


      1 //! Fixed-memory admission of decrypted ingress before WebSocket decoding.
      2 
      3 use crate::source_ingress::IngressConnection;
      4 use std::io;
      5 use std::pin::Pin;
      6 use std::task::{Context, Poll};
      7 use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
      8 
      9 const MAX_UPGRADE_BYTES: usize = 16 * 1024;
     10 const READ_CHUNK_BYTES: usize = 4096;
     11 
     12 /// Counts HTTP response and frame overhead conservatively in the same raw-byte
     13 /// allowance as payloads. No message or payload is retained by this layer.
     14 pub(crate) struct MeteredIo<S> {
     15     inner: S,
     16     ingress: IngressConnection,
     17     parser: WireParser,
     18     failed: bool,
     19 }
     20 
     21 impl<S> MeteredIo<S> {
     22     pub(crate) fn new(inner: S, ingress: IngressConnection) -> Self {
     23         Self {
     24             inner,
     25             ingress,
     26             parser: WireParser::default(),
     27             failed: false,
     28         }
     29     }
     30 
     31     fn admitted(&self, context: &Context<'_>) -> io::Result<()> {
     32         if self.failed || !self.ingress.admitted(context) {
     33             Err(denied())
     34         } else {
     35             Ok(())
     36         }
     37     }
     38 }
     39 
     40 impl<S: AsyncRead + Unpin> AsyncRead for MeteredIo<S> {
     41     fn poll_read(
     42         self: Pin<&mut Self>,
     43         context: &mut Context<'_>,
     44         output: &mut ReadBuf<'_>,
     45     ) -> Poll<io::Result<()>> {
     46         let this = self.get_mut();
     47         if let Err(error) = this.admitted(context) {
     48             return Poll::Ready(Err(error));
     49         }
     50         if output.remaining() == 0 {
     51             return Poll::Ready(Ok(()));
     52         }
     53         let mut scratch = [0; READ_CHUNK_BYTES];
     54         let length = output.remaining().min(scratch.len());
     55         let mut input = ReadBuf::new(&mut scratch[..length]);
     56         match Pin::new(&mut this.inner).poll_read(context, &mut input) {
     57             Poll::Pending => Poll::Pending,
     58             Poll::Ready(Err(error)) => {
     59                 this.failed = true;
     60                 Poll::Ready(Err(error))
     61             }
     62             Poll::Ready(Ok(())) => {
     63                 let bytes = input.filled();
     64                 let result = if !this.ingress.charge(bytes.len(), 0, 0) {
     65                     Err(denied())
     66                 } else if bytes.is_empty() {
     67                     this.parser.eof()
     68                 } else {
     69                     this.parser
     70                         .consume(bytes, |data| this.ingress.charge(0, 1, usize::from(data)))
     71                 };
     72                 if let Err(error) = result {
     73                     this.failed = true;
     74                     return Poll::Ready(Err(error));
     75                 }
     76                 // Cancellation may race a ready read. Revalidate before any
     77                 // plaintext becomes visible to the WebSocket implementation.
     78                 if let Err(error) = this.admitted(context) {
     79                     this.failed = true;
     80                     return Poll::Ready(Err(error));
     81                 }
     82                 output.put_slice(bytes);
     83                 Poll::Ready(Ok(()))
     84             }
     85         }
     86     }
     87 }
     88 
     89 impl<S: AsyncWrite + Unpin> AsyncWrite for MeteredIo<S> {
     90     fn poll_write(
     91         self: Pin<&mut Self>,
     92         context: &mut Context<'_>,
     93         bytes: &[u8],
     94     ) -> Poll<io::Result<usize>> {
     95         let this = self.get_mut();
     96         if this.failed || !this.ingress.admitted_write(context) {
     97             return Poll::Ready(Err(denied()));
     98         }
     99         Pin::new(&mut this.inner).poll_write(context, bytes)
    100     }
    101 
    102     fn poll_flush(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<io::Result<()>> {
    103         let this = self.get_mut();
    104         if this.failed || !this.ingress.admitted_write(context) {
    105             return Poll::Ready(Err(denied()));
    106         }
    107         Pin::new(&mut this.inner).poll_flush(context)
    108     }
    109 
    110     fn poll_shutdown(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<io::Result<()>> {
    111         // Closing the underlying transport remains possible after revocation.
    112         Pin::new(&mut self.get_mut().inner).poll_shutdown(context)
    113     }
    114 }
    115 
    116 #[derive(Default)]
    117 struct WireParser {
    118     upgrade_bytes: usize,
    119     delimiter: u32,
    120     upgraded: bool,
    121     header: [u8; 10],
    122     header_len: usize,
    123     header_needed: usize,
    124     payload_remaining: usize,
    125 }
    126 
    127 impl WireParser {
    128     fn consume(
    129         &mut self,
    130         mut bytes: &[u8],
    131         mut admit_frame: impl FnMut(bool) -> bool,
    132     ) -> io::Result<()> {
    133         while let Some((&byte, rest)) = bytes.split_first() {
    134             if !self.upgraded {
    135                 self.upgrade_bytes += 1;
    136                 if self.upgrade_bytes > MAX_UPGRADE_BYTES {
    137                     return Err(invalid("WebSocket upgrade response exceeds its limit"));
    138                 }
    139                 self.delimiter = (self.delimiter << 8) | u32::from(byte);
    140                 self.upgraded = self.delimiter == u32::from_be_bytes(*b"\r\n\r\n");
    141                 bytes = rest;
    142             } else if self.payload_remaining != 0 {
    143                 let consumed = self.payload_remaining.min(bytes.len());
    144                 self.payload_remaining -= consumed;
    145                 bytes = &bytes[consumed..];
    146             } else {
    147                 if self.header_len == 0 {
    148                     let opcode = byte & 0x0f;
    149                     if !admit_frame(matches!(opcode, 1 | 2)) {
    150                         return Err(denied());
    151                     }
    152                     if byte & 0x70 != 0 || !matches!(opcode, 0 | 1 | 2 | 8 | 9 | 10) {
    153                         return Err(invalid("unsupported WebSocket frame flags or opcode"));
    154                     }
    155                     if opcode >= 8 && byte & 0x80 == 0 {
    156                         return Err(invalid("fragmented WebSocket control frame"));
    157                     }
    158                     self.header_needed = 2;
    159                 }
    160                 self.header[self.header_len] = byte;
    161                 self.header_len += 1;
    162                 bytes = rest;
    163                 if self.header_len == 2 {
    164                     if byte & 0x80 != 0 {
    165                         return Err(invalid("masked server WebSocket frame"));
    166                     }
    167                     self.header_needed = match byte {
    168                         126 => 4,
    169                         127 => 10,
    170                         _ => 2,
    171                     };
    172                 }
    173                 if self.header_len == self.header_needed {
    174                     self.payload_remaining = self.payload_length()?;
    175                     self.header_len = 0;
    176                 }
    177             }
    178         }
    179         Ok(())
    180     }
    181 
    182     fn payload_length(&self) -> io::Result<usize> {
    183         let (length, minimum) = match self.header[1] {
    184             126 => (
    185                 u64::from(u16::from_be_bytes([self.header[2], self.header[3]])),
    186                 126,
    187             ),
    188             127 => {
    189                 let mut extended = [0; 8];
    190                 extended.copy_from_slice(&self.header[2..10]);
    191                 (u64::from_be_bytes(extended), 65536)
    192             }
    193             length => (u64::from(length), 0),
    194         };
    195         if length < minimum || length > crate::relay::MAX_WIRE_MESSAGE_BYTES as u64 {
    196             return Err(invalid(
    197                 "WebSocket frame length is noncanonical or exceeds its limit",
    198             ));
    199         }
    200         if self.header[0] & 0x0f >= 8 && length > 125 {
    201             return Err(invalid("WebSocket control frame exceeds its limit"));
    202         }
    203         usize::try_from(length).map_err(|_| invalid("WebSocket frame length cannot be represented"))
    204     }
    205 
    206     fn eof(&self) -> io::Result<()> {
    207         if !self.upgraded || self.header_len != 0 || self.payload_remaining != 0 {
    208             Err(io::Error::new(
    209                 io::ErrorKind::UnexpectedEof,
    210                 "truncated WebSocket ingress",
    211             ))
    212         } else {
    213             Ok(())
    214         }
    215     }
    216 }
    217 
    218 fn invalid(message: &'static str) -> io::Error {
    219     io::Error::new(io::ErrorKind::InvalidData, message)
    220 }
    221 
    222 fn denied() -> io::Error {
    223     io::Error::new(
    224         io::ErrorKind::ConnectionAborted,
    225         "relay ingress budget or generation revoked",
    226     )
    227 }
    228 
    229 #[cfg(test)]
    230 #[path = "source_wire_tests.rs"]
    231 mod tests;