use std::fmt::{self, Formatter};
use zerocopy::byteorder::{BigEndian, U16};
use zerocopy::{FromBytes, IntoBytes, Unaligned};
use crate::packet::protocol::EtherProto;
use crate::packet::{HeaderParser, PacketHeader};
#[repr(C, packed)]
#[derive(
FromBytes, IntoBytes, Unaligned, Debug, Clone, Copy, zerocopy::KnownLayout, zerocopy::Immutable,
)]
pub struct GreHeader {
flags_version: U16<BigEndian>,
protocol_type: U16<BigEndian>,
}
impl GreHeader {
pub const FLAG_CHECKSUM: u16 = 0x8000; pub const FLAG_ROUTING: u16 = 0x4000; pub const FLAG_KEY: u16 = 0x2000; pub const FLAG_SEQUENCE: u16 = 0x1000; pub const FLAG_STRICT_ROUTE: u16 = 0x0800; pub const FLAG_ACK: u16 = 0x0080;
pub const VERSION_MASK: u16 = 0x0007; pub const RECUR_MASK: u16 = 0x0700; pub const FLAGS_MASK: u16 = 0x00F8;
pub const VERSION_0: u16 = 0x0000; pub const VERSION_1: u16 = 0x0001;
#[allow(unused)]
const NAME: &'static str = "GreHeader";
#[inline]
pub fn flags_version(&self) -> u16 {
self.flags_version.get()
}
#[inline]
pub fn version(&self) -> u8 {
(self.flags_version() & Self::VERSION_MASK) as u8
}
#[inline]
pub fn protocol_type(&self) -> EtherProto {
self.protocol_type.get().into()
}
#[inline]
pub fn has_checksum(&self) -> bool {
self.flags_version() & Self::FLAG_CHECKSUM != 0
}
#[inline]
pub fn has_routing(&self) -> bool {
self.flags_version() & Self::FLAG_ROUTING != 0
}
#[inline]
pub fn has_key(&self) -> bool {
self.flags_version() & Self::FLAG_KEY != 0
}
#[inline]
pub fn has_sequence(&self) -> bool {
self.flags_version() & Self::FLAG_SEQUENCE != 0
}
#[inline]
pub fn has_strict_route(&self) -> bool {
self.flags_version() & Self::FLAG_STRICT_ROUTE != 0
}
#[inline]
pub fn has_ack(&self) -> bool {
self.flags_version() & Self::FLAG_ACK != 0
}
#[inline]
pub fn recursion_control(&self) -> u8 {
((self.flags_version() & Self::RECUR_MASK) >> 8) as u8
}
#[inline]
fn is_valid(&self) -> bool {
let version = self.version();
if version > 1 {
return false;
}
if version == 0 {
let reserved = self.flags_version() & Self::FLAGS_MASK;
if reserved != 0 {
return false;
}
}
true
}
#[inline]
pub fn header_length(&self) -> usize {
let mut len = Self::FIXED_LEN;
if self.has_checksum() || self.has_routing() {
len += 4;
}
if self.has_key() {
len += 4;
}
if self.has_sequence() {
len += 4;
}
if self.version() == 1 && self.has_ack() {
len += 4;
}
len
}
pub fn flags_string(&self) -> String {
let mut flags = Vec::new();
if self.has_checksum() {
flags.push("C");
}
if self.has_routing() {
flags.push("R");
}
if self.has_key() {
flags.push("K");
}
if self.has_sequence() {
flags.push("S");
}
if self.has_strict_route() {
flags.push("s");
}
if self.has_ack() {
flags.push("A");
}
if flags.is_empty() {
"none".to_string()
} else {
flags.join("")
}
}
}
#[derive(Debug, Clone)]
pub struct GreHeaderOpt<'a> {
pub header: &'a GreHeader,
pub raw_options: &'a [u8],
}
impl<'a> GreHeaderOpt<'a> {
pub fn checksum(&self) -> Option<u16> {
if !self.header.has_checksum() && !self.header.has_routing() {
return None;
}
if self.raw_options.len() < 2 {
return None;
}
Some(u16::from_be_bytes([
self.raw_options[0],
self.raw_options[1],
]))
}
pub fn offset(&self) -> Option<u16> {
if !self.header.has_checksum() && !self.header.has_routing() {
return None;
}
if self.raw_options.len() < 4 {
return None;
}
Some(u16::from_be_bytes([
self.raw_options[2],
self.raw_options[3],
]))
}
pub fn key(&self) -> Option<u32> {
if !self.header.has_key() {
return None;
}
let mut offset = 0;
if self.header.has_checksum() || self.header.has_routing() {
offset += 4;
}
if self.raw_options.len() < offset + 4 {
return None;
}
let key_bytes = &self.raw_options[offset..offset + 4];
Some(u32::from_be_bytes([
key_bytes[0],
key_bytes[1],
key_bytes[2],
key_bytes[3],
]))
}
pub fn sequence_number(&self) -> Option<u32> {
if !self.header.has_sequence() {
return None;
}
let mut offset = 0;
if self.header.has_checksum() || self.header.has_routing() {
offset += 4;
}
if self.header.has_key() {
offset += 4;
}
if self.raw_options.len() < offset + 4 {
return None;
}
let seq_bytes = &self.raw_options[offset..offset + 4];
Some(u32::from_be_bytes([
seq_bytes[0],
seq_bytes[1],
seq_bytes[2],
seq_bytes[3],
]))
}
pub fn acknowledgment_number(&self) -> Option<u32> {
if self.header.version() != 1 || !self.header.has_ack() {
return None;
}
let mut offset = 0;
if self.header.has_checksum() || self.header.has_routing() {
offset += 4;
}
if self.header.has_key() {
offset += 4;
}
if self.header.has_sequence() {
offset += 4;
}
if self.raw_options.len() < offset + 4 {
return None;
}
let ack_bytes = &self.raw_options[offset..offset + 4];
Some(u32::from_be_bytes([
ack_bytes[0],
ack_bytes[1],
ack_bytes[2],
ack_bytes[3],
]))
}
}
impl std::ops::Deref for GreHeaderOpt<'_> {
type Target = GreHeader;
#[inline]
fn deref(&self) -> &Self::Target {
self.header
}
}
impl PacketHeader for GreHeader {
const NAME: &'static str = "GreHeader";
type InnerType = EtherProto;
#[inline]
fn inner_type(&self) -> Self::InnerType {
self.protocol_type()
}
#[inline]
fn total_len(&self, _buf: &[u8]) -> usize {
self.header_length()
}
#[inline]
fn is_valid(&self) -> bool {
self.is_valid()
}
}
impl HeaderParser for GreHeader {
type Output<'a> = GreHeaderOpt<'a>;
#[inline]
fn into_view<'a>(header: &'a Self, raw_options: &'a [u8]) -> Self::Output<'a> {
GreHeaderOpt {
header,
raw_options,
}
}
}
impl fmt::Display for GreHeader {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(
f,
"GRE v{} proto={}(0x{:04x}) flags={}",
self.version(),
self.protocol_type(),
self.protocol_type().0,
self.flags_string()
)
}
}
impl fmt::Display for GreHeaderOpt<'_> {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(
f,
"GRE v{} proto={} flags={}",
self.version(),
self.protocol_type(),
self.flags_string()
)?;
if let Some(key) = self.key() {
write!(f, " key={}", key)?;
}
if let Some(seq) = self.sequence_number() {
write!(f, " seq={}", seq)?;
}
if let Some(ack) = self.acknowledgment_number() {
write!(f, " ack={}", ack)?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_gre_header_size() {
assert_eq!(std::mem::size_of::<GreHeader>(), 4);
assert_eq!(GreHeader::FIXED_LEN, 4);
}
#[test]
fn test_gre_basic_header() {
let header = GreHeader {
flags_version: U16::new(0x0000), protocol_type: U16::new(0x0800), };
assert_eq!(header.version(), 0);
assert_eq!(header.protocol_type(), EtherProto::IPV4);
assert!(!header.has_checksum());
assert!(!header.has_key());
assert!(!header.has_sequence());
assert!(header.is_valid());
assert_eq!(header.header_length(), 4);
}
#[test]
fn test_gre_with_key() {
let header = GreHeader {
flags_version: U16::new(0x2000), protocol_type: U16::new(0x0800), };
assert!(header.has_key());
assert!(!header.has_checksum());
assert!(!header.has_sequence());
assert_eq!(header.header_length(), 8); }
#[test]
fn test_gre_with_sequence() {
let header = GreHeader {
flags_version: U16::new(0x1000), protocol_type: U16::new(0x0800), };
assert!(header.has_sequence());
assert!(!header.has_checksum());
assert!(!header.has_key());
assert_eq!(header.header_length(), 8); }
#[test]
fn test_gre_with_checksum() {
let header = GreHeader {
flags_version: U16::new(0x8000), protocol_type: U16::new(0x0800), };
assert!(header.has_checksum());
assert!(!header.has_key());
assert!(!header.has_sequence());
assert_eq!(header.header_length(), 8); }
#[test]
fn test_gre_all_flags() {
let header = GreHeader {
flags_version: U16::new(0xB000), protocol_type: U16::new(0x0800), };
assert!(header.has_checksum());
assert!(header.has_key());
assert!(header.has_sequence());
assert_eq!(header.header_length(), 16); }
#[test]
fn test_gre_version_validation() {
let header_v0 = GreHeader {
flags_version: U16::new(0x0000),
protocol_type: U16::new(0x0800),
};
assert!(header_v0.is_valid());
assert_eq!(header_v0.version(), 0);
let header_v1 = GreHeader {
flags_version: U16::new(0x0001),
protocol_type: U16::new(0x880B), };
assert!(header_v1.is_valid());
assert_eq!(header_v1.version(), 1);
let header_invalid = GreHeader {
flags_version: U16::new(0x0002), protocol_type: U16::new(0x0800),
};
assert!(!header_invalid.is_valid());
}
#[test]
fn test_gre_parsing_basic() {
let mut packet = Vec::new();
packet.extend_from_slice(&0x0000u16.to_be_bytes()); packet.extend_from_slice(&0x0800u16.to_be_bytes());
packet.extend_from_slice(b"payload");
let result = GreHeader::from_bytes(&packet);
assert!(result.is_ok());
let (header, payload) = result.unwrap();
assert_eq!(header.version(), 0);
assert_eq!(header.protocol_type(), EtherProto::IPV4);
assert_eq!(payload, b"payload");
}
#[test]
fn test_gre_parsing_with_key() {
let mut packet = Vec::new();
packet.extend_from_slice(&0x2000u16.to_be_bytes()); packet.extend_from_slice(&0x0800u16.to_be_bytes()); packet.extend_from_slice(&0x12345678u32.to_be_bytes());
packet.extend_from_slice(b"test");
let result = GreHeader::from_bytes(&packet);
assert!(result.is_ok());
let (header, payload) = result.unwrap();
assert!(header.has_key());
assert_eq!(header.key().unwrap(), 0x12345678);
assert_eq!(payload, b"test");
}
#[test]
fn test_gre_parsing_with_sequence() {
let mut packet = Vec::new();
packet.extend_from_slice(&0x1000u16.to_be_bytes()); packet.extend_from_slice(&0x0800u16.to_be_bytes()); packet.extend_from_slice(&0x00000042u32.to_be_bytes());
let result = GreHeader::from_bytes(&packet);
assert!(result.is_ok());
let (header, _) = result.unwrap();
assert!(header.has_sequence());
assert_eq!(header.sequence_number().unwrap(), 0x42);
}
#[test]
fn test_gre_parsing_with_checksum() {
let mut packet = Vec::new();
packet.extend_from_slice(&0x8000u16.to_be_bytes()); packet.extend_from_slice(&0x0800u16.to_be_bytes()); packet.extend_from_slice(&0xABCDu16.to_be_bytes()); packet.extend_from_slice(&0x0000u16.to_be_bytes());
let result = GreHeader::from_bytes(&packet);
assert!(result.is_ok());
let (header, _) = result.unwrap();
assert!(header.has_checksum());
assert_eq!(header.checksum().unwrap(), 0xABCD);
}
#[test]
fn test_gre_parsing_all_options() {
let mut packet = Vec::new();
packet.extend_from_slice(&0xB000u16.to_be_bytes()); packet.extend_from_slice(&0x0800u16.to_be_bytes()); packet.extend_from_slice(&0x1234u16.to_be_bytes()); packet.extend_from_slice(&0x0000u16.to_be_bytes()); packet.extend_from_slice(&0xDEADBEEFu32.to_be_bytes()); packet.extend_from_slice(&0x00000100u32.to_be_bytes());
let result = GreHeader::from_bytes(&packet);
assert!(result.is_ok());
let (header, _) = result.unwrap();
assert!(header.has_checksum());
assert!(header.has_key());
assert!(header.has_sequence());
assert_eq!(header.checksum().unwrap(), 0x1234);
assert_eq!(header.key().unwrap(), 0xDEADBEEF);
assert_eq!(header.sequence_number().unwrap(), 0x100);
}
#[test]
fn test_gre_parsing_too_small() {
let packet = vec![0u8; 3];
let result = GreHeader::from_bytes(&packet);
assert!(result.is_err());
}
#[test]
fn test_gre_flags_string() {
let header1 = GreHeader {
flags_version: U16::new(0x0000),
protocol_type: U16::new(0x0800),
};
assert_eq!(header1.flags_string(), "none");
let header2 = GreHeader {
flags_version: U16::new(0x8000), protocol_type: U16::new(0x0800),
};
assert_eq!(header2.flags_string(), "C");
let header3 = GreHeader {
flags_version: U16::new(0xB000), protocol_type: U16::new(0x0800),
};
assert_eq!(header3.flags_string(), "CKS");
}
#[test]
fn test_gre_nvgre_scenario() {
let mut packet = Vec::new();
packet.extend_from_slice(&0x2000u16.to_be_bytes()); packet.extend_from_slice(&0x6558u16.to_be_bytes()); packet.extend_from_slice(&0x00010001u32.to_be_bytes());
let result = GreHeader::from_bytes(&packet);
assert!(result.is_ok());
let (header, _) = result.unwrap();
assert_eq!(header.protocol_type(), EtherProto::TEB);
assert!(header.has_key());
assert_eq!(header.key().unwrap(), 0x00010001);
}
#[test]
fn test_gre_enhanced_version_1() {
let mut packet = Vec::new();
packet.extend_from_slice(&0x3081u16.to_be_bytes()); packet.extend_from_slice(&0x880Bu16.to_be_bytes()); packet.extend_from_slice(&0x0004u16.to_be_bytes()); packet.extend_from_slice(&0x0001u16.to_be_bytes()); packet.extend_from_slice(&0x00000001u32.to_be_bytes()); packet.extend_from_slice(&0x00000000u32.to_be_bytes());
let result = GreHeader::from_bytes(&packet);
assert!(result.is_ok());
let (header, _) = result.unwrap();
assert_eq!(header.version(), 1);
assert!(header.has_key());
assert!(header.has_sequence());
assert!(header.has_ack());
}
#[test]
fn test_gre_protocol_types() {
let protocols = vec![
(EtherProto::IPV4, "IPv4"),
(EtherProto::IPV6, "IPv6"),
(EtherProto::TEB, "TEB"),
(EtherProto::PPP_MP, "PPP"),
(EtherProto::MPLS_UC, "MPLS unicast"),
];
for (proto_type, _name) in protocols {
let header = GreHeader {
flags_version: U16::new(0x0000),
protocol_type: U16::new(proto_type.0.get()),
};
assert_eq!(header.protocol_type(), proto_type);
assert!(header.is_valid());
}
}
#[test]
fn test_gre_header_length_calculation() {
let h1 = GreHeader {
flags_version: U16::new(0x0000),
protocol_type: U16::new(0x0800),
};
assert_eq!(h1.header_length(), 4);
let h2 = GreHeader {
flags_version: U16::new(0x8000),
protocol_type: U16::new(0x0800),
};
assert_eq!(h2.header_length(), 8);
let h3 = GreHeader {
flags_version: U16::new(0x2000),
protocol_type: U16::new(0x0800),
};
assert_eq!(h3.header_length(), 8);
let h4 = GreHeader {
flags_version: U16::new(0x1000),
protocol_type: U16::new(0x0800),
};
assert_eq!(h4.header_length(), 8);
let h5 = GreHeader {
flags_version: U16::new(0xA000),
protocol_type: U16::new(0x0800),
};
assert_eq!(h5.header_length(), 12);
let h6 = GreHeader {
flags_version: U16::new(0xB000),
protocol_type: U16::new(0x0800),
};
assert_eq!(h6.header_length(), 16);
}
}