diff --git a/sys/windivert.c b/sys/windivert.c index 34c10c1..4f5fca6 100644 --- a/sys/windivert.c +++ b/sys/windivert.c @@ -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 -