internet 0.1.0

Network library for rust
Documentation
//! TCP Header Mapping.
//!
//! Provides zero-copy read and write access to TCP header fields.
//!
//! As defined in [RFC 9293].
//!
//! [IETF RFC 9293]: https://datatracker.ietf.org/doc/html/rfc9293

use crate::ietf::tcp::{Checksum, Port};
use crate::{Buf, BufError, BufMut, BufResult};

/// A zero-copy mapping for a TCP header.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct HeaderMapping<T> {
    buffer: T,
}

impl<T> HeaderMapping<T> {
    /// Creates a new header mapping.
    pub fn new(buffer: T) -> Self {
        Self { buffer }
    }

    /// Consumes the mapping and returns the underlying buffer.
    pub fn into_inner(self) -> T {
        self.buffer
    }

    /// Returns a reference to the underlying buffer.
    pub fn as_inner(&self) -> &T {
        &self.buffer
    }
}

impl<T: Buf> HeaderMapping<T> {
    /// Reads the source port.
    pub fn read_source_port(&self) -> Port {
        unsafe { Port::new(self.buffer.get_u16_be_unchecked(0)) }
    }

    /// Reads the destination port.
    pub fn read_destination_port(&self) -> Port {
        unsafe { Port::new(self.buffer.get_u16_be_unchecked(2)) }
    }

    /// Reads the sequence number.
    pub fn read_sequence_number(&self) -> u32 {
        unsafe { self.buffer.get_u32_be_unchecked(4) }
    }

    /// Reads the acknowledgment number.
    pub fn read_acknowledgment_number(&self) -> u32 {
        unsafe { self.buffer.get_u32_be_unchecked(8) }
    }

    /// Reads the data offset in 32-bit words.
    pub fn read_data_offset(&self) -> u8 {
        unsafe { ((self.buffer.get_u16_be_unchecked(12) >> 12) & 0x0F) as u8 }
    }

    /// Reads the NS flag.
    pub fn read_ns(&self) -> bool {
        unsafe { (self.buffer.get_u16_be_unchecked(12) & 0x0100) != 0 }
    }

    /// Reads the CWR flag.
    pub fn read_cwr(&self) -> bool {
        unsafe { (self.buffer.get_u16_be_unchecked(12) & 0x0080) != 0 }
    }

    /// Reads the ECE flag.
    pub fn read_ece(&self) -> bool {
        unsafe { (self.buffer.get_u16_be_unchecked(12) & 0x0040) != 0 }
    }

    /// Reads the URG flag.
    pub fn read_urg(&self) -> bool {
        unsafe { (self.buffer.get_u16_be_unchecked(12) & 0x0020) != 0 }
    }

    /// Reads the ACK flag.
    pub fn read_ack(&self) -> bool {
        unsafe { (self.buffer.get_u16_be_unchecked(12) & 0x0010) != 0 }
    }

    /// Reads the PSH flag.
    pub fn read_psh(&self) -> bool {
        unsafe { (self.buffer.get_u16_be_unchecked(12) & 0x0008) != 0 }
    }

    /// Reads the RST flag.
    pub fn read_rst(&self) -> bool {
        unsafe { (self.buffer.get_u16_be_unchecked(12) & 0x0004) != 0 }
    }

    /// Reads the SYN flag.
    pub fn read_syn(&self) -> bool {
        unsafe { (self.buffer.get_u16_be_unchecked(12) & 0x0002) != 0 }
    }

    /// Reads the FIN flag.
    pub fn read_fin(&self) -> bool {
        unsafe { (self.buffer.get_u16_be_unchecked(12) & 0x0001) != 0 }
    }

    /// Reads the window size.
    pub fn read_window(&self) -> u16 {
        unsafe { self.buffer.get_u16_be_unchecked(14) }
    }

    /// Reads the checksum.
    pub fn read_checksum(&self) -> Checksum {
        unsafe { Checksum(self.buffer.get_u16_be_unchecked(16)) }
    }

    /// Reads the urgent pointer.
    pub fn read_urgent_pointer(&self) -> u16 {
        unsafe { self.buffer.get_u16_be_unchecked(18) }
    }
}

