Ensure offloaded checksums are zero.

WinDivert now treats all offloaded checksum fields as zero:
- WinDivertRecv() will zero all offloaded checksum fields.
- Filter matching will treat offloaded fields as zero.
This commit is contained in:
basil00
2015-07-25 12:05:23 +08:00
parent 3bcf1ae7a0
commit 3fcb692478
+163 -17
View File
@@ -193,11 +193,17 @@ WDF_DECLARE_CONTEXT_TYPE_WITH_NAME(req_context_s, windivert_req_context_get);
* WinDivert packet structure.
*/
#define WINDIVERT_PACKET_SIZE (sizeof(struct packet_s))
#define WINDIVERT_IP_CHECKSUM 0x01
#define WINDIVERT_TCP_CHECKSUM 0x02
#define WINDIVERT_UDP_CHECKSUM 0x04
#define WINDIVERT_ALL_CHECKSUMS \
(WINDIVERT_IP_CHECKSUM | WINDIVERT_TCP_CHECKSUM | WINDIVERT_UDP_CHECKSUM)
struct packet_s
{
LIST_ENTRY entry; // Entry for queue.
PNET_BUFFER_LIST net_buffer_list; // Clone of the net buffer list.
size_t data_len; // Length of `data'.
UINT8 checksums; // Which checksums are valid.
UINT8 direction; // Packet direction.
UINT32 if_idx; // Interface index.
UINT32 sub_if_idx; // Sub-interface index.
@@ -390,8 +396,7 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx,
IN OUT void *data, const FWPS_FILTER0 *filter, IN UINT64 flow_context,
OUT FWPS_CLASSIFY_OUT0 *result);
static BOOL windivert_queue_packet(context_t context, PNET_BUFFER buffer,
PNET_BUFFER_LIST buffers, UINT8 direction, UINT32 if_idx,
UINT32 sub_if_idx, BOOL isloopback);
UINT8 direction, UINT32 if_idx, UINT32 sub_if_idx, UINT8 checksums);
static BOOL windivert_reinject_packet(context_t context, UINT8 direction,
BOOL isipv4, UINT32 if_idx, UINT32 sub_if_idx, UINT32 priority,
PNET_BUFFER buffer);
@@ -400,7 +405,10 @@ static void NTAPI windivert_reinject_complete(VOID *context,
static void windivert_free_packet(packet_t packet);
static UINT8 windivert_skip_headers(UINT8 proto, UINT8 **header, size_t *len);
static BOOL windivert_filter(PNET_BUFFER buffer, UINT32 if_idx,
UINT32 sub_if_idx, BOOL outbound, BOOL isipv4, filter_t filter);
UINT32 sub_if_idx, BOOL outbound, BOOL isipv4, UINT8 checksums,
filter_t filter);
static void windivert_zero_checksums(void *header, size_t len,
UINT8 checksums);
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,
@@ -1384,6 +1392,9 @@ static void windivert_read_service(context_t context)
addr->Direction = packet->direction;
}
// Zero the IP/TCP/UDP checksums here (if required).
windivert_zero_checksums(dst, dst_len, packet->checksums);
status = STATUS_SUCCESS;
windivert_read_service_complete:
@@ -2119,6 +2130,8 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx,
UINT32 priority;
PNET_BUFFER_LIST buffers;
PNET_BUFFER buffer, buffer_fst, buffer_itr;
NDIS_TCP_IP_CHECKSUM_NET_BUFFER_LIST_INFO checksums_info;
UINT8 checksums;
BOOL outbound, queued;
context_t context;
packet_t packet;
@@ -2168,6 +2181,28 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx,
}
}
/*
* Determine which checksum fields are present or not.
*/
if (isloopback)
{
// Loopback packets appear to have bogus checksums, so do not trust.
checksums = 0;
}
else if (direction == WINDIVERT_DIRECTION_OUTBOUND)
{
checksums_info.Value = NET_BUFFER_LIST_INFO(buffers,
TcpIpChecksumNetBufferListInfo);
checksums =
(isipv4? 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;
}
/*
* This code is complicated by the fact the a single NET_BUFFER_LIST
* may contain several NET_BUFFER structures. Each NET_BUFFER needs to
@@ -2184,7 +2219,7 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx,
do
{
if (windivert_filter(buffer_fst, if_idx, sub_if_idx, outbound,
isipv4, context->filter))
isipv4, checksums, context->filter))
{
break;
}
@@ -2220,8 +2255,8 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx,
queued = FALSE;
if ((context->flags & WINDIVERT_FLAG_DROP) == 0)
{
if (!windivert_queue_packet(context, buffer_itr, buffers, direction,
if_idx, sub_if_idx, isloopback))
if (!windivert_queue_packet(context, buffer_itr, direction, if_idx,
sub_if_idx, checksums))
{
goto windivert_classify_callout_exit;
}
@@ -2233,12 +2268,12 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx,
while (buffer_itr != NULL)
{
if (windivert_filter(buffer_itr, if_idx, sub_if_idx, outbound,
isipv4, context->filter))
isipv4, checksums, context->filter))
{
if ((context->flags & WINDIVERT_FLAG_DROP) == 0)
{
if (!windivert_queue_packet(context, buffer_itr, buffers,
direction, if_idx, sub_if_idx, isloopback))
if (!windivert_queue_packet(context, buffer_itr, direction,
if_idx, sub_if_idx, checksums))
{
goto windivert_classify_callout_exit;
}
@@ -2296,11 +2331,9 @@ windivert_classify_callout_exit:
* Queue a NET_BUFFER.
*/
static BOOL windivert_queue_packet(context_t context, PNET_BUFFER buffer,
PNET_BUFFER_LIST buffers, UINT8 direction, UINT32 if_idx,
UINT32 sub_if_idx, BOOL isloopback)
UINT8 direction, UINT32 if_idx, UINT32 sub_if_idx, UINT8 checksums)
{
KLOCK_QUEUE_HANDLE lock_handle;
NDIS_TCP_IP_CHECKSUM_NET_BUFFER_LIST_INFO checksum_info;
PVOID data;
PLIST_ENTRY entry;
packet_t packet;
@@ -2325,7 +2358,7 @@ static BOOL windivert_queue_packet(context_t context, PNET_BUFFER buffer,
{
RtlCopyMemory(packet->data, data, data_len);
}
packet->checksums = checksums;
packet->direction = direction;
packet->if_idx = if_idx;
packet->sub_if_idx = sub_if_idx;
@@ -2533,11 +2566,100 @@ static UINT8 windivert_skip_headers(UINT8 proto, UINT8 **header, size_t *len)
}
}
/*
* Zero invalid checksum fields.
*/
static void windivert_zero_checksums(void *header, size_t len,
UINT8 checksums)
{
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;
if (checksums == WINDIVERT_ALL_CHECKSUMS)
{
return;
}
if (len < sizeof(struct iphdr))
{
return;
}
switch (ip_header->Version)
{
case 4:
ip_header_len = ip_header->HdrLength*sizeof(UINT32);
if (len < ip_header_len)
{
return;
}
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;
case 6:
if (len < sizeof(struct ipv6hdr))
{
return;
}
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;
default:
return;
}
switch (proto)
{
case IPPROTO_TCP:
if ((checksums & WINDIVERT_TCP_CHECKSUM) != 0)
{
return;
}
tcp_header = (struct tcphdr *)trans_header;
if (trans_len < sizeof(struct tcphdr))
{
return;
}
tcp_header->Checksum = 0;
break;
case IPPROTO_UDP:
if ((checksums & WINDIVERT_UDP_CHECKSUM) != 0)
{
return;
}
udp_header = (struct udphdr *)trans_header;
if (trans_len < sizeof(struct udphdr))
{
return;
}
udp_header->Checksum = 0;
break;
}
}
/*
* Checks if the given packet is of interest.
*/
static BOOL windivert_filter(PNET_BUFFER buffer, UINT32 if_idx,
UINT32 sub_if_idx, BOOL outbound, BOOL isipv4, filter_t filter)
UINT32 sub_if_idx, BOOL outbound, BOOL isipv4, UINT8 checksums,
filter_t filter)
{
size_t tot_len, ip_header_len;
struct iphdr *ip_header = NULL;
@@ -2792,7 +2914,15 @@ static BOOL windivert_filter(PNET_BUFFER buffer, UINT32 if_idx,
field[0] = (UINT32)ip_header->Protocol;
break;
case WINDIVERT_FILTER_FIELD_IP_CHECKSUM:
field[0] = (UINT32)RtlUshortByteSwap(ip_header->Checksum);
if ((checksums & WINDIVERT_IP_CHECKSUM) != 0)
{
field[0] =
(UINT32)RtlUshortByteSwap(ip_header->Checksum);
}
else
{
field[0] = 0;
}
break;
case WINDIVERT_FILTER_FIELD_IP_SRCADDR:
field[0] = (UINT32)RtlUlongByteSwap(ip_header->SrcAddr);
@@ -2900,7 +3030,15 @@ static BOOL windivert_filter(PNET_BUFFER buffer, UINT32 if_idx,
field[0] = (UINT32)RtlUshortByteSwap(tcp_header->Window);
break;
case WINDIVERT_FILTER_FIELD_TCP_CHECKSUM:
field[0] = (UINT32)RtlUshortByteSwap(tcp_header->Checksum);
if ((checksums & WINDIVERT_TCP_CHECKSUM) != 0)
{
field[0] =
(UINT32)RtlUshortByteSwap(tcp_header->Checksum);
}
else
{
field[0] = 0;
}
break;
case WINDIVERT_FILTER_FIELD_TCP_URGPTR:
field[0] = (UINT32)RtlUshortByteSwap(tcp_header->UrgPtr);
@@ -2919,7 +3057,15 @@ static BOOL windivert_filter(PNET_BUFFER buffer, UINT32 if_idx,
field[0] = (UINT32)RtlUshortByteSwap(udp_header->Length);
break;
case WINDIVERT_FILTER_FIELD_UDP_CHECKSUM:
field[0] = (UINT32)RtlUshortByteSwap(udp_header->Checksum);
if ((checksums & WINDIVERT_UDP_CHECKSUM) != 0)
{
field[0] =
(UINT32)RtlUshortByteSwap(udp_header->Checksum);
}
else
{
field[0] = 0;
}
break;
case WINDIVERT_FILTER_FIELD_UDP_PAYLOADLENGTH:
field[0] = (UINT32)(tot_len - ip_header_len -