Skip to main content

arcbox_virtio_net/
header.rs

1//! `VirtIO`-net wire format — header struct and `NetPacket` envelope.
2
3use arcbox_virtio_core::virtio_bindings;
4
5/// `VirtIO` network header.
6#[repr(C)]
7#[derive(Debug, Clone, Copy, Default)]
8pub struct VirtioNetHeader {
9    /// Flags.
10    pub flags: u8,
11    /// GSO type.
12    pub gso_type: u8,
13    /// Header length.
14    pub hdr_len: u16,
15    /// GSO size.
16    pub gso_size: u16,
17    /// Checksum start.
18    pub csum_start: u16,
19    /// Checksum offset.
20    pub csum_offset: u16,
21    /// Number of buffers.
22    pub num_buffers: u16,
23}
24
25impl VirtioNetHeader {
26    /// Size of the header in bytes.
27    pub const SIZE: usize = 12;
28
29    /// No GSO.
30    pub const GSO_NONE: u8 = virtio_bindings::virtio_net::VIRTIO_NET_HDR_GSO_NONE as u8;
31    /// TCP/IPv4 GSO.
32    pub const GSO_TCPV4: u8 = virtio_bindings::virtio_net::VIRTIO_NET_HDR_GSO_TCPV4 as u8;
33    /// UDP GSO.
34    pub const GSO_UDP: u8 = virtio_bindings::virtio_net::VIRTIO_NET_HDR_GSO_UDP as u8;
35    /// TCP/IPv6 GSO.
36    pub const GSO_TCPV6: u8 = virtio_bindings::virtio_net::VIRTIO_NET_HDR_GSO_TCPV6 as u8;
37    /// ECN flag.
38    pub const GSO_ECN: u8 = virtio_bindings::virtio_net::VIRTIO_NET_HDR_GSO_ECN as u8;
39
40    /// Header needs checksum.
41    pub const FLAG_NEEDS_CSUM: u8 = virtio_bindings::virtio_net::VIRTIO_NET_HDR_F_NEEDS_CSUM as u8;
42    /// Data is valid.
43    pub const FLAG_DATA_VALID: u8 = virtio_bindings::virtio_net::VIRTIO_NET_HDR_F_DATA_VALID as u8;
44
45    /// Creates a new header.
46    #[must_use]
47    pub fn new() -> Self {
48        Self::default()
49    }
50
51    /// Parses from bytes.
52    #[must_use]
53    pub fn from_bytes(bytes: &[u8]) -> Option<Self> {
54        if bytes.len() < Self::SIZE {
55            return None;
56        }
57        Some(Self {
58            flags: bytes[0],
59            gso_type: bytes[1],
60            hdr_len: u16::from_le_bytes([bytes[2], bytes[3]]),
61            gso_size: u16::from_le_bytes([bytes[4], bytes[5]]),
62            csum_start: u16::from_le_bytes([bytes[6], bytes[7]]),
63            csum_offset: u16::from_le_bytes([bytes[8], bytes[9]]),
64            num_buffers: u16::from_le_bytes([bytes[10], bytes[11]]),
65        })
66    }
67
68    /// Converts to bytes.
69    #[must_use]
70    pub fn to_bytes(&self) -> [u8; Self::SIZE] {
71        let mut bytes = [0u8; Self::SIZE];
72        bytes[0] = self.flags;
73        bytes[1] = self.gso_type;
74        bytes[2..4].copy_from_slice(&self.hdr_len.to_le_bytes());
75        bytes[4..6].copy_from_slice(&self.gso_size.to_le_bytes());
76        bytes[6..8].copy_from_slice(&self.csum_start.to_le_bytes());
77        bytes[8..10].copy_from_slice(&self.csum_offset.to_le_bytes());
78        bytes[10..12].copy_from_slice(&self.num_buffers.to_le_bytes());
79        bytes
80    }
81}
82
83/// Network packet.
84#[derive(Debug, Clone)]
85pub struct NetPacket {
86    /// `VirtIO` header.
87    pub header: VirtioNetHeader,
88    /// Packet data (Ethernet frame).
89    pub data: Vec<u8>,
90}
91
92impl NetPacket {
93    /// Creates a new packet.
94    #[must_use]
95    pub fn new(data: Vec<u8>) -> Self {
96        Self {
97            header: VirtioNetHeader::new(),
98            data,
99        }
100    }
101
102    /// Returns the total size including header.
103    #[must_use]
104    pub fn total_size(&self) -> usize {
105        VirtioNetHeader::SIZE + self.data.len()
106    }
107}
108
109#[cfg(test)]
110mod tests {
111    use super::*;
112
113    #[test]
114    fn test_header_size() {
115        assert_eq!(VirtioNetHeader::SIZE, 12);
116    }
117
118    #[test]
119    fn test_header_constants() {
120        assert_eq!(VirtioNetHeader::GSO_NONE, 0);
121        assert_eq!(VirtioNetHeader::GSO_TCPV4, 1);
122        assert_eq!(VirtioNetHeader::GSO_UDP, 3);
123        assert_eq!(VirtioNetHeader::GSO_TCPV6, 4);
124        assert_eq!(VirtioNetHeader::GSO_ECN, 0x80);
125        assert_eq!(VirtioNetHeader::FLAG_NEEDS_CSUM, 1);
126        assert_eq!(VirtioNetHeader::FLAG_DATA_VALID, 2);
127    }
128
129    #[test]
130    fn test_header_new() {
131        let header = VirtioNetHeader::new();
132        assert_eq!(header.flags, 0);
133        assert_eq!(header.gso_type, 0);
134        assert_eq!(header.hdr_len, 0);
135        assert_eq!(header.gso_size, 0);
136        assert_eq!(header.csum_start, 0);
137        assert_eq!(header.csum_offset, 0);
138        assert_eq!(header.num_buffers, 0);
139    }
140
141    #[test]
142    fn test_header_serialization() {
143        let header = VirtioNetHeader {
144            flags: 1,
145            gso_type: 2,
146            hdr_len: 0x1234,
147            gso_size: 0x5678,
148            csum_start: 0x9ABC,
149            csum_offset: 0xDEF0,
150            num_buffers: 0x1111,
151        };
152
153        let bytes = header.to_bytes();
154        let parsed = VirtioNetHeader::from_bytes(&bytes).unwrap();
155
156        assert_eq!(parsed.flags, header.flags);
157        assert_eq!(parsed.gso_type, header.gso_type);
158        assert_eq!(parsed.hdr_len, header.hdr_len);
159        assert_eq!(parsed.gso_size, header.gso_size);
160        assert_eq!(parsed.csum_start, header.csum_start);
161        assert_eq!(parsed.csum_offset, header.csum_offset);
162        assert_eq!(parsed.num_buffers, header.num_buffers);
163    }
164
165    #[test]
166    fn test_header_from_bytes_too_small() {
167        let bytes = [0u8; 11];
168        assert!(VirtioNetHeader::from_bytes(&bytes).is_none());
169    }
170
171    #[test]
172    fn test_header_from_bytes_exact_size() {
173        let bytes = [0u8; 12];
174        assert!(VirtioNetHeader::from_bytes(&bytes).is_some());
175    }
176
177    #[test]
178    fn test_header_from_bytes_larger() {
179        let bytes = [0u8; 100];
180        let header = VirtioNetHeader::from_bytes(&bytes).unwrap();
181        assert_eq!(header.flags, 0);
182    }
183
184    #[test]
185    fn test_header_endianness() {
186        let header = VirtioNetHeader {
187            flags: 0,
188            gso_type: 0,
189            hdr_len: 0x0102,
190            gso_size: 0,
191            csum_start: 0,
192            csum_offset: 0,
193            num_buffers: 0,
194        };
195
196        let bytes = header.to_bytes();
197        // Little-endian: 0x0102 should be stored as [0x02, 0x01]
198        assert_eq!(bytes[2], 0x02);
199        assert_eq!(bytes[3], 0x01);
200    }
201
202    #[test]
203    fn test_packet_new() {
204        let data = vec![0xAA, 0xBB, 0xCC];
205        let packet = NetPacket::new(data.clone());
206
207        assert_eq!(packet.data, data);
208        assert_eq!(packet.header.flags, 0);
209    }
210
211    #[test]
212    fn test_packet_total_size() {
213        let data = vec![0u8; 100];
214        let packet = NetPacket::new(data);
215
216        assert_eq!(packet.total_size(), VirtioNetHeader::SIZE + 100);
217    }
218
219    #[test]
220    fn test_packet_empty() {
221        let packet = NetPacket::new(vec![]);
222        assert_eq!(packet.total_size(), VirtioNetHeader::SIZE);
223        assert!(packet.data.is_empty());
224    }
225
226    #[test]
227    fn test_packet_large() {
228        let data = vec![0u8; 9000];
229        let packet = NetPacket::new(data);
230        assert_eq!(packet.total_size(), VirtioNetHeader::SIZE + 9000);
231    }
232
233    #[test]
234    #[allow(clippy::clone_on_copy)]
235    fn test_header_clone_copy() {
236        let header = VirtioNetHeader {
237            flags: 1,
238            gso_type: 2,
239            hdr_len: 3,
240            gso_size: 4,
241            csum_start: 5,
242            csum_offset: 6,
243            num_buffers: 7,
244        };
245
246        let cloned = header.clone();
247        let copied = header;
248
249        assert_eq!(cloned.flags, 1);
250        assert_eq!(copied.flags, 1);
251    }
252
253    #[test]
254    fn test_packet_clone() {
255        let packet = NetPacket {
256            header: VirtioNetHeader::new(),
257            data: vec![1, 2, 3],
258        };
259
260        let cloned = packet.clone();
261        assert_eq!(cloned.data, packet.data);
262    }
263}