use crate::checksum::{checksum, transport_checksum};
use crate::l4::TcpFlags;
use crate::{EtherType, Frame, Protocol};
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
pub fn build_ipv4(
src: Ipv4Addr,
dst: Ipv4Addr,
proto: Protocol,
ttl: u8,
payload: &[u8],
) -> Vec<u8> {
let total = 20 + payload.len();
let mut v = Vec::with_capacity(total);
v.push(0x45); v.push(0); v.extend_from_slice(&(total.min(u16::MAX as usize) as u16).to_be_bytes());
v.extend_from_slice(&[0, 0]); v.extend_from_slice(&[0, 0]); v.push(ttl);
v.push(proto.as_u8());
v.extend_from_slice(&[0, 0]); v.extend_from_slice(&src.octets());
v.extend_from_slice(&dst.octets());
let sum = checksum(&v[..20]);
v[10..12].copy_from_slice(&sum.to_be_bytes());
v.extend_from_slice(payload);
v
}
pub fn build_ipv6(
src: Ipv6Addr,
dst: Ipv6Addr,
next_header: Protocol,
hop_limit: u8,
payload: &[u8],
) -> Vec<u8> {
let mut v = Vec::with_capacity(40 + payload.len());
v.extend_from_slice(&[0x60, 0, 0, 0]); v.extend_from_slice(&(payload.len().min(u16::MAX as usize) as u16).to_be_bytes());
v.push(next_header.as_u8());
v.push(hop_limit);
v.extend_from_slice(&src.octets());
v.extend_from_slice(&dst.octets());
v.extend_from_slice(payload);
v
}
pub fn build_ip(
src: IpAddr,
dst: IpAddr,
proto: Protocol,
hop_limit: u8,
payload: &[u8],
) -> Option<Vec<u8>> {
match (src, dst) {
(IpAddr::V4(s), IpAddr::V4(d)) => Some(build_ipv4(s, d, proto, hop_limit, payload)),
(IpAddr::V6(s), IpAddr::V6(d)) => Some(build_ipv6(s, d, proto, hop_limit, payload)),
_ => None,
}
}
pub fn build_udp(
src: IpAddr,
dst: IpAddr,
src_port: u16,
dst_port: u16,
payload: &[u8],
) -> Vec<u8> {
let len = 8 + payload.len();
let mut v = Vec::with_capacity(len);
v.extend_from_slice(&src_port.to_be_bytes());
v.extend_from_slice(&dst_port.to_be_bytes());
v.extend_from_slice(&(len.min(u16::MAX as usize) as u16).to_be_bytes());
v.extend_from_slice(&[0, 0]); v.extend_from_slice(payload);
let sum = transport_checksum(Protocol::UDP, src, dst, &v);
let sum = if sum == 0 { 0xFFFF } else { sum };
v[6..8].copy_from_slice(&sum.to_be_bytes());
v
}
#[allow(clippy::too_many_arguments)]
pub fn build_tcp(
src: IpAddr,
dst: IpAddr,
src_port: u16,
dst_port: u16,
seq: u32,
ack: u32,
flags: TcpFlags,
window: u16,
payload: &[u8],
) -> Vec<u8> {
let mut v = Vec::with_capacity(20 + payload.len());
v.extend_from_slice(&src_port.to_be_bytes());
v.extend_from_slice(&dst_port.to_be_bytes());
v.extend_from_slice(&seq.to_be_bytes());
v.extend_from_slice(&ack.to_be_bytes());
v.push(5 << 4); v.push(flags.bits());
v.extend_from_slice(&window.to_be_bytes());
v.extend_from_slice(&[0, 0]); v.extend_from_slice(&[0, 0]); v.extend_from_slice(payload);
let sum = transport_checksum(Protocol::TCP, src, dst, &v);
v[16..18].copy_from_slice(&sum.to_be_bytes());
v
}
pub fn build_icmpv4(
message_type: u8,
code: u8,
rest_of_header: [u8; 4],
payload: &[u8],
) -> Vec<u8> {
let mut v = Vec::with_capacity(8 + payload.len());
v.push(message_type);
v.push(code);
v.extend_from_slice(&[0, 0]); v.extend_from_slice(&rest_of_header);
v.extend_from_slice(payload);
let sum = checksum(&v);
v[2..4].copy_from_slice(&sum.to_be_bytes());
v
}
pub fn build_icmpv6(
src: Ipv6Addr,
dst: Ipv6Addr,
message_type: u8,
code: u8,
rest_of_header: [u8; 4],
payload: &[u8],
) -> Vec<u8> {
let mut v = Vec::with_capacity(8 + payload.len());
v.push(message_type);
v.push(code);
v.extend_from_slice(&[0, 0]); v.extend_from_slice(&rest_of_header);
v.extend_from_slice(payload);
let sum = transport_checksum(Protocol::ICMPV6, src.into(), dst.into(), &v);
v[2..4].copy_from_slice(&sum.to_be_bytes());
v
}
pub fn push_vlan(frame: &Frame, vid: u16, pcp: u8) -> Vec<u8> {
let bytes = frame.as_bytes();
if bytes.len() < 14 || frame.has_vlan() {
return bytes.to_vec();
}
let tci = ((pcp as u16 & 0x07) << 13) | (vid & 0x0FFF);
let mut v = Vec::with_capacity(bytes.len() + 4);
v.extend_from_slice(&bytes[0..12]); v.extend_from_slice(&EtherType::VLAN.as_u16().to_be_bytes());
v.extend_from_slice(&tci.to_be_bytes());
v.extend_from_slice(&bytes[12..]); v
}
pub fn pop_vlan(frame: &Frame) -> Vec<u8> {
let bytes = frame.as_bytes();
if !frame.has_vlan() {
return bytes.to_vec();
}
let mut v = Vec::with_capacity(bytes.len() - 4);
v.extend_from_slice(&bytes[0..12]);
v.extend_from_slice(&bytes[16..]);
v
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{MacAddr, Packet, build_frame};
const V4_A: Ipv4Addr = Ipv4Addr::new(10, 0, 0, 1);
const V4_B: Ipv4Addr = Ipv4Addr::new(10, 0, 0, 2);
fn v6_a() -> Ipv6Addr {
"2001:db8::1".parse().unwrap()
}
fn v6_b() -> Ipv6Addr {
"2001:db8::2".parse().unwrap()
}
#[test]
fn ipv4_udp_roundtrip() {
let udp = build_udp(V4_A.into(), V4_B.into(), 5000, 53, b"hello");
let buf = build_ipv4(V4_A, V4_B, Protocol::UDP, 64, &udp);
let p = Packet::from_slice(&buf);
assert_eq!(p.version(), 4);
assert_eq!(p.ipv4_total_len() as usize, buf.len());
assert!(p.verify_ipv4_checksum());
assert_eq!(p.verify_transport_checksum(), Some(true));
let dg = p.udp().unwrap();
assert_eq!(dg.src_port(), 5000);
assert_eq!(dg.dst_port(), 53);
assert_eq!(dg.payload(), b"hello");
}
#[test]
fn ipv6_udp_roundtrip() {
let udp = build_udp(v6_a().into(), v6_b().into(), 1, 2, b"x");
let buf = build_ipv6(v6_a(), v6_b(), Protocol::UDP, 64, &udp);
let p = Packet::from_slice(&buf);
assert_eq!(p.version(), 6);
assert_eq!(p.transport_protocol(), Protocol::UDP);
assert_eq!(p.verify_transport_checksum(), Some(true));
assert_eq!(p.udp().unwrap().payload(), b"x");
}
#[test]
fn ipv4_tcp_roundtrip() {
let tcp = build_tcp(
V4_A.into(),
V4_B.into(),
1234,
80,
1000,
2000,
TcpFlags::SYN | TcpFlags::ACK,
65535,
b"body",
);
let buf = build_ipv4(V4_A, V4_B, Protocol::TCP, 64, &tcp);
let p = Packet::from_slice(&buf);
assert_eq!(p.verify_transport_checksum(), Some(true));
let seg = p.tcp().unwrap();
assert_eq!(seg.src_port(), 1234);
assert_eq!(seg.seq(), 1000);
assert_eq!(seg.ack(), 2000);
assert!(seg.flags().contains(TcpFlags::SYN | TcpFlags::ACK));
assert_eq!(seg.payload(), b"body");
}
#[test]
fn icmpv4_checksum_is_self_verifying() {
let msg = build_icmpv4(8, 0, [0, 1, 0, 2], b"ping");
let buf = build_ipv4(V4_A, V4_B, Protocol::ICMP, 64, &msg);
let p = Packet::from_slice(&buf);
let icmp = p.icmp().unwrap();
assert!(icmp.verify_icmpv4_checksum());
assert_eq!(icmp.echo_id(), 1);
assert_eq!(icmp.echo_seq(), 2);
assert_eq!(p.verify_transport_checksum(), Some(true));
}
#[test]
fn icmpv6_checksum_covers_pseudo_header() {
let msg = build_icmpv6(v6_a(), v6_b(), 128, 0, [0, 1, 0, 1], b"ping");
let buf = build_ipv6(v6_a(), v6_b(), Protocol::ICMPV6, 64, &msg);
let p = Packet::from_slice(&buf);
assert_eq!(p.verify_transport_checksum(), Some(true));
let other = build_ipv6(
v6_a(),
"2001:db8::99".parse().unwrap(),
Protocol::ICMPV6,
64,
&msg,
);
assert_eq!(
Packet::from_slice(&other).verify_transport_checksum(),
Some(false)
);
}
#[test]
fn build_ip_rejects_mixed_families() {
assert!(build_ip(V4_A.into(), v6_b().into(), Protocol::UDP, 64, &[]).is_none());
assert!(build_ip(V4_A.into(), V4_B.into(), Protocol::UDP, 64, &[]).is_some());
}
#[test]
fn udp_zero_checksum_becomes_all_ones() {
for i in 0..2000u16 {
let udp = build_udp(
V4_A.into(),
V4_B.into(),
i,
i.wrapping_mul(7),
&i.to_be_bytes(),
);
assert_ne!(&udp[6..8], &[0, 0]);
}
}
#[test]
fn recompute_after_rewrite_matches_builder() {
let udp = build_udp(V4_A.into(), V4_B.into(), 5000, 53, b"hello");
let mut buf = build_ipv4(V4_A, V4_B, Protocol::UDP, 64, &udp);
let new_src = Ipv4Addr::new(192, 168, 0, 7);
let p = Packet::from_mut(&mut buf);
p.set_ipv4_src_addr(new_src);
p.recompute_checksums();
assert!(p.verify_ipv4_checksum());
assert_eq!(p.verify_transport_checksum(), Some(true));
let expect_udp = build_udp(new_src.into(), V4_B.into(), 5000, 53, b"hello");
let expect = build_ipv4(new_src, V4_B, Protocol::UDP, 64, &expect_udp);
assert_eq!(buf, expect);
}
#[test]
fn vlan_push_and_pop() {
let payload = [1u8, 2, 3, 4];
let eth = build_frame(
MacAddr::broadcast(),
MacAddr::zero(),
EtherType::IPV4,
&payload,
);
let f = Frame::from_slice(ð);
assert!(!f.has_vlan());
let tagged = push_vlan(f, 42, 3);
let tf = Frame::from_slice(&tagged);
assert!(tf.has_vlan());
assert_eq!(tf.vlan_id(), 42);
assert_eq!(tf.vlan_pcp(), 3);
assert_eq!(tf.ether_type(), EtherType::IPV4);
assert_eq!(tf.payload(), &payload);
assert_eq!(tf.dst_mac(), f.dst_mac());
assert_eq!(push_vlan(tf, 7, 0), tagged);
assert_eq!(pop_vlan(tf), eth);
assert_eq!(pop_vlan(f), eth);
}
}