use crate::{
Buf, BufMut, BufResult, Codec, Cursor,
ietf::{ipv4, transport},
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct Header {
pub source_port: Port,
pub destination_port: Port,
pub length: u16,
pub checksum: Checksum,
}
impl Codec for Header {
fn encode<W: BufMut>(&self, writer: &mut Cursor<W>, _: ()) -> BufResult<()> {
self.source_port.encode(writer, ())?;
self.destination_port.encode(writer, ())?;
self.length.encode(writer, ())?;
self.checksum.encode(writer, ())
}
fn decode<R: Buf>(reader: &mut Cursor<R>, _: ()) -> BufResult<Self> {
Ok(Self {
source_port: Port::decode(reader, ())?,
destination_port: Port::decode(reader, ())?,
length: u16::decode(reader, ())?,
checksum: Checksum::decode(reader, ())?,
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
#[repr(transparent)]
pub struct Port(pub transport::Port);
impl Port {
pub const DNS: Self = Self(transport::Port(53));
pub const DHCP_SERVER: Self = Self(transport::Port(67));
pub const DHCP_CLIENT: Self = Self(transport::Port(68));
pub const TFTP: Self = Self(transport::Port(69));
pub const HTTP: Self = Self(transport::Port(80));
pub const NTP: Self = Self(transport::Port(123));
pub const HTTPS: Self = Self(transport::Port(443));
pub const fn is_system(self) -> bool {
self.0.is_system()
}
pub const fn is_user(self) -> bool {
self.0.is_user()
}
pub const fn is_dynamic(self) -> bool {
self.0.is_dynamic()
}
}
impl Default for Port {
fn default() -> Self {
Self(transport::Port::default())
}
}
impl From<u16> for Port {
#[inline]
fn from(val: u16) -> Self {
Port(transport::Port::new(val))
}
}
impl From<Port> for u16 {
#[inline]
fn from(port: Port) -> Self {
port.0.as_u16()
}
}
impl From<transport::Port> for Port {
#[inline]
fn from(port: transport::Port) -> Self {
Port(port)
}
}
impl From<Port> for transport::Port {
#[inline]
fn from(port: Port) -> Self {
port.0
}
}
impl Codec for Port {
fn encode<W: BufMut>(&self, writer: &mut Cursor<W>, _: ()) -> BufResult<()> {
self.0.encode(writer, ())
}
fn decode<R: Buf>(reader: &mut Cursor<R>, _: ()) -> BufResult<Self> {
Ok(Self(transport::Port::decode(reader, ())?))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
#[repr(transparent)]
pub struct Checksum(pub u16);
impl Checksum {
pub fn calculate_ipv4(
pseudo_header: &Ipv4PseudoHeader,
header: &Header,
payload: &[u8],
) -> Self {
let mut sum: u32 = 0;
let source_ip_bytes: [u8; 4] = pseudo_header.source_address.into();
sum += u16::from_be_bytes([source_ip_bytes[0], source_ip_bytes[1]]) as u32;
sum += u16::from_be_bytes([source_ip_bytes[2], source_ip_bytes[3]]) as u32;
let destination_ip_bytes: [u8; 4] = pseudo_header.destination_address.into();
sum += u16::from_be_bytes([destination_ip_bytes[0], destination_ip_bytes[1]]) as u32;
sum += u16::from_be_bytes([destination_ip_bytes[2], destination_ip_bytes[3]]) as u32;
sum += 0x0011;
sum += pseudo_header.udp_length as u32;
sum += u16::from(header.source_port) as u32;
sum += u16::from(header.destination_port) as u32;
sum += header.length as u32;
let mut i = 0;
while i + 1 < payload.len() {
sum += u16::from_be_bytes([payload[i], payload[i + 1]]) as u32;
i += 2;
}
if i < payload.len() {
sum += (payload[i] as u32) << 8;
}
while (sum >> 16) != 0 {
sum = (sum & 0xFFFF) + (sum >> 16);
}
let checksum = !(sum as u16);
Checksum(if checksum == 0 { 0xFFFF } else { checksum })
}
pub fn calculate_ipv6(
pseudo_header: &Ipv6PseudoHeader,
header: &Header,
payload: &[u8],
) -> Self {
let mut sum: u32 = 0;
for i in 0..8 {
sum += u16::from_be_bytes([
pseudo_header.source_ip[i * 2],
pseudo_header.source_ip[i * 2 + 1],
]) as u32;
}
for i in 0..8 {
sum += u16::from_be_bytes([
pseudo_header.destination_ip[i * 2],
pseudo_header.destination_ip[i * 2 + 1],
]) as u32;
}
sum += ((pseudo_header.udp_length >> 16) & 0xFFFF) as u32;
sum += (pseudo_header.udp_length & 0xFFFF) as u32;
sum += 0x0011;
sum += u16::from(header.source_port) as u32;
sum += u16::from(header.destination_port) as u32;
sum += header.length as u32;
let mut i = 0;
while i + 1 < payload.len() {
sum += u16::from_be_bytes([payload[i], payload[i + 1]]) as u32;
i += 2;
}
if i < payload.len() {
sum += (payload[i] as u32) << 8;
}
while (sum >> 16) != 0 {
sum = (sum & 0xFFFF) + (sum >> 16);
}
let checksum = !(sum as u16);
Checksum(checksum)
}
}
impl Codec for Checksum {
fn encode<W: BufMut>(&self, writer: &mut Cursor<W>, _: ()) -> BufResult<()> {
self.0.encode(writer, ())
}
fn decode<R: Buf>(reader: &mut Cursor<R>, _: ()) -> BufResult<Self> {
Ok(Self(u16::decode(reader, ())?))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct Ipv4PseudoHeader {
pub source_address: ipv4::Address,
pub destination_address: ipv4::Address,
pub udp_length: u16,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct Ipv6PseudoHeader {
pub source_ip: [u8; 16],
pub destination_ip: [u8; 16],
pub udp_length: u32,
}
#[cfg(test)]
mod tests {
use core::fmt::Debug;
use super::{Checksum, Header, Ipv4PseudoHeader, Port};
use crate::{Codec, Cursor, ietf::ipv4};
fn codec_roundtrip<T: Codec<C> + Debug + Eq, C: Copy>(
etalon_struct: T,
etalon_bytes: &[u8],
context: C,
) {
let mut encoded_bytes = vec![];
{
let writer = &mut Cursor::new(&mut encoded_bytes);
etalon_struct.encode(writer, context).unwrap();
}
assert_eq!(etalon_bytes, &encoded_bytes);
let decoded_struct = {
let reader = &mut Cursor::new(&mut encoded_bytes);
T::decode(reader, context).unwrap()
};
assert_eq!(etalon_struct, decoded_struct);
encoded_bytes.fill(0x00);
{
let writer = &mut Cursor::new(&mut encoded_bytes);
decoded_struct.encode(writer, context).unwrap();
}
assert_eq!(etalon_bytes, &encoded_bytes);
}
#[test]
fn port() {
let etalon_bytes = &[0x00, 0x35]; let etalon_struct = Port::DNS;
codec_roundtrip(etalon_struct, etalon_bytes, ());
}
#[test]
fn header() {
let etalon_bytes = &[
0x00, 0x35, 0xC3, 0x50, 0x00, 0x10, 0x00, 0x00, ];
let etalon_struct = Header {
source_port: Port::DNS,
destination_port: Port::from(50000u16),
length: 16,
checksum: Checksum(0),
};
codec_roundtrip(etalon_struct, etalon_bytes, ());
}
#[test]
fn checksum_calculation() {
let pseudo_header = Ipv4PseudoHeader {
source_address: ipv4::Address::from([192, 168, 1, 1]),
destination_address: ipv4::Address::from([192, 168, 1, 2]),
udp_length: 16,
};
let header = Header {
source_port: Port::DNS,
destination_port: Port::from(50000u16),
length: 16,
checksum: Checksum(0),
};
let payload = &[];
let checksum = Checksum::calculate_ipv4(&pseudo_header, &header, payload);
assert_ne!(checksum.0, 0); }
}