diff --git a/src/main.c b/src/main.c index 87466d3..f7801d2 100644 --- a/src/main.c +++ b/src/main.c @@ -145,15 +145,16 @@ static __always_inline const __u8 *get_tcp_data(const __u16 EtherType, return NULL; } -static __u16 tcp_checksum_fold(__u32 sum) { +static __always_inline __u16 tcp_checksum_fold(__u32 sum) { sum = (sum & 0xFFFF) + (sum >> 16); sum = (sum & 0xFFFF) + (sum >> 16); sum = (sum & 0xFFFF) + (sum >> 16); return (__u16)~sum; } -static __u32 tcp_checksum_add(const void *data, const void *data_end, - __u32 sum) { +static __always_inline __u32 tcp_checksum_add(const void *data, + const void *data_end, + __u32 sum) { const __u8 *p = data; for (int i = 0; i < 740; i++) { // MTU/2, upper bound for the verifier if (p + 2 > (const __u8 *)data_end) { @@ -227,9 +228,11 @@ int xdp_drop_prog(struct xdp_md *ctx) { *tcp_flags_ptr = tcp_flags | ACK_MASK; + *(__u16 *)(tcp_packet_data + TCP_CHECKSUM_OFFSET) = 0; // must be zero for calculation + const __u8 tcp_data_offset = - bpf_ntohs(((*(__u8 *)(tcp_packet_data + TCP_DATA_OFFSET_OFFSET))) >> - TCP_DATA_OFFSET_SHIFT) * + ((*(__u8 *)(tcp_packet_data + TCP_DATA_OFFSET_OFFSET)) >> + TCP_DATA_OFFSET_SHIFT) * sizeof(__u32); __u16 tcp_checksum; // big endian @@ -246,25 +249,31 @@ int xdp_drop_prog(struct xdp_md *ctx) { { // swap src and dst ip __u32 src_ip_BE = *src_ip_ptr; *src_ip_ptr = *dst_ip_ptr; - *dst_ip_ptr = *src_ip_ptr; + *dst_ip_ptr = src_ip_BE; } *(__u8 *)(ip_packet_data + IPV4_TTL_OFFSET) = DEFAULT_TTL; { // calculate ipv4 header checksum + if ((*ip_packet_data & 0b00001111) != 5) { + return XDP_PASS; + } + if (ip_packet_data + 20 > data_end) { + return XDP_PASS; + } __u16 *ipv4_checksum_ptr = (__u16 *)(ip_packet_data + IPV4_CHECKSUM_OFFSET); *ipv4_checksum_ptr = 0; // must be zero for calculation - __u16 new_checksum = 0; - const __u8 *p = ip_packet_data; - for (int i = 0; i < 30; i++) { - if (p + 2 > data_end || p + 2 > tcp_packet_data) { + __u32 new_checksum = 0; + for (int i = 0; i < 10; i++) { + if (ip_packet_data + i * 2 + 2 > data_end) { break; } - new_checksum += bpf_ntohs(*(const __u16 *)p); - p += 2; + new_checksum += bpf_ntohs(*(__u16 *)(ip_packet_data + i * 2)); } - *ipv4_checksum_ptr = bpf_htons(new_checksum); + new_checksum = (new_checksum & 0xFFFF) + (new_checksum >> 16); + new_checksum = (new_checksum & 0xFFFF) + (new_checksum >> 16); + *ipv4_checksum_ptr = bpf_htons((__u16)new_checksum); } struct tcp_v4_pseudo_header pseudo_header = { @@ -275,6 +284,12 @@ int xdp_drop_prog(struct xdp_md *ctx) { &pseudo_header, (const __u8 *)&pseudo_header + sizeof(pseudo_header), tmp_sum); } else { + if (EtherType != ETHER_TYPE_IPV6) { + return XDP_PASS; + } + if (ip_packet_data + IPV6_HEADER_SIZE > data_end) { + return XDP_PASS; + } __u128 *src_ip_ptr = (__u128 *)(ip_packet_data + IPV6_SRC_ADDR_OFFSET); __u128 *dst_ip_ptr = @@ -283,7 +298,7 @@ int xdp_drop_prog(struct xdp_md *ctx) { { // swap src and dst ip __u128 src_ip_BE = *src_ip_ptr; *src_ip_ptr = *dst_ip_ptr; - *dst_ip_ptr = *src_ip_ptr; + *dst_ip_ptr = src_ip_BE; } *(__u8 *)(ip_packet_data + IPV6_HOP_LIMIT_OFFSET) = DEFAULT_TTL; @@ -300,10 +315,7 @@ int xdp_drop_prog(struct xdp_md *ctx) { (const __u8 *)&pseudo_header + sizeof(pseudo_header), tmp_sum); } - tmp_sum = tcp_checksum_add(tcp_packet_data, - tcp_packet_data + tcp_data_offset, tmp_sum); - tmp_sum = tcp_checksum_add(tcp_packet_data + tcp_data_offset, data_end, - tmp_sum); + tmp_sum = tcp_checksum_add(tcp_packet_data, data_end, tmp_sum); tcp_checksum = bpf_htons(tcp_checksum_fold(tmp_sum)); } *(__u16 *)(tcp_packet_data + TCP_CHECKSUM_OFFSET) = tcp_checksum;