impl HeaderMapping<&[u8]> {
    /// Reads the TCP options as a byte slice.
    ///
    /// Returns an empty slice if the data offset is 5 (no options).
    pub fn read_options(&self) -> BufResult<&[u8]> {
        let offset = self.read_data_offset() as usize;
        if offset < 5 {
            return Ok(&[]);
        }
        let len = (offset * 4).saturating_sub(20);
        if self.buffer.len() < 20 + len {
            return Err(BufError::UnexpectedEof);
        }
        Ok(&self.buffer[20..20 + len])
    }
}

impl<T: BufMut> HeaderMapping<T> {
    fn set_flag_checked(&mut self, mask: u16, value: bool) -> BufResult<()> {
        if self.buffer.length() < 14 {
            return Err(BufError::UnexpectedEof);
        }
        let current = unsafe { self.buffer.get_u16_be_unchecked(12) };
        let new_val = if value {
            current | mask
        } else {
            current & !mask
        };
        unsafe { self.buffer.set_u16_be_unchecked(12, new_val) }
        Ok(())
    }

    /// Writes the source port.
    pub fn write_source_port(&mut self, port: Port) -> BufResult<()> {
        if self.buffer.length() < 2 {
            return Err(BufError::UnexpectedEof);
        }
        unsafe { self.buffer.set_u16_be_unchecked(0, port.as_u16()) }
        Ok(())
    }

    /// Writes the destination port.
    pub fn write_destination_port(&mut self, port: Port) -> BufResult<()> {
        if self.buffer.length() < 4 {
            return Err(BufError::UnexpectedEof);
        }
        unsafe { self.buffer.set_u16_be_unchecked(2, port.as_u16()) }
        Ok(())
    }

    /// Writes the sequence number.
    pub fn write_sequence_number(&mut self, seq: u32) -> BufResult<()> {
        if self.buffer.length() < 8 {
            return Err(BufError::UnexpectedEof);
        }
        unsafe { self.buffer.set_u32_be_unchecked(4, seq) }
        Ok(())
    }

    /// Writes the acknowledgment number.
    pub fn write_acknowledgment_number(&mut self, ack: u32) -> BufResult<()> {
        if self.buffer.length() < 12 {
            return Err(BufError::UnexpectedEof);
        }
        unsafe { self.buffer.set_u32_be_unchecked(8, ack) }
        Ok(())
    }

    /// Writes the NS flag.
    pub fn write_ns(&mut self, value: bool) -> BufResult<()> {
        self.set_flag_checked(0x0100, value)
    }

    /// Writes the CWR flag.
    pub fn write_cwr(&mut self, value: bool) -> BufResult<()> {
        self.set_flag_checked(0x0080, value)
    }

    /// Writes the ECE flag.
    pub fn write_ece(&mut self, value: bool) -> BufResult<()> {
        self.set_flag_checked(0x0040, value)
    }

    /// Writes the URG flag.
    pub fn write_urg(&mut self, value: bool) -> BufResult<()> {
        self.set_flag_checked(0x0020, value)
    }

    /// Writes the ACK flag.
    pub fn write_ack(&mut self, value: bool) -> BufResult<()> {
        self.set_flag_checked(0x0010, value)
    }

    /// Writes the PSH flag.
    pub fn write_psh(&mut self, value: bool) -> BufResult<()> {
        self.set_flag_checked(0x0008, value)
    }

    /// Writes the RST flag.
    pub fn write_rst(&mut self, value: bool) -> BufResult<()> {
        self.set_flag_checked(0x0004, value)
    }

    /// Writes the SYN flag.
    pub fn write_syn(&mut self, value: bool) -> BufResult<()> {
        self.set_flag_checked(0x0002, value)
    }

    /// Writes the FIN flag.
    pub fn write_fin(&mut self, value: bool) -> BufResult<()> {
        self.set_flag_checked(0x0001, value)
    }

    /// Writes the window size.
    pub fn write_window(&mut self, window: u16) -> BufResult<()> {
        if self.buffer.length() < 16 {
            return Err(BufError::UnexpectedEof);
        }
        unsafe { self.buffer.set_u16_be_unchecked(14, window) }
        Ok(())
    }

    /// Writes the checksum.
    pub fn write_checksum(&mut self, checksum: Checksum) -> BufResult<()> {
        if self.buffer.length() < 18 {
            return Err(BufError::UnexpectedEof);
        }
        unsafe { self.buffer.set_u16_be_unchecked(16, checksum.0) }
        Ok(())
    }

