From aea3a3a858d3c5e1ad3545eb473de98354575dc3 Mon Sep 17 00:00:00 2001 From: basil00 Date: Thu, 9 Nov 2017 22:11:13 +0800 Subject: [PATCH] Add support for pseudo IP/TCP/UDP checksums. Most NIC cards support checksum offloading, meaning that it is not necessary to calculate the full IP/TCP/UDP checksums for injected packets. WINDIVERT_ADDRESS has been extended to include 3 extra flags that indicate if the packet has full or pseudo checksums. This is a WIP. --- dll/windivert_helper.c | 168 +++++++++++------------ examples/netdump/netdump.c | 4 - examples/netfilter/netfilter.c | 20 ++- examples/passthru/passthru.c | 6 +- examples/streamdump/streamdump.c | 4 +- examples/webfilter/webfilter.c | 8 +- include/windivert.h | 16 +-- sys/windivert.c | 220 +++++++++++++------------------ 8 files changed, 194 insertions(+), 252 deletions(-) diff --git a/dll/windivert_helper.c b/dll/windivert_helper.c index 44d8b45..44e3cb8 100644 --- a/dll/windivert_helper.c +++ b/dll/windivert_helper.c @@ -1,6 +1,6 @@ /* * windivert_helper.c - * (C) 2016, all rights reserved, + * (C) 2017, all rights reserved, * * This program is free software: you can redistribute it and/or modify * it under the terms of the GNU Lesser General Public License as published by @@ -173,6 +173,8 @@ typedef UINT64 ERROR; #define IS_ERROR(err) \ (GET_CODE(err) != WINDIVERT_ERROR_NONE) +#define MAX(a, b) ((a) > (b)? (a): (b)) + /* * Compiler memory pool: */ @@ -188,10 +190,9 @@ typedef struct POOL */ static PEXPR WinDivertParseFilter(PPOOL pool, TOKEN *toks, UINT *i, INT depth, BOOL and); -static void WinDivertInitPseudoHeader(PWINDIVERT_IPHDR ip_header, - PWINDIVERT_PSEUDOHDR pseudo_header, UINT8 protocol, UINT len); -static void WinDivertInitPseudoHeaderV6(PWINDIVERT_IPV6HDR ipv6_header, - PWINDIVERT_PSEUDOV6HDR pseudov6_header, UINT8 protocol, UINT len); +static UINT16 WinDivertInitPseudoHeader(PWINDIVERT_IPHDR ip_header, + PWINDIVERT_IPV6HDR ipv6_header, UINT8 protocol, UINT len, + void *pseudo_header); static UINT16 WinDivertHelperCalcChecksum(PVOID pseudo_header, UINT16 pseudo_header_len, PVOID data, UINT len); @@ -412,10 +413,11 @@ WinDivertHelperParsePacketExit: * Calculate IPv4/IPv6/ICMP/ICMPv6/TCP/UDP checksums. */ extern UINT WinDivertHelperCalcChecksums(PVOID pPacket, UINT packetLen, - UINT64 flags) + PWINDIVERT_ADDRESS pAddr, UINT64 flags) { - WINDIVERT_PSEUDOHDR pseudo_header; - WINDIVERT_PSEUDOV6HDR pseudov6_header; + UINT8 pseudo_header[ + MAX(sizeof(WINDIVERT_PSEUDOHDR), sizeof(WINDIVERT_PSEUDOV6HDR))]; + UINT16 pseudo_header_len; PWINDIVERT_IPHDR ip_header; PWINDIVERT_IPV6HDR ipv6_header; PWINDIVERT_ICMPHDR icmp_header; @@ -424,36 +426,25 @@ extern UINT WinDivertHelperCalcChecksums(PVOID pPacket, UINT packetLen, PWINDIVERT_UDPHDR udp_header; UINT payload_len, checksum_len; UINT count = 0; - UINT64 flags_all = - (WINDIVERT_HELPER_NO_IP_CHECKSUM | - WINDIVERT_HELPER_NO_ICMP_CHECKSUM | - WINDIVERT_HELPER_NO_ICMPV6_CHECKSUM | - WINDIVERT_HELPER_NO_TCP_CHECKSUM | - WINDIVERT_HELPER_NO_UDP_CHECKSUM); - - if ((flags & flags_all) == flags_all) - { - return 0; - } WinDivertHelperParsePacket(pPacket, packetLen, &ip_header, &ipv6_header, &icmp_header, &icmpv6_header, &tcp_header, &udp_header, NULL, &payload_len); - if (ip_header != NULL && !(flags & WINDIVERT_HELPER_NO_IP_CHECKSUM) && - (!(flags & WINDIVERT_HELPER_NO_REPLACE) || ip_header->Checksum == 0)) + if (ip_header != NULL && !(flags & WINDIVERT_HELPER_NO_IP_CHECKSUM)) { ip_header->Checksum = 0; - ip_header->Checksum = WinDivertHelperCalcChecksum(NULL, 0, - ip_header, ip_header->HdrLength*sizeof(UINT32)); + if (pAddr == NULL || pAddr->IPv4Checksum != 0) + { + ip_header->Checksum = WinDivertHelperCalcChecksum(NULL, 0, + ip_header, ip_header->HdrLength*sizeof(UINT32)); + } count++; } if (icmp_header != NULL) { - if ((flags & WINDIVERT_HELPER_NO_ICMP_CHECKSUM) || - ((flags & WINDIVERT_HELPER_NO_REPLACE) && - icmp_header->Checksum != 0)) + if ((flags & WINDIVERT_HELPER_NO_ICMP_CHECKSUM) != 0) { return count; } @@ -466,46 +457,49 @@ extern UINT WinDivertHelperCalcChecksums(PVOID pPacket, UINT packetLen, if (icmpv6_header != NULL) { - if ((flags & WINDIVERT_HELPER_NO_ICMPV6_CHECKSUM) || - ((flags & WINDIVERT_HELPER_NO_REPLACE) && - icmpv6_header->Checksum != 0)) + if ((flags & WINDIVERT_HELPER_NO_ICMPV6_CHECKSUM) != 0) { return count; } checksum_len = payload_len + sizeof(WINDIVERT_ICMPV6HDR); - WinDivertInitPseudoHeaderV6(ipv6_header, &pseudov6_header, - IPPROTO_ICMPV6, checksum_len); + pseudo_header_len = WinDivertInitPseudoHeader(NULL, ipv6_header, + IPPROTO_ICMPV6, checksum_len, pseudo_header); icmpv6_header->Checksum = 0; - icmpv6_header->Checksum = WinDivertHelperCalcChecksum(&pseudov6_header, - sizeof(pseudov6_header), icmpv6_header, checksum_len); + icmpv6_header->Checksum = WinDivertHelperCalcChecksum(pseudo_header, + pseudo_header_len, icmpv6_header, checksum_len); count++; return count; } if (tcp_header != NULL) { - if ((flags & WINDIVERT_HELPER_NO_TCP_CHECKSUM) || - ((flags & WINDIVERT_HELPER_NO_REPLACE) && - tcp_header->Checksum != 0)) + if ((flags & WINDIVERT_HELPER_NO_TCP_CHECKSUM) != 0) { return count; } - checksum_len = payload_len + tcp_header->HdrLength*sizeof(UINT32); - if (ip_header != NULL) + if (pAddr == NULL || pAddr->TCPChecksum != 0) { - WinDivertInitPseudoHeader(ip_header, &pseudo_header, IPPROTO_TCP, - checksum_len); + // Full TCP checksum + checksum_len = payload_len + tcp_header->HdrLength*sizeof(UINT32); + pseudo_header_len = WinDivertInitPseudoHeader(ip_header, + ipv6_header, IPPROTO_TCP, checksum_len, pseudo_header); tcp_header->Checksum = 0; - tcp_header->Checksum = WinDivertHelperCalcChecksum(&pseudo_header, - sizeof(pseudo_header), tcp_header, checksum_len); + tcp_header->Checksum = WinDivertHelperCalcChecksum( + pseudo_header, pseudo_header_len, tcp_header, checksum_len); + } + else if (pAddr->Direction == WINDIVERT_DIRECTION_OUTBOUND) + { + // Pseudo TCP checksum + checksum_len = payload_len + tcp_header->HdrLength*sizeof(UINT32); + pseudo_header_len = WinDivertInitPseudoHeader(ip_header, + ipv6_header, IPPROTO_TCP, checksum_len, pseudo_header); + tcp_header->Checksum = ~WinDivertHelperCalcChecksum( + pseudo_header, pseudo_header_len, NULL, 0); } else { - WinDivertInitPseudoHeaderV6(ipv6_header, &pseudov6_header, - IPPROTO_TCP, checksum_len); + // Don't care checksum tcp_header->Checksum = 0; - tcp_header->Checksum = WinDivertHelperCalcChecksum(&pseudov6_header, - sizeof(pseudov6_header), tcp_header, checksum_len); } count++; return count; @@ -513,32 +507,37 @@ extern UINT WinDivertHelperCalcChecksums(PVOID pPacket, UINT packetLen, if (udp_header != NULL) { - if ((flags & WINDIVERT_HELPER_NO_UDP_CHECKSUM) || - ((flags & WINDIVERT_HELPER_NO_REPLACE) && - udp_header->Checksum != 0)) + if ((flags & WINDIVERT_HELPER_NO_UDP_CHECKSUM) != 0) { return count; } - checksum_len = payload_len + sizeof(WINDIVERT_UDPHDR); - if (ip_header != NULL) + if (pAddr == NULL || pAddr->UDPChecksum != 0) { - WinDivertInitPseudoHeader(ip_header, &pseudo_header, IPPROTO_UDP, - checksum_len); + // Full UDP checksum + checksum_len = payload_len + sizeof(WINDIVERT_UDPHDR); + pseudo_header_len = WinDivertInitPseudoHeader(ip_header, + ipv6_header, IPPROTO_UDP, checksum_len, pseudo_header); udp_header->Checksum = 0; - udp_header->Checksum = WinDivertHelperCalcChecksum(&pseudo_header, - sizeof(pseudo_header), udp_header, checksum_len); + udp_header->Checksum = WinDivertHelperCalcChecksum( + pseudo_header, pseudo_header_len, udp_header, checksum_len); if (udp_header->Checksum == 0) { udp_header->Checksum = 0xFFFF; } } + else if (pAddr->Direction == WINDIVERT_DIRECTION_OUTBOUND) + { + // Pseudo UDP checksum + checksum_len = payload_len + sizeof(WINDIVERT_UDPHDR); + pseudo_header_len = WinDivertInitPseudoHeader(ip_header, + ipv6_header, IPPROTO_UDP, checksum_len, pseudo_header); + udp_header->Checksum = ~WinDivertHelperCalcChecksum( + pseudo_header, pseudo_header_len, NULL, 0); + } else { - WinDivertInitPseudoHeaderV6(ipv6_header, &pseudov6_header, - IPPROTO_UDP, checksum_len); + // Don't care checksum udp_header->Checksum = 0; - udp_header->Checksum = WinDivertHelperCalcChecksum(&pseudov6_header, - sizeof(pseudov6_header), udp_header, checksum_len); } count++; } @@ -546,31 +545,36 @@ extern UINT WinDivertHelperCalcChecksums(PVOID pPacket, UINT packetLen, } /* - * Initialize the IP pseudo header. + * Initialize the IP/IPv6 pseudo header. */ -static void WinDivertInitPseudoHeader(PWINDIVERT_IPHDR ip_header, - PWINDIVERT_PSEUDOHDR pseudo_header, UINT8 protocol, UINT len) +static UINT16 WinDivertInitPseudoHeader(PWINDIVERT_IPHDR ip_header, + PWINDIVERT_IPV6HDR ipv6_header, UINT8 protocol, UINT len, + void *pseudo_header) { - pseudo_header->SrcAddr = ip_header->SrcAddr; - pseudo_header->DstAddr = ip_header->DstAddr; - pseudo_header->Zero = 0; - pseudo_header->Protocol = protocol; - pseudo_header->Length = htons((UINT16)len); -} - -/* - * Initialize the IPv6 pseudo header. - */ -static void WinDivertInitPseudoHeaderV6(PWINDIVERT_IPV6HDR ipv6_header, - PWINDIVERT_PSEUDOV6HDR pseudov6_header, UINT8 protocol, UINT len) -{ - memcpy(pseudov6_header->SrcAddr, ipv6_header->SrcAddr, - sizeof(pseudov6_header->SrcAddr)); - memcpy(pseudov6_header->DstAddr, ipv6_header->DstAddr, - sizeof(pseudov6_header->DstAddr)); - pseudov6_header->Length = htonl((UINT32)len); - pseudov6_header->NextHdr = protocol; - pseudov6_header->Zero = 0; + if (ip_header != NULL) + { + PWINDIVERT_PSEUDOHDR pseudo_header_v4 = + (PWINDIVERT_PSEUDOHDR)pseudo_header; + pseudo_header_v4->SrcAddr = ip_header->SrcAddr; + pseudo_header_v4->DstAddr = ip_header->DstAddr; + pseudo_header_v4->Zero = 0; + pseudo_header_v4->Protocol = protocol; + pseudo_header_v4->Length = htons((UINT16)len); + return sizeof(WINDIVERT_PSEUDOHDR); + } + else + { + PWINDIVERT_PSEUDOV6HDR pseudo_header_v6 = + (PWINDIVERT_PSEUDOV6HDR)pseudo_header; + memcpy(pseudo_header_v6->SrcAddr, ipv6_header->SrcAddr, + sizeof(pseudo_header_v6->SrcAddr)); + memcpy(pseudo_header_v6->DstAddr, ipv6_header->DstAddr, + sizeof(pseudo_header_v6->DstAddr)); + pseudo_header_v6->Length = htonl((UINT32)len); + pseudo_header_v6->NextHdr = protocol; + pseudo_header_v6->Zero = 0; + return sizeof(WINDIVERT_PSEUDOV6HDR); + } } /* diff --git a/examples/netdump/netdump.c b/examples/netdump/netdump.c index e06dd55..0bb609d 100644 --- a/examples/netdump/netdump.c +++ b/examples/netdump/netdump.c @@ -124,10 +124,6 @@ int __cdecl main(int argc, char **argv) continue; } - // Calculate checksums. - WinDivertHelperCalcChecksums(packet, packet_len, - WINDIVERT_HELPER_NO_REPLACE); - // Print info about the matching packet. WinDivertHelperParsePacket(packet, packet_len, &ip_header, &ipv6_header, &icmp_header, &icmpv6_header, &tcp_header, diff --git a/examples/netfilter/netfilter.c b/examples/netfilter/netfilter.c index 6862eab..1897040 100644 --- a/examples/netfilter/netfilter.c +++ b/examples/netfilter/netfilter.c @@ -1,6 +1,6 @@ /* * netfilter.c - * (C) 2016, all rights reserved, + * (C) 2017, all rights reserved, * * This program is free software: you can redistribute it and/or modify * it under the terms of the GNU Lesser General Public License as published by @@ -270,11 +270,10 @@ int __cdecl main(int argc, char **argv) htonl(ntohl(tcp_header->SeqNum) + 1): htonl(ntohl(tcp_header->SeqNum) + payload_len)); - WinDivertHelperCalcChecksums((PVOID)reset, sizeof(TCPPACKET), - 0); - memcpy(&send_addr, &recv_addr, sizeof(send_addr)); send_addr.Direction = !recv_addr.Direction; + WinDivertHelperCalcChecksums((PVOID)reset, sizeof(TCPPACKET), + &send_addr, 0); if (!WinDivertSend(handle, (PVOID)reset, sizeof(TCPPACKET), &send_addr, NULL)) { @@ -298,11 +297,10 @@ int __cdecl main(int argc, char **argv) htonl(ntohl(tcp_header->SeqNum) + 1): htonl(ntohl(tcp_header->SeqNum) + payload_len)); - WinDivertHelperCalcChecksums((PVOID)resetv6, - sizeof(TCPV6PACKET), 0); - memcpy(&send_addr, &recv_addr, sizeof(send_addr)); send_addr.Direction = !recv_addr.Direction; + WinDivertHelperCalcChecksums((PVOID)resetv6, + sizeof(TCPV6PACKET), &send_addr, 0); if (!WinDivertSend(handle, (PVOID)resetv6, sizeof(TCPV6PACKET), &send_addr, NULL)) { @@ -325,10 +323,10 @@ int __cdecl main(int argc, char **argv) dnr->ip.SrcAddr = ip_header->DstAddr; dnr->ip.DstAddr = ip_header->SrcAddr; - WinDivertHelperCalcChecksums((PVOID)dnr, icmp_length, 0); - memcpy(&send_addr, &recv_addr, sizeof(send_addr)); send_addr.Direction = !recv_addr.Direction; + WinDivertHelperCalcChecksums((PVOID)dnr, icmp_length, + &send_addr, 0); if (!WinDivertSend(handle, (PVOID)dnr, icmp_length, &send_addr, NULL)) { @@ -348,10 +346,10 @@ int __cdecl main(int argc, char **argv) memcpy(dnrv6->ipv6.DstAddr, ipv6_header->SrcAddr, sizeof(dnrv6->ipv6.DstAddr)); - WinDivertHelperCalcChecksums((PVOID)dnrv6, icmpv6_length, 0); - memcpy(&send_addr, &recv_addr, sizeof(send_addr)); send_addr.Direction = !recv_addr.Direction; + WinDivertHelperCalcChecksums((PVOID)dnrv6, icmpv6_length, + &send_addr, 0); if (!WinDivertSend(handle, (PVOID)dnrv6, icmpv6_length, &send_addr, NULL)) { diff --git a/examples/passthru/passthru.c b/examples/passthru/passthru.c index 7817ba3..b86a69b 100644 --- a/examples/passthru/passthru.c +++ b/examples/passthru/passthru.c @@ -1,6 +1,6 @@ /* * passthru.c - * (C) 2016, all rights reserved, + * (C) 2017, all rights reserved, * * This program is free software: you can redistribute it and/or modify * it under the terms of the GNU Lesser General Public License as published by @@ -108,10 +108,6 @@ static DWORD passthru(LPVOID arg) } // Re-inject the matching packet. - // NOTE: Only use the WINDIVERT_HELPER_NO_REPLACE flag if the packet - // was not modified. - WinDivertHelperCalcChecksums(packet, packet_len, - WINDIVERT_HELPER_NO_REPLACE); if (!WinDivertSend(handle, packet, packet_len, &addr, NULL)) { fprintf(stderr, "warning: failed to reinject packet (%d)\n", diff --git a/examples/streamdump/streamdump.c b/examples/streamdump/streamdump.c index c0f84aa..884342b 100644 --- a/examples/streamdump/streamdump.c +++ b/examples/streamdump/streamdump.c @@ -1,6 +1,6 @@ /* * streamdump.c - * (C) 2016 basil, all rights reserved, + * (C) 2017 basil, all rights reserved, * * This program is free software: you can redistribute it and/or modify * it under the terms of the GNU Lesser General Public License as published by @@ -278,7 +278,7 @@ read_failed: break; } - WinDivertHelperCalcChecksums(packet, packet_len, 0); + WinDivertHelperCalcChecksums(packet, packet_len, &addr, 0); poverlapped = (OVERLAPPED *)malloc(sizeof(OVERLAPPED)); if (poverlapped == NULL) { diff --git a/examples/webfilter/webfilter.c b/examples/webfilter/webfilter.c index 89e5b19..dba66e0 100644 --- a/examples/webfilter/webfilter.c +++ b/examples/webfilter/webfilter.c @@ -189,8 +189,6 @@ int __cdecl main(int argc, char **argv) !BlackListPayloadMatch(blacklist, payload, (UINT16)payload_len)) { // Packet does not match the blacklist; simply reinject it. - WinDivertHelperCalcChecksums(packet, packet_len, - WINDIVERT_HELPER_NO_REPLACE); if (!WinDivertSend(handle, packet, packet_len, &addr, NULL)) { fprintf(stderr, "warning: failed to reinject packet (%d)\n", @@ -210,7 +208,7 @@ int __cdecl main(int argc, char **argv) reset->tcp.DstPort = htons(80); reset->tcp.SeqNum = tcp_header->SeqNum; reset->tcp.AckNum = tcp_header->AckNum; - WinDivertHelperCalcChecksums((PVOID)reset, sizeof(PACKET), 0); + WinDivertHelperCalcChecksums((PVOID)reset, sizeof(PACKET), &addr, 0); if (!WinDivertSend(handle, (PVOID)reset, sizeof(PACKET), &addr, NULL)) { fprintf(stderr, "warning: failed to send reset packet (%d)\n", @@ -224,8 +222,8 @@ int __cdecl main(int argc, char **argv) blockpage->header.tcp.SeqNum = tcp_header->AckNum; blockpage->header.tcp.AckNum = htonl(ntohl(tcp_header->SeqNum) + payload_len); - WinDivertHelperCalcChecksums((PVOID)blockpage, blockpage_len, 0); addr.Direction = !addr.Direction; // Reverse direction. + WinDivertHelperCalcChecksums((PVOID)blockpage, blockpage_len, &addr, 0); if (!WinDivertSend(handle, (PVOID)blockpage, blockpage_len, &addr, NULL)) { @@ -243,7 +241,7 @@ int __cdecl main(int argc, char **argv) htonl(ntohl(tcp_header->AckNum) + sizeof(block_data) - 1); finish->tcp.AckNum = htonl(ntohl(tcp_header->SeqNum) + payload_len); - WinDivertHelperCalcChecksums((PVOID)finish, sizeof(PACKET), 0); + WinDivertHelperCalcChecksums((PVOID)finish, sizeof(PACKET), &addr, 0); if (!WinDivertSend(handle, (PVOID)finish, sizeof(PACKET), &addr, NULL)) { fprintf(stderr, "warning: failed to send finish packet (%d)\n", diff --git a/include/windivert.h b/include/windivert.h index fe173f8..cb38122 100644 --- a/include/windivert.h +++ b/include/windivert.h @@ -62,7 +62,10 @@ typedef struct UINT32 SubIfIdx; /* Packet's sub-interface index. */ UINT8 Direction:1; /* Packet's direction. */ UINT8 Loopback:1; /* Packet is loopback? */ - UINT8 Reserved:6; + UINT8 IPv4Checksum:1; /* Packet has full IPv4 checksum? */ + UINT8 TCPChecksum:1; /* Packet has full TCP checksum? */ + UINT8 UDPChecksum:1; /* Packet has full UDP checksum? */ + UINT8 Reserved:2; } WINDIVERT_ADDRESS, *PWINDIVERT_ADDRESS; #define WINDIVERT_DIRECTION_OUTBOUND 0 @@ -318,7 +321,6 @@ typedef struct #define WINDIVERT_HELPER_NO_ICMPV6_CHECKSUM 4 #define WINDIVERT_HELPER_NO_TCP_CHECKSUM 8 #define WINDIVERT_HELPER_NO_UDP_CHECKSUM 16 -#define WINDIVERT_HELPER_NO_REPLACE 2048 /* * Parse IPv4/IPv6/ICMP/ICMPv6/TCP/UDP headers from a raw packet. @@ -355,6 +357,7 @@ extern WINDIVERTEXPORT BOOL WinDivertHelperParseIPv6Address( extern WINDIVERTEXPORT UINT WinDivertHelperCalcChecksums( __inout PVOID pPacket, __in UINT packetLen, + __in_opt PWINDIVERT_ADDRESS pAddr, __in UINT64 flags); /* @@ -376,15 +379,6 @@ extern WINDIVERTEXPORT BOOL WinDivertHelperEvalFilter( __in UINT packetLen, __in PWINDIVERT_ADDRESS pAddr); -/****************************************************************************/ -/* WINDIVERT LEGACY API */ -/****************************************************************************/ - -/* - * Deprecated API: - */ -#define WINDIVERT_FLAG_NO_CHECKSUM 0 - #endif /* WINDIVERT_KERNEL */ #ifdef __cplusplus diff --git a/sys/windivert.c b/sys/windivert.c index 904faee..8b13c51 100644 --- a/sys/windivert.c +++ b/sys/windivert.c @@ -238,18 +238,26 @@ struct packet_s }; typedef struct packet_s *packet_t; -#if 0 /* - * WinDivert address definition. + * IPv4/IPv6 pseudo headers. */ -struct windivert_addr_s +typedef struct { - UINT32 IfIdx; - UINT32 SubIfIdx; - UINT8 Direction; -}; -typedef struct windivert_addr_s *windivert_addr_t; -#endif + UINT32 SrcAddr; + UINT32 DstAddr; + UINT8 Zero; + UINT8 Protocol; + UINT16 Length; +} WINDIVERT_PSEUDOHDR, *PWINDIVERT_PSEUDOHDR; + +typedef struct +{ + UINT32 SrcAddr[4]; + UINT32 DstAddr[4]; + UINT32 Length; + UINT32 Zero:24; + UINT32 NextHdr:8; +} WINDIVERT_PSEUDOV6HDR, *PWINDIVERT_PSEUDOV6HDR; /* * Header definitions. @@ -441,9 +449,9 @@ static UINT8 windivert_skip_headers(UINT8 proto, UINT8 **header, size_t *len); static int windivert_big_num_compare(const UINT32 *a, const UINT32 *b); static BOOL windivert_filter(PNET_BUFFER buffer, UINT32 if_idx, UINT32 sub_if_idx, BOOL outbound, BOOL isipv4, BOOL hop, BOOL loopback, - UINT8 checksums, filter_t filter); + filter_t filter); static NTSTATUS windivert_finalize_packet(void *header, size_t len, - BOOL hop, UINT8 checksums); + BOOL hop); static filter_t windivert_filter_compile(windivert_ioctl_filter_t ioctl_filter, size_t ioctl_filter_len); static void windivert_filter_analyze(filter_t filter, BOOL *is_inbound, @@ -1544,11 +1552,14 @@ static void windivert_read_service_request(packet_t packet, addr->SubIfIdx = sub_if_idx; addr->Direction = direction; addr->Loopback = (loopback? 1: 0); + addr->IPv4Checksum = ((checksums & WINDIVERT_IP_CHECKSUM) != 0? 1: 0); + addr->TCPChecksum = ((checksums & WINDIVERT_TCP_CHECKSUM) != 0? 1: 0); + addr->UDPChecksum = ((checksums & WINDIVERT_UDP_CHECKSUM) != 0? 1: 0); addr->Reserved = 0; } // Zero the IP/TCP/UDP checksums and/or decrement the TTL (if required). - status = windivert_finalize_packet(dst, dst_len, hop, checksums); + status = windivert_finalize_packet(dst, dst_len, hop); windivert_read_service_request_exit: if (NT_SUCCESS(status)) @@ -1630,7 +1641,7 @@ static NTSTATUS windivert_write(context_t context, WDFREQUEST request, struct iphdr *ip_header; struct ipv6hdr *ipv6_header; BOOL isipv4; - UINT8 layer; + UINT8 layer, checksums; UINT32 priority; UINT64 flags; HANDLE handle, compl_handle; @@ -1740,6 +1751,32 @@ windivert_write_bad_packet: flags = context->flags; KeReleaseInStackQueuedSpinLock(&lock_handle); + if (layer != WINDIVERT_LAYER_NETWORK_FORWARD) + { + checksums_info.Value = NET_BUFFER_LIST_INFO(buffers, + TcpIpChecksumNetBufferListInfo); + if (addr->Direction == WINDIVERT_DIRECTION_OUTBOUND) + { + checksums_info.Transmit.TcpChecksum = + (addr->TCPChecksum != 0? 0: 1); + checksums_info.Transmit.UdpChecksum = + (addr->UDPChecksum != 0? 0: 1); + checksums_info.Transmit.IpHeaderChecksum = + (addr->IPv4Checksum != 0? 0: 1); + } + else + { + checksums_info.Receive.TcpChecksumSucceeded = + (addr->TCPChecksum != 0? 0: 1); + checksums_info.Receive.UdpChecksumSucceeded = + (addr->UDPChecksum != 0? 0: 1); + checksums_info.Receive.IpChecksumSucceeded = + (addr->IPv4Checksum != 0? 0: 1); + } + NET_BUFFER_LIST_INFO(buffers, TcpIpChecksumNetBufferListInfo) = + checksums_info.Value; + } + handle = (isipv4? inject_handle: injectv6_handle); compl_handle = ((flags & WINDIVERT_FLAG_DEBUG) != 0? (HANDLE)request: NULL); if (layer == WINDIVERT_LAYER_NETWORK_FORWARD) @@ -2372,7 +2409,7 @@ static void windivert_classify_callout(context_t context, IN UINT8 direction, PNET_BUFFER_LIST buffers; PNET_BUFFER buffer, buffer_fst, buffer_itr; NDIS_TCP_IP_CHECKSUM_NET_BUFFER_LIST_INFO checksums_info; - UINT8 checksums; + UINT8 layer, checksums; BOOL outbound, hop; WDFOBJECT object; work_t work; @@ -2414,6 +2451,7 @@ static void windivert_classify_callout(context_t context, IN UINT8 direction, result->actionType = FWP_ACTION_CONTINUE; return; } + layer = context->layer; priority = context->priority; filter = context->filter; object = (WDFOBJECT)context->object; @@ -2449,27 +2487,33 @@ static void windivert_classify_callout(context_t context, IN UINT8 direction, } // Determine which checksum fields are present or not. + checksums_info.Value = NET_BUFFER_LIST_INFO(buffers, + TcpIpChecksumNetBufferListInfo); if (loopback) { - // Loopback packets appear to have bogus checksums, so do not trust. checksums = 0; } + else if (layer == WINDIVERT_LAYER_NETWORK_FORWARD) + { + checksums = WINDIVERT_ALL_CHECKSUMS; + } else if (direction == WINDIVERT_DIRECTION_OUTBOUND) { - checksums_info.Value = NET_BUFFER_LIST_INFO(buffers, - TcpIpChecksumNetBufferListInfo); checksums = - (isipv4? 0: WINDIVERT_IP_CHECKSUM) | + (checksums_info.Transmit.IpHeaderChecksum? 0: + WINDIVERT_IP_CHECKSUM) | (checksums_info.Transmit.TcpChecksum? 0: WINDIVERT_TCP_CHECKSUM) | (checksums_info.Transmit.UdpChecksum? 0: WINDIVERT_UDP_CHECKSUM); } else { - checksums = WINDIVERT_ALL_CHECKSUMS; - } - if (hop && isipv4) - { - checksums &= ~WINDIVERT_IP_CHECKSUM; + checksums = + (checksums_info.Receive.IpChecksumSucceeded? 0: + WINDIVERT_IP_CHECKSUM) | + (checksums_info.Receive.TcpChecksumSucceeded? 0: + WINDIVERT_TCP_CHECKSUM) | + (checksums_info.Receive.UdpChecksumSucceeded? 0: + WINDIVERT_UDP_CHECKSUM); } timestamp = KeQueryPerformanceCounter(NULL).QuadPart; @@ -2505,7 +2549,7 @@ static void windivert_classify_callout(context_t context, IN UINT8 direction, do { BOOL match = windivert_filter(buffer_fst, if_idx, sub_if_idx, outbound, - isipv4, hop, loopback, checksums, filter); + isipv4, hop, loopback, filter); if (match) { break; @@ -2677,7 +2721,7 @@ VOID windivert_worker(IN WDFWORKITEM item) { match = windivert_filter(buffer_itr, work->if_idx, work->sub_if_idx, outbound, work->is_ipv4, work->hop, - work->loopback, work->checksums, filter); + work->loopback, filter); if (match) { ok = windivert_queue_packet(context, buffer_itr, @@ -3092,19 +3136,14 @@ static UINT8 windivert_skip_headers(UINT8 proto, UINT8 **header, size_t *len) /* * Zero the IP/TCP/UDP checksums and/or decrement the TTL (if required) */ -static NTSTATUS windivert_finalize_packet(void *header, size_t len, - BOOL hop, UINT8 checksums) +static NTSTATUS windivert_finalize_packet(void *header, size_t len, BOOL hop) { struct iphdr *ip_header = (struct iphdr *)header; struct ipv6hdr *ipv6_header = (struct ipv6hdr *)header; - size_t ip_header_len, trans_len; - void *trans_header; - struct tcphdr *tcp_header; - struct udphdr *udp_header; - UINT8 proto; + size_t ip_header_len; NTSTATUS status = STATUS_SUCCESS; - if (checksums == WINDIVERT_ALL_CHECKSUMS && !hop) + if (!hop) { return status; } @@ -3116,93 +3155,34 @@ static NTSTATUS windivert_finalize_packet(void *header, size_t len, switch (ip_header->Version) { case 4: - ip_header_len = ip_header->HdrLength*sizeof(UINT32); - if (len < ip_header_len) + if (ip_header->TTL <= 1) { - return status; + status = STATUS_HOPLIMIT_EXCEEDED; } - - if (hop) + if (ip_header->TTL != 0) { - if (ip_header->TTL <= 1) - { - status = STATUS_HOPLIMIT_EXCEEDED; - } - if (ip_header->TTL != 0) - { - ip_header->TTL--; - } + ip_header->TTL--; } - - if ((checksums & WINDIVERT_IP_CHECKSUM) != 0) - { - ip_header->Checksum = 0; - } - proto = ip_header->Protocol; - trans_len = len - ip_header_len; - trans_header = (UINT8 *)ip_header + ip_header_len; - break; + return status; case 6: if (len < sizeof(struct ipv6hdr)) { return status; } - - if (hop) + if (ipv6_header->HopLimit <= 1) { - if (ipv6_header->HopLimit <= 1) - { - status = STATUS_HOPLIMIT_EXCEEDED; - } - if (ipv6_header->HopLimit != 0) - { - ipv6_header->HopLimit--; - } + status = STATUS_HOPLIMIT_EXCEEDED; } - - trans_len = len - sizeof(struct ipv6hdr); - trans_header = (UINT8 *)(ipv6_header + 1); - - // Skip extension headers: - proto = windivert_skip_headers(ipv6_header->NextHdr, - (UINT8 **)&trans_header, &trans_len); - break; + if (ipv6_header->HopLimit != 0) + { + ipv6_header->HopLimit--; + } + return status; default: return status; } - - switch (proto) - { - case IPPROTO_TCP: - if ((checksums & WINDIVERT_TCP_CHECKSUM) != 0) - { - return status; - } - tcp_header = (struct tcphdr *)trans_header; - if (trans_len < sizeof(struct tcphdr)) - { - return status; - } - tcp_header->Checksum = 0; - break; - - case IPPROTO_UDP: - if ((checksums & WINDIVERT_UDP_CHECKSUM) != 0) - { - return status; - } - udp_header = (struct udphdr *)trans_header; - if (trans_len < sizeof(struct udphdr)) - { - return status; - } - udp_header->Checksum = 0; - break; - } - - return status; } /* @@ -3250,7 +3230,7 @@ static int windivert_big_num_compare(const UINT32 *a, const UINT32 *b) */ static BOOL windivert_filter(PNET_BUFFER buffer, UINT32 if_idx, UINT32 sub_if_idx, BOOL outbound, BOOL isipv4, BOOL hop, BOOL loopback, - UINT8 checksums, filter_t filter) + filter_t filter) { size_t tot_len, ip_header_len; struct iphdr *ip_header = NULL; @@ -3513,15 +3493,7 @@ static BOOL windivert_filter(PNET_BUFFER buffer, UINT32 if_idx, field[0] = (UINT32)ip_header->Protocol; break; case WINDIVERT_FILTER_FIELD_IP_CHECKSUM: - if ((checksums & WINDIVERT_IP_CHECKSUM) != 0) - { - field[0] = - (UINT32)RtlUshortByteSwap(ip_header->Checksum); - } - else - { - field[0] = 0; - } + field[0] = (UINT32)RtlUshortByteSwap(ip_header->Checksum); break; case WINDIVERT_FILTER_FIELD_IP_SRCADDR: field[0] = (UINT32)RtlUlongByteSwap(ip_header->SrcAddr); @@ -3633,15 +3605,7 @@ static BOOL windivert_filter(PNET_BUFFER buffer, UINT32 if_idx, field[0] = (UINT32)RtlUshortByteSwap(tcp_header->Window); break; case WINDIVERT_FILTER_FIELD_TCP_CHECKSUM: - if ((checksums & WINDIVERT_TCP_CHECKSUM) != 0) - { - field[0] = - (UINT32)RtlUshortByteSwap(tcp_header->Checksum); - } - else - { - field[0] = 0; - } + field[0] = (UINT32)RtlUshortByteSwap(tcp_header->Checksum); break; case WINDIVERT_FILTER_FIELD_TCP_URGPTR: field[0] = (UINT32)RtlUshortByteSwap(tcp_header->UrgPtr); @@ -3660,15 +3624,7 @@ static BOOL windivert_filter(PNET_BUFFER buffer, UINT32 if_idx, field[0] = (UINT32)RtlUshortByteSwap(udp_header->Length); break; case WINDIVERT_FILTER_FIELD_UDP_CHECKSUM: - if ((checksums & WINDIVERT_UDP_CHECKSUM) != 0) - { - field[0] = - (UINT32)RtlUshortByteSwap(udp_header->Checksum); - } - else - { - field[0] = 0; - } + field[0] = (UINT32)RtlUshortByteSwap(udp_header->Checksum); break; case WINDIVERT_FILTER_FIELD_UDP_PAYLOADLENGTH: field[0] = (UINT32)(tot_len - ip_header_len -