use std::fmt::{self, Formatter};
use std::mem;
use zerocopy::byteorder::{BigEndian, U16};
use zerocopy::{FromBytes, IntoBytes, Unaligned};
use crate::packet::{HeaderParser, PacketHeader};
#[repr(C, packed)]
#[derive(
FromBytes, IntoBytes, Unaligned, Debug, Clone, Copy, zerocopy::KnownLayout, zerocopy::Immutable,
)]
pub struct UdpHeader {
src_port: U16<BigEndian>,
dst_port: U16<BigEndian>,
length: U16<BigEndian>,
checksum: U16<BigEndian>,
}
impl UdpHeader {
#[inline]
pub fn src_port(&self) -> u16 {
self.src_port.get()
}
#[inline]
pub fn dst_port(&self) -> u16 {
self.dst_port.get()
}
#[inline]
pub fn length(&self) -> u16 {
self.length.get()
}
#[inline]
pub fn checksum(&self) -> u16 {
self.checksum.get()
}
#[inline]
pub fn header_len(&self) -> usize {
mem::size_of::<UdpHeader>()
}
#[inline]
pub fn payload_len(&self) -> usize {
let total = self.length() as usize;
total.saturating_sub(Self::FIXED_LEN)
}
#[inline]
pub fn is_valid(&self) -> bool {
self.length() >= Self::FIXED_LEN as u16
}
pub fn verify_checksum(&self, src_ip: u32, dst_ip: u32, udp_data: &[u8]) -> bool {
let checksum = self.checksum();
if checksum == 0 {
return true;
}
let computed = Self::compute_checksum(src_ip, dst_ip, udp_data);
computed == checksum
}
pub fn compute_checksum(src_ip: u32, dst_ip: u32, udp_data: &[u8]) -> u16 {
let mut sum: u32 = 0;
sum += (src_ip >> 16) & 0xFFFF;
sum += src_ip & 0xFFFF;
sum += (dst_ip >> 16) & 0xFFFF;
sum += dst_ip & 0xFFFF;
sum += 17;
sum += udp_data.len() as u32;
let mut i = 0;
while i < udp_data.len() {
if i + 1 < udp_data.len() {
let word = u16::from_be_bytes([udp_data[i], udp_data[i + 1]]);
sum += word as u32;
i += 2;
} else {
let word = u16::from_be_bytes([udp_data[i], 0]);
sum += word as u32;
i += 1;
}
}
while sum >> 16 != 0 {
sum = (sum & 0xFFFF) + (sum >> 16);
}
!sum as u16
}
}
impl PacketHeader for UdpHeader {
const NAME: &'static str = "UdpHeader";
#[inline]
fn is_valid(&self) -> bool {
self.is_valid()
}
type InnerType = ();
#[inline]
fn inner_type(&self) -> Self::InnerType {}
}
impl HeaderParser for UdpHeader {
type Output<'a> = &'a UdpHeader;
#[inline]
fn into_view<'a>(header: &'a Self, _: &'a [u8]) -> Self::Output<'a> {
header
}
}
impl fmt::Display for UdpHeader {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(
f,
"UDP {} -> {} len={}",
self.src_port(),
self.dst_port(),
self.length()
)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_udp_header_basic() {
let header = UdpHeader {
src_port: U16::new(53),
dst_port: U16::new(12345),
length: U16::new(16), checksum: U16::new(0),
};
assert_eq!(header.src_port(), 53);
assert_eq!(header.dst_port(), 12345);
assert_eq!(header.length(), 16);
assert_eq!(header.header_len(), 8);
assert_eq!(header.payload_len(), 8);
assert!(header.is_valid());
}
#[test]
fn test_udp_header_validation() {
let invalid_header = UdpHeader {
src_port: U16::new(53),
dst_port: U16::new(12345),
length: U16::new(7), checksum: U16::new(0),
};
assert!(!invalid_header.is_valid());
let valid_header = UdpHeader {
src_port: U16::new(53),
dst_port: U16::new(12345),
length: U16::new(8), checksum: U16::new(0),
};
assert!(valid_header.is_valid());
}
#[test]
fn test_udp_checksum_zero() {
let header = UdpHeader {
src_port: U16::new(53),
dst_port: U16::new(12345),
length: U16::new(8),
checksum: U16::new(0),
};
assert!(header.verify_checksum(0x7f000001, 0x7f000001, &[]));
}
#[test]
fn test_udp_header_size() {
assert_eq!(mem::size_of::<UdpHeader>(), 8);
assert_eq!(UdpHeader::FIXED_LEN, 8);
}
#[test]
fn test_udp_parsing_basic() {
let packet = create_test_packet();
let result = UdpHeader::from_bytes(&packet);
assert!(result.is_ok());
let (header, payload) = result.unwrap();
assert_eq!(header.src_port(), 12345);
assert_eq!(header.dst_port(), 53);
assert_eq!(header.length(), 16);
assert_eq!(payload.len(), 8); assert!(header.is_valid());
}
#[test]
fn test_udp_parsing_too_small() {
let packet = vec![0u8; 7];
let result = UdpHeader::from_bytes(&packet);
assert!(result.is_err());
}
#[test]
fn test_udp_total_len() {
let packet = create_test_packet();
let (header, _) = UdpHeader::from_bytes(&packet).unwrap();
assert_eq!(header.total_len(&packet), 8);
}
#[test]
fn test_udp_from_bytes_with_payload() {
let mut packet = Vec::new();
packet.extend_from_slice(&5000u16.to_be_bytes()); packet.extend_from_slice(&8080u16.to_be_bytes());
let payload_data = b"Hello, UDP!";
let total_length = 8 + payload_data.len();
packet.extend_from_slice(&(total_length as u16).to_be_bytes()); packet.extend_from_slice(&0u16.to_be_bytes());
packet.extend_from_slice(payload_data);
let result = UdpHeader::from_bytes(&packet);
assert!(result.is_ok());
let (header, payload) = result.unwrap();
assert_eq!(header.src_port(), 5000);
assert_eq!(header.dst_port(), 8080);
assert_eq!(header.length(), total_length as u16);
assert_eq!(header.payload_len(), payload_data.len());
assert_eq!(payload.len(), payload_data.len());
assert_eq!(payload, payload_data);
}
#[test]
fn test_udp_payload_length_calculation() {
let header1 = UdpHeader {
src_port: U16::new(1234),
dst_port: U16::new(5678),
length: U16::new(8), checksum: U16::new(0),
};
assert_eq!(header1.payload_len(), 0);
let header2 = UdpHeader {
src_port: U16::new(1234),
dst_port: U16::new(5678),
length: U16::new(100), checksum: U16::new(0),
};
assert_eq!(header2.payload_len(), 92);
let header3 = UdpHeader {
src_port: U16::new(1234),
dst_port: U16::new(5678),
length: U16::new(5), checksum: U16::new(0),
};
assert_eq!(header3.payload_len(), 0);
}
#[test]
fn test_udp_dns_packet() {
let mut packet = Vec::new();
packet.extend_from_slice(&54321u16.to_be_bytes()); packet.extend_from_slice(&53u16.to_be_bytes());
let dns_payload = vec![
0xab, 0xcd, 0x01, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, ];
let total_length = 8 + dns_payload.len();
packet.extend_from_slice(&(total_length as u16).to_be_bytes()); packet.extend_from_slice(&0u16.to_be_bytes());
packet.extend_from_slice(&dns_payload);
let (header, payload) = UdpHeader::from_bytes(&packet).unwrap();
assert_eq!(header.src_port(), 54321);
assert_eq!(header.dst_port(), 53);
assert_eq!(header.length(), total_length as u16);
assert_eq!(payload.len(), dns_payload.len());
assert_eq!(payload, dns_payload.as_slice());
}
#[test]
fn test_udp_checksum_computation() {
let src_ip = 0xC0A80101; let dst_ip = 0xC0A80102;
let mut udp_packet = Vec::new();
udp_packet.extend_from_slice(&12345u16.to_be_bytes()); udp_packet.extend_from_slice(&80u16.to_be_bytes()); udp_packet.extend_from_slice(&12u16.to_be_bytes()); udp_packet.extend_from_slice(&0u16.to_be_bytes());
udp_packet.extend_from_slice(b"test");
let checksum = UdpHeader::compute_checksum(src_ip, dst_ip, &udp_packet);
assert_ne!(checksum, 0);
}
#[test]
fn test_udp_multiple_packets() {
let packets: Vec<(u16, u16, Vec<u8>)> = vec![
(1234, 5678, b"payload1".to_vec()),
(80, 54321, b"HTTP response".to_vec()),
(53, 12345, b"DNS".to_vec()),
];
for (src, dst, payload_data) in packets {
let mut packet = Vec::new();
packet.extend_from_slice(&src.to_be_bytes());
packet.extend_from_slice(&dst.to_be_bytes());
packet.extend_from_slice(&((8 + payload_data.len()) as u16).to_be_bytes());
packet.extend_from_slice(&0u16.to_be_bytes());
packet.extend_from_slice(&payload_data);
let (header, payload) = UdpHeader::from_bytes(&packet).unwrap();
assert_eq!(header.src_port(), src);
assert_eq!(header.dst_port(), dst);
assert_eq!(payload, payload_data.as_slice());
}
}
fn create_test_packet() -> Vec<u8> {
let mut packet = Vec::new();
packet.extend_from_slice(&12345u16.to_be_bytes());
packet.extend_from_slice(&53u16.to_be_bytes());
packet.extend_from_slice(&16u16.to_be_bytes());
packet.extend_from_slice(&0u16.to_be_bytes());
packet.extend_from_slice(b"DNS data");
packet
}
}