    /// Writes the urgent pointer.
    pub fn write_urgent_pointer(&mut self, ptr: u16) -> BufResult<()> {
        if self.buffer.length() < 20 {
            return Err(BufError::UnexpectedEof);
        }
        unsafe { self.buffer.set_u16_be_unchecked(18, ptr) }
        Ok(())
    }

    /// Consumes the mapping, writes the data offset, and returns the buffer.
    pub fn consume_write_data_offset(mut self, offset: u8) -> T {
        if self.buffer.length() < 14 {
            return self.buffer;
        }
        let current = unsafe { self.buffer.get_u16_be_unchecked(12) };
        let new_val = (current & 0x0FFF) | (((offset as u16) & 0x0F) << 12);
        unsafe { self.buffer.set_u16_be_unchecked(12, new_val) }
        self.buffer
    }
}

impl HeaderMapping<&mut [u8]> {
    /// Writes the TCP options from a byte slice.
    ///
    /// Pads with zeros if the slice is shorter than the space allocated by the data offset.
    pub fn write_options(&mut self, options: &[u8]) -> BufResult<()> {
        let offset = self.read_data_offset() as usize;
        if offset < 5 {
            return Ok(());
        }
        let len = (offset * 4).saturating_sub(20);
        if options.len() > len {
            return Err(BufError::UnexpectedEof);
        }
        if self.buffer.len() < 20 + len {
            return Err(BufError::UnexpectedEof);
        }

        self.buffer[20..20 + options.len()].copy_from_slice(options);
        if options.len() < len {
            self.buffer[20 + options.len()..20 + len].fill(0);
        }
        Ok(())
    }
}

#[cfg(test)]
mod tests {
    use super::{Checksum, HeaderMapping, Port};

    #[test]
    fn read_fields() {
        let buffer = [
            0x00, 0x50, 0xC0, 0x00, 0xDE, 0xAD, 0xBE, 0xEF, 0xCA, 0xFE, 0xBA, 0xBE, 0x50, 0x12,
            0xFF, 0xFF, 0x12, 0x34, 0x00, 0x00,
        ];
        let mapping = HeaderMapping::new(&buffer[..]);
        assert_eq!(mapping.read_source_port(), Port::HTTP);
        assert_eq!(mapping.read_destination_port(), Port::new(49152));
        assert_eq!(mapping.read_sequence_number(), 0xDEADBEEF);
        assert_eq!(mapping.read_acknowledgment_number(), 0xCAFEBABE);
        assert_eq!(mapping.read_data_offset(), 5);
        assert!(mapping.read_syn());
        assert!(mapping.read_ack());
        assert!(!mapping.read_fin());
    }

    #[test]
    fn write_fields() {
        let mut buffer = [0u8; 20];
        {
            let mut mapping = HeaderMapping::new(&mut buffer[..]);
            mapping.write_source_port(Port::HTTP).unwrap();
            mapping.write_destination_port(Port::new(49152)).unwrap();
            mapping.write_sequence_number(0xDEADBEEF).unwrap();
            mapping.write_acknowledgment_number(0xCAFEBABE).unwrap();
            mapping.write_syn(true).unwrap();
            mapping.write_ack(true).unwrap();
            mapping.write_window(0xFFFF).unwrap();
            mapping.write_checksum(Checksum(0x1234)).unwrap();
        }
        assert_eq!(&buffer[0..2], &[0x00, 0x50]);
        assert_eq!(&buffer[2..4], &[0xC0, 0x00]);
    }

    // #[test]
    // fn options_roundtrip() {
    //   let mut buffer = [0u8; 28];
    //   {
    //     let mut mapping = HeaderMapping::new(&mut buffer[..]);
    //     mapping.write_data_offset(7).unwrap(); // 28 bytes total, 8 bytes for options
    //     let options = [0x02, 0x04, 0x05, 0xB4, 0x01, 0x00, 0x00, 0x00];
    //     mapping.write_options(&options).unwrap();
    //   }
    //   {
    //     let mapping = HeaderMapping::new(&buffer[..]);
    //     assert_eq!(mapping.read_data_offset(), 7);
    //     assert_eq!(
    //       mapping.read_options().unwrap(),
    //       &[0x02, 0x04, 0x05, 0xB4, 0x01, 0x00, 0x00, 0x00]
    //     );
    //   }
    // }
}