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.
This commit is contained in:
basil00
2017-11-09 22:11:13 +08:00
parent bbf6a34aa6
commit aea3a3a858
8 changed files with 194 additions and 252 deletions
+86 -82
View File
@@ -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);
}
}
/*
-4
View File
@@ -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,
+9 -11
View File
@@ -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))
{
+1 -5
View File
@@ -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",
+2 -2
View File
@@ -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)
{
+3 -5
View File
@@ -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",
+5 -11
View File
@@ -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
+88 -132
View File
@@ -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 -