use bnb::{BitEnum, EncodeExt, bin, bitfield, u2, u4, u6, u13};
use std::net::Ipv4Addr;
use tracing::info;
#[bitfield(u8, bits = msb)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct VersionIhl {
version: u4,
ihl: u4,
}
#[bitfield(u8, bits = msb)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct Tos {
dscp: u6,
ecn: Ecn,
}
#[derive(BitEnum, Clone, Copy, Debug, PartialEq, Eq)]
#[bit_enum(u2)]
enum Ecn {
NotEct,
Ect1,
Ect0,
Ce,
}
#[bitfield(u16, bits = msb)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct FlagsFrag {
reserved: bool, dont_fragment: bool,
more_fragments: bool,
fragment_offset: u13,
}
#[derive(BitEnum, Clone, Copy, Debug, PartialEq, Eq)]
#[bit_enum(u8)]
#[repr(u8)]
enum Protocol {
Icmp = 1,
Tcp = 6,
Udp = 17,
#[catch_all]
Other(u8),
}
#[bin(big)]
#[derive(Debug, Clone, PartialEq, Eq)]
struct Ipv4Header {
ver_ihl: VersionIhl,
tos: Tos,
total_length: u16,
identification: u16,
flags_frag: FlagsFrag,
ttl: u8,
protocol: Protocol,
#[bw(calc = self.header_checksum())]
#[builder(default)]
checksum: u16,
#[br(map = |raw: u32| Ipv4Addr::from(raw))]
#[bw(map = |ip: &Ipv4Addr| u32::from(*ip))]
src: Ipv4Addr,
#[br(map = |raw: u32| Ipv4Addr::from(raw))]
#[bw(map = |ip: &Ipv4Addr| u32::from(*ip))]
dst: Ipv4Addr,
}
impl Ipv4Header {
fn header_checksum(&self) -> u16 {
let words = [
(u16::from(self.ver_ihl.to_raw()) << 8) | u16::from(self.tos.to_raw()),
self.total_length,
self.identification,
self.flags_frag.to_raw(),
(u16::from(self.ttl) << 8) | u16::from(u8::from(self.protocol)),
(u32::from(self.src) >> 16) as u16,
u32::from(self.src) as u16,
(u32::from(self.dst) >> 16) as u16,
u32::from(self.dst) as u16,
];
let mut sum: u32 = words.iter().map(|&w| u32::from(w)).sum();
while (sum >> 16) != 0 {
sum = (sum & 0xFFFF) + (sum >> 16);
}
!(sum as u16)
}
}
fn hex(bytes: &[u8]) -> String {
bytes
.iter()
.map(|b| format!("{b:02x}"))
.collect::<Vec<_>>()
.join(" ")
}
fn checksum_of(bytes: &[u8]) -> u16 {
u16::from_be_bytes([bytes[10], bytes[11]])
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
tracing_subscriber::fmt()
.with_max_level(tracing::Level::INFO)
.with_target(false)
.without_time()
.init();
let wire: [u8; 20] = [
0x45, 0x00, 0x00, 0x73, 0x00, 0x00, 0x40, 0x00, 0x40, 0x11, 0xb8, 0x61, 0xc0, 0xa8, 0x00,
0x01, 0xc0, 0xa8, 0x00, 0xc7,
];
info!(len = wire.len(), bytes = %hex(&wire), "decoding IPv4 header");
let hdr = Ipv4Header::decode_exact(&wire)?;
info!("the full decoded structure:\n{hdr:#?}"); info!(
version = %hdr.ver_ihl.version(),
ihl = %hdr.ver_ihl.ihl(),
ttl = hdr.ttl,
protocol = ?hdr.protocol,
df = hdr.flags_frag.dont_fragment(),
src = %hdr.src,
dst = %hdr.dst,
checksum = %format!("0x{:04x}", hdr.checksum),
is_canonical = hdr.is_canonical(), "decoded header",
);
assert!(hdr.is_canonical());
assert!(hdr.canonical_diff().is_empty());
let verbatim = hdr.to_bytes()?;
info!(bytes = %hex(&verbatim), "to_bytes (verbatim) → byte-identical to the input");
assert_eq!(verbatim, wire);
let built = Ipv4Header::builder()
.ver_ihl(
VersionIhl::new()
.with_version(u4::new(4))
.with_ihl(u4::new(5)),
)
.tos(Tos::new())
.total_length(40)
.identification(0x1c46)
.flags_frag(FlagsFrag::new().with_dont_fragment(true))
.ttl(64)
.protocol(Protocol::Tcp)
.src(Ipv4Addr::new(10, 0, 0, 1))
.dst(Ipv4Addr::new(10, 0, 0, 2))
.build()?;
info!(
stored_checksum = built.checksum, is_canonical = built.is_canonical(), diff = ?built.canonical_diff(), "built a header (checksum unset)",
);
assert!(!built.is_canonical());
let canonical = built.to_canonical_bytes()?;
info!(
bytes = %hex(&canonical),
checksum = %format!("0x{:04x}", checksum_of(&canonical)),
"to_canonical_bytes → checksum filled in",
);
let mut tampered = built.clone();
tampered.checksum = 0xBAD0;
let raw = tampered.to_bytes()?; let fixed = tampered.to_canonical_bytes()?; info!(
diff = ?tampered.canonical_diff(),
verbatim = %format!("0x{:04x}", checksum_of(&raw)),
canonical = %format!("0x{:04x}", checksum_of(&fixed)),
"to_bytes keeps 0xBAD0; to_canonical_bytes recomputes the real value",
);
assert_eq!(checksum_of(&raw), 0xBAD0);
assert_ne!(checksum_of(&fixed), 0xBAD0);
assert!(tampered.clone().to_canonical().is_canonical());
for (label, value) in [
("verbatim", tampered.clone()),
("canonical", tampered.clone().to_canonical()),
] {
let mut socket: Vec<u8> = Vec::new();
value.encode(&mut socket)?;
info!(mode = label, checksum = %format!("0x{:04x}", checksum_of(&socket)), "encode(w)");
}
let mut exotic = wire;
exotic[9] = 0xFD; let parsed = Ipv4Header::decode_exact(&exotic)?;
info!(protocol = ?parsed.protocol, "unknown protocol preserved, not rejected (catch_all)");
assert_eq!(parsed.protocol, Protocol::Other(0xFD));
info!("all checks passed");
Ok(())
}