diff --git a/common/src/ebpf/parsing.rs b/common/src/ebpf/parsing.rs index 861ead3..9b1ca8b 100644 --- a/common/src/ebpf/parsing.rs +++ b/common/src/ebpf/parsing.rs @@ -1,3 +1,5 @@ +use core::mem::size_of; + use aya_ebpf::helpers::bpf_ktime_get_ns; use network_types::eth::{EthHdr, EtherType}; use network_types::ip::{IpProto, Ipv4Hdr, Ipv6Hdr}; @@ -7,6 +9,7 @@ use network_types::udp::UdpHdr; use crate::define::offset::*; use crate::model::parsed_packet::ParsedPacket; +#[allow(clippy::result_unit_err, clippy::not_unsafe_ptr_arg_deref)] pub fn parse_packet(start: usize, end: usize, target: *mut ParsedPacket) -> Result<(), ()> { unsafe { if start + ETHER_HEADER_END > end { @@ -30,6 +33,8 @@ unsafe fn parse_ipv4_packet(start: usize, end: usize, target: *mut ParsedPacket) unsafe { let ipv4 = &*((start + IPV4_HEADER_START) as *const Ipv4Hdr); + let ipv4_header_len = parse_ipv4_header_len(start, end)?; + let l4_start = IPV4_HEADER_START + ipv4_header_len; let packet_length = (end - start) as u32; let t = &mut *target; @@ -41,12 +46,12 @@ unsafe fn parse_ipv4_packet(start: usize, end: usize, target: *mut ParsedPacket) t.protocol = ipv4.proto; let (src_port, dst_port, tcp_flags, l4_header_len) = match ipv4.proto { - IpProto::Tcp => parse_tcp(start, end, IPV4_TCP_HEADER_START, IPV4_TCP_HEADER_END)?, - IpProto::Udp => parse_udp(start, end, IPV4_UDP_HEADER_START, IPV4_UDP_HEADER_END)?, + IpProto::Tcp => parse_tcp(start, end, l4_start)?, + IpProto::Udp => parse_udp(start, end, l4_start)?, _ => (0, 0, 0, 0), }; - t.payload_length = packet_length.saturating_sub((IPV4_HEADER_END + l4_header_len) as u32); + t.payload_length = packet_length.saturating_sub((l4_start + l4_header_len) as u32); t.src_port = src_port; t.dst_port = dst_port; t.tcp_flags = tcp_flags; @@ -74,8 +79,8 @@ unsafe fn parse_ipv6_packet(start: usize, end: usize, target: *mut ParsedPacket) t.protocol = ipv6.next_hdr; let (src_port, dst_port, tcp_flags, l4_header_len) = match ipv6.next_hdr { - IpProto::Tcp => parse_tcp(start, end, IPV6_TCP_HEADER_START, IPV6_TCP_HEADER_END)?, - IpProto::Udp => parse_udp(start, end, IPV6_UDP_HEADER_START, IPV6_UDP_HEADER_END)?, + IpProto::Tcp => parse_tcp(start, end, IPV6_TCP_HEADER_START)?, + IpProto::Udp => parse_udp(start, end, IPV6_UDP_HEADER_START)?, _ => (0, 0, 0, 0), }; @@ -89,8 +94,31 @@ unsafe fn parse_ipv6_packet(start: usize, end: usize, target: *mut ParsedPacket) } #[inline(always)] -unsafe fn parse_tcp(start: usize, end: usize, tcp_start: usize, tcp_end: usize) -> Result<(u16, u16, u8, usize), ()> { - if start + tcp_end > end { +#[allow(clippy::manual_range_contains)] +unsafe fn parse_ipv4_header_len(start: usize, end: usize) -> Result { + if start + IPV4_HEADER_START + 1 > end { + return Err(()); + } + + let version_ihl = unsafe { *((start + IPV4_HEADER_START) as *const u8) }; + let version = version_ihl >> 4; + let ihl = (version_ihl & 0x0f) as usize; + if version != 4 || ihl < 5 || ihl > 15 { + return Err(()); + } + + let header_len = ihl * 4; + if start + IPV4_HEADER_START + header_len > end { + return Err(()); + } + + Ok(header_len) +} + +#[inline(always)] +#[allow(clippy::manual_range_contains)] +unsafe fn parse_tcp(start: usize, end: usize, tcp_start: usize) -> Result<(u16, u16, u8, usize), ()> { + if start + tcp_start + size_of::() > end { return Err(()); } @@ -115,8 +143,8 @@ unsafe fn parse_tcp(start: usize, end: usize, tcp_start: usize, tcp_end: usize) } #[inline(always)] -unsafe fn parse_udp(start: usize, end: usize, udp_start: usize, udp_end: usize) -> Result<(u16, u16, u8, usize), ()> { - if start + udp_end > end { +unsafe fn parse_udp(start: usize, end: usize, udp_start: usize) -> Result<(u16, u16, u8, usize), ()> { + if start + udp_start + size_of::() > end { return Err(()); }