From 2e1bfa8ca5216273061227b386c9b21fff27d186 Mon Sep 17 00:00:00 2001 From: basil00 Date: Tue, 10 Oct 2017 23:14:37 +0800 Subject: [PATCH] WinDivert driver overhaul. This is a major update designed to modernize the WinDivert driver, including optimizations, design improvements and bug fixes. The new version has not been fully tested and should be considered **UNSTABLE**. - Most of the packet processing is now (almost) fully out-of-band. This is a good since the classify function runs at DISPATCH_LEVEL. The driver will still try to match at least one packet before moving the work out-of-band. - Re-injected non-matching packets are now clones rather than copies. - Queued packets are still copied. This is because the driver should avoid keeping a reference to the original packet for very long. Since we do not trust the user application to handle the packet in a timely fashion, it is better to copy rather than keep a reference. That said, the driver now implements an optimization where it will service a read request immediately if possible (saving 1 packet copy). - SNIFF mode also now works differently. Previously, SNIFF mode would not block the original packet. However, this is problematic under the new design since WinDivert cannot permit the packet and retain a reference to it at the same time. The new version will block & absorb the original packet and re-inject a clone out-of-band. - Packet time management has been replaced. Previously, a timer was used to periodically wake up a function that would sweep away expired packets. The new version explicitly timestamps every packet, and expired packets are cleaned up by the read service routine. - The context->filter is now deallocated in the destroy callback to avoid possible a race condition with the callout function. It is unclear if this is really necessary, however. --- sys/windivert.c | 927 ++++++++++++++++++++++++++++-------------------- 1 file changed, 545 insertions(+), 382 deletions(-) diff --git a/sys/windivert.c b/sys/windivert.c index 44d9641..9603a80 100644 --- a/sys/windivert.c +++ b/sys/windivert.c @@ -35,10 +35,10 @@ EVT_WDF_DRIVER_UNLOAD windivert_unload; EVT_WDF_IO_IN_CALLER_CONTEXT windivert_caller_context; EVT_WDF_IO_QUEUE_IO_DEVICE_CONTROL windivert_ioctl; EVT_WDF_DEVICE_FILE_CREATE windivert_create; -EVT_WDF_TIMER windivert_timer; EVT_WDF_FILE_CLEANUP windivert_cleanup; EVT_WDF_FILE_CLOSE windivert_close; -EVT_WDF_WORKITEM windivert_read_service_work_item; +EVT_WDF_OBJECT_CONTEXT_DESTROY windivert_destroy; +EVT_WDF_WORKITEM windivert_worker; /* * Debugging macros. @@ -60,6 +60,10 @@ static void DEBUG(PCCH format, ...) DbgPrint("WINDIVERT: %s\n", buf); va_end(args); } +#else // DEBUG_ON +#define DEBUG(format, ...) +#endif + static void DEBUG_ERROR(PCCH format, NTSTATUS status, ...) { va_list args; @@ -73,10 +77,6 @@ static void DEBUG_ERROR(PCCH format, NTSTATUS status, ...) DbgPrint("WINDIVERT: *** ERROR ***: (status = %x): %s\n", status, buf); va_end(args); } -#else // DEBUG_ON -#define DEBUG(format, ...) -#define DEBUG_ERROR(format, status, ...) -#endif // DEBUG_ON #define WINDIVERT_TAG 'viDW' @@ -126,14 +126,15 @@ struct context_s context_state_t state; // Context's state. KSPIN_LOCK lock; // Context-wide lock. WDFDEVICE device; // Context's device. + WDFFILEOBJECT object; // Context's parent object. + LIST_ENTRY work_queue; // Work queue. LIST_ENTRY packet_queue; // Packet queue. ULONG packet_queue_length; // Packet queue length. ULONG packet_queue_maxlength; // Packet queue max length. ULONG packet_queue_size; // Packet queue size (in bytes). ULONG packet_queue_maxsize; // Packet queue max size. - WDFTIMER timer; // Packet timer. - UINT timer_timeout; // Packet timeout (in ms). - BOOL timer_ticktock; // Packet timer ticktock. + ULONGLONG packet_queue_maxcounts; // Packet queue max counts. + ULONG packet_queue_maxtime; // Packet queue max time. WDFQUEUE read_queue; // Read queue. WDFWORKITEM workers[WINDIVERT_CONTEXT_MAXWORKERS]; // Read workers. @@ -191,6 +192,25 @@ typedef struct req_context_s req_context_s; typedef struct req_context_s *req_context_t; WDF_DECLARE_CONTEXT_TYPE_WITH_NAME(req_context_s, windivert_req_context_get); +/* + * WinDivert work structure. + */ +struct work_s +{ + LIST_ENTRY entry; // Entry for queue. + PNET_BUFFER_LIST buffers; // Packet list. + PNET_BUFFER buffer; // First matching packet. + UINT advance; // Bytes to retreat/advance. + BOOL is_ipv4; // Is IPv4? + UINT8 checksums; // Which checksums are valid. + UINT8 direction; // Packet direction. + UINT32 if_idx; // Interface index. + UINT32 sub_if_idx; // Sub-interface index. + UINT32 priority; // WinDivert priority. + ULONGLONG timestamp; // Packet timestamp. +}; +typedef struct work_s *work_t; + /* * WinDivert packet structure. */ @@ -208,7 +228,7 @@ struct packet_s UINT8 direction; // Packet direction. UINT32 if_idx; // Interface index. UINT32 sub_if_idx; // Sub-interface index. - BOOL timer_ticktock; // Time-out ticktock. + ULONGLONG timestamp; // Packet timestamp. char data[]; // Packet data. }; typedef struct packet_s *packet_t; @@ -318,6 +338,7 @@ HANDLE injectv6_handle = NULL; NDIS_HANDLE pool_handle = NULL; HANDLE engine_handle = NULL; LONG priority_counter = 0; +static ULONGLONG counts_per_ms = 0; /* * Priorities. @@ -341,7 +362,7 @@ static void windivert_driver_unload(void); extern VOID windivert_ioctl(IN WDFQUEUE queue, IN WDFREQUEST request, IN size_t in_length, IN size_t out_len, IN ULONG code); static NTSTATUS windivert_read(context_t context, WDFREQUEST request); -extern VOID windivert_read_service_work_item(IN WDFWORKITEM item); +extern VOID windivert_worker(IN WDFWORKITEM item); static void windivert_read_service(context_t context); static BOOLEAN windivert_context_verify(context_t context, context_state_t state); @@ -353,9 +374,9 @@ static NTSTATUS windivert_install_callouts(context_t context, BOOL is_inbound, static NTSTATUS windivert_install_callout(context_t context, UINT idx, layer_t layer); static void windivert_uninstall_callouts(context_t context); -extern VOID windivert_timer(IN WDFTIMER timer); extern VOID windivert_cleanup(IN WDFFILEOBJECT object); extern VOID windivert_close(IN WDFFILEOBJECT object); +extern VOID windivert_destroy(IN WDFOBJECT object); extern NTSTATUS windivert_write(context_t context, WDFREQUEST request, windivert_addr_t addr); extern void NTAPI windivert_inject_complete(VOID *context, @@ -392,15 +413,16 @@ static void windivert_classify_forward_network_v6_callout( IN const FWPS_INCOMING_METADATA_VALUES0 *meta_vals, IN OUT void *data, const FWPS_FILTER0 *filter, IN UINT64 flow_context, OUT FWPS_CLASSIFY_OUT0 *result); -static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx, - IN UINT32 sub_if_idx, IN BOOL isipv4, IN BOOL isloopback, - IN OUT void *data, const FWPS_FILTER0 *filter, IN UINT64 flow_context, - OUT FWPS_CLASSIFY_OUT0 *result); +static void windivert_classify_callout(context_t context, IN UINT8 direction, + IN UINT32 if_idx, IN UINT32 sub_if_idx, IN BOOL isipv4, + IN BOOL isloopback, IN UINT advance, IN OUT void *data, + IN UINT64 flow_context, OUT FWPS_CLASSIFY_OUT0 *result); static BOOL windivert_queue_packet(context_t context, PNET_BUFFER buffer, - 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); + UINT8 direction, UINT32 if_idx, UINT32 sub_if_idx, UINT8 checksums, + ULONGLONG timestamp); +static BOOL windivert_reinject_packet(BOOL sniff_mode, BOOL foward, + UINT8 direction, BOOL isipv4, UINT32 if_idx, UINT32 sub_if_idx, + UINT32 priority, PNET_BUFFER_LIST buffers, PNET_BUFFER buffer); static void NTAPI windivert_reinject_complete(VOID *context, NET_BUFFER_LIST *buffers, BOOLEAN dispatch_level); static void windivert_free_packet(packet_t packet); @@ -542,6 +564,7 @@ extern NTSTATUS DriverEntry(IN PDRIVER_OBJECT driver_obj, WDFQUEUE queue; WDF_OBJECT_ATTRIBUTES obj_attrs; NET_BUFFER_LIST_POOL_PARAMETERS pool_params; + LARGE_INTEGER freq; NTSTATUS status; DECLARE_CONST_UNICODE_STRING(device_name, L"\\Device\\" WINDIVERT_DEVICE_NAME); @@ -550,6 +573,11 @@ extern NTSTATUS DriverEntry(IN PDRIVER_OBJECT driver_obj, DEBUG("LOAD: loading WinDivert driver"); + // Initialize timer info. + KeQueryPerformanceCounter(&freq); + counts_per_ms = (ULONGLONG)freq.QuadPart / 1000; + counts_per_ms = (counts_per_ms == 0? 1: counts_per_ms); + // Initialize the layers. layer_inbound_network_ipv4->layer_guid = FWPM_LAYER_INBOUND_IPPACKET_V4; layer_outbound_network_ipv4->layer_guid = FWPM_LAYER_OUTBOUND_IPPACKET_V4; @@ -602,6 +630,9 @@ extern NTSTATUS DriverEntry(IN PDRIVER_OBJECT driver_obj, WDF_FILEOBJECT_CONFIG_INIT(&file_config, windivert_create, windivert_close, windivert_cleanup); WDF_OBJECT_ATTRIBUTES_INIT_CONTEXT_TYPE(&obj_attrs, context_s); + obj_attrs.ExecutionLevel = WdfExecutionLevelPassive; + obj_attrs.SynchronizationScope = WdfSynchronizationScopeNone; + obj_attrs.EvtDestroyCallback = windivert_destroy; WdfDeviceInitSetFileObjectConfig(device_init, &file_config, &obj_attrs); WdfDeviceInitSetIoInCallerContextCallback(device_init, windivert_caller_context); @@ -853,7 +884,6 @@ extern VOID windivert_create(IN WDFDEVICE device, IN WDFREQUEST request, IN WDFFILEOBJECT object) { WDF_IO_QUEUE_CONFIG queue_config; - WDF_TIMER_CONFIG timer_config; WDF_WORKITEM_CONFIG item_config; WDF_OBJECT_ATTRIBUTES obj_attrs; FWPM_SESSION0 session; @@ -867,12 +897,14 @@ extern VOID windivert_create(IN WDFDEVICE device, IN WDFREQUEST request, context->magic = WINDIVERT_CONTEXT_MAGIC; context->state = WINDIVERT_CONTEXT_STATE_OPENING; context->device = device; + context->object = object; context->packet_queue_length = 0; context->packet_queue_maxlength = WINDIVERT_PARAM_QUEUE_LEN_DEFAULT; context->packet_queue_size = 0; context->packet_queue_maxsize = WINDIVERT_PARAM_QUEUE_SIZE_DEFAULT; - context->timer = NULL; - context->timer_timeout = WINDIVERT_PARAM_QUEUE_TIME_DEFAULT; + context->packet_queue_maxcounts = + WINDIVERT_PARAM_QUEUE_TIME_DEFAULT * counts_per_ms; + context->packet_queue_maxtime = WINDIVERT_PARAM_QUEUE_TIME_DEFAULT; context->layer_0 = WINDIVERT_LAYER_DEFAULT; context->layer = WINDIVERT_LAYER_DEFAULT; context->flags_0 = 0; @@ -892,6 +924,7 @@ extern VOID windivert_create(IN WDFDEVICE device, IN WDFREQUEST request, } context->filter_on = FALSE; KeInitializeSpinLock(&context->lock); + InitializeListHead(&context->work_queue); InitializeListHead(&context->packet_queue); for (i = 0; i < WINDIVERT_CONTEXT_MAXLAYERS; i++) { @@ -916,17 +949,7 @@ extern VOID windivert_create(IN WDFDEVICE device, IN WDFREQUEST request, DEBUG_ERROR("failed to create I/O read queue", status); goto windivert_create_exit; } - WDF_TIMER_CONFIG_INIT(&timer_config, windivert_timer); - timer_config.AutomaticSerialization = TRUE; - WDF_OBJECT_ATTRIBUTES_INIT(&obj_attrs); - obj_attrs.ParentObject = (WDFOBJECT)object; - status = WdfTimerCreate(&timer_config, &obj_attrs, &context->timer); - if (!NT_SUCCESS(status)) - { - DEBUG_ERROR("failed to create packet time-out timer", status); - goto windivert_create_exit; - } - WDF_WORKITEM_CONFIG_INIT(&item_config, windivert_read_service_work_item); + WDF_WORKITEM_CONFIG_INIT(&item_config, windivert_worker); item_config.AutomaticSerialization = FALSE; WDF_OBJECT_ATTRIBUTES_INIT(&obj_attrs); obj_attrs.ParentObject = (WDFOBJECT)object; @@ -961,10 +984,6 @@ windivert_create_exit: { WdfObjectDelete(context->read_queue); } - if (context->timer != NULL) - { - WdfObjectDelete(context->timer); - } for (i = 0; i < WINDIVERT_CONTEXT_MAXWORKERS; i++) { if (context->workers[i] != NULL) @@ -1185,54 +1204,6 @@ unregister_callouts: } } -/* - * WinDivert old-packet cleanup routine. - */ -extern VOID windivert_timer(IN WDFTIMER timer) -{ - KLOCK_QUEUE_HANDLE lock_handle; - PLIST_ENTRY entry; - WDFFILEOBJECT object = (WDFFILEOBJECT)WdfTimerGetParentObject(timer); - context_t context = windivert_context_get(object); - packet_t packet; - - if (!windivert_context_verify(context, WINDIVERT_CONTEXT_STATE_OPEN)) - { - return; - } - - // DEBUG("TIMER (context=%p, ticktock=%u)", context, - // context->timer_ticktock); - - // Sweep away old packets. - KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); - while (!IsListEmpty(&context->packet_queue)) - { - entry = RemoveHeadList(&context->packet_queue); - packet = CONTAINING_RECORD(entry, struct packet_s, entry); - if (packet->timer_ticktock == context->timer_ticktock) - { - InsertHeadList(&context->packet_queue, entry); - break; - } - context->packet_queue_length--; - context->packet_queue_size -= packet->data_len; - KeReleaseInStackQueuedSpinLock(&lock_handle); - - // Packet is old, dispose of it. - DEBUG("TIMEOUT (context=%p, packet=%p)", context, packet); - windivert_free_packet(packet); - KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); - } - - KeReleaseInStackQueuedSpinLock(&lock_handle); - context->timer_ticktock = !context->timer_ticktock; - - // Restart the timer. - WdfTimerStart(context->timer, - WDF_REL_TIMEOUT_IN_MS(context->timer_timeout)); -} - /* * Divert cleanup routine. */ @@ -1242,6 +1213,7 @@ extern VOID windivert_cleanup(IN WDFFILEOBJECT object) PLIST_ENTRY entry; UINT i; context_t context = windivert_context_get(object); + work_t work; packet_t packet; NTSTATUS status; @@ -1251,9 +1223,17 @@ extern VOID windivert_cleanup(IN WDFFILEOBJECT object) { return; } - WdfTimerStop(context->timer, TRUE); KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); context->state = WINDIVERT_CONTEXT_STATE_CLOSING; + while (!IsListEmpty(&context->work_queue)) + { + entry = RemoveHeadList(&context->work_queue); + KeReleaseInStackQueuedSpinLock(&lock_handle); + work = CONTAINING_RECORD(entry, struct work_s, entry); + FwpsDereferenceNetBufferList(work->buffers, FALSE); + ExFreePoolWithTag(work, WINDIVERT_TAG); + KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); + } while (!IsListEmpty(&context->packet_queue)) { entry = RemoveHeadList(&context->packet_queue); @@ -1265,7 +1245,6 @@ extern VOID windivert_cleanup(IN WDFFILEOBJECT object) KeReleaseInStackQueuedSpinLock(&lock_handle); WdfIoQueuePurge(context->read_queue, NULL, NULL); WdfObjectDelete(context->read_queue); - WdfObjectDelete(context->timer); for (i = 0; i < WINDIVERT_CONTEXT_MAXWORKERS; i++) { WdfWorkItemFlush(context->workers[i]); @@ -1273,11 +1252,6 @@ extern VOID windivert_cleanup(IN WDFFILEOBJECT object) } windivert_uninstall_callouts(context); FwpmEngineClose0(context->engine_handle); - if (context->filter != NULL) - { - ExFreePoolWithTag(context->filter, WINDIVERT_TAG); - context->filter = NULL; - } } /* @@ -1285,6 +1259,7 @@ extern VOID windivert_cleanup(IN WDFFILEOBJECT object) */ extern VOID windivert_close(IN WDFFILEOBJECT object) { + KLOCK_QUEUE_HANDLE lock_handle; context_t context = windivert_context_get(object); DEBUG("CLOSE: closing WinDivert context (context=%p)", context); @@ -1293,7 +1268,28 @@ extern VOID windivert_close(IN WDFFILEOBJECT object) { return; } + KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); context->state = WINDIVERT_CONTEXT_STATE_CLOSED; + KeReleaseInStackQueuedSpinLock(&lock_handle); +} + +/* + * WinDivert destroy routine. + */ +extern VOID windivert_destroy(IN WDFOBJECT object) +{ + context_t context = windivert_context_get((WDFFILEOBJECT)object); + + DEBUG("DESTROY: destroying WinDivert context (context=%p)", context); + + if (!windivert_context_verify(context, WINDIVERT_CONTEXT_STATE_CLOSED)) + { + return; + } + if (context->filter != NULL) + { + ExFreePoolWithTag(context->filter, WINDIVERT_TAG); + } } /* @@ -1321,19 +1317,83 @@ static NTSTATUS windivert_read(context_t context, WDFREQUEST request) } /* - * WinDivert read service worker. + * WinDivert service a single read request. */ -VOID windivert_read_service_work_item(IN WDFWORKITEM item) +static void windivert_read_service_request(packet_t packet, + PNET_BUFFER buffer, UINT8 direction, UINT32 if_idx, UINT32 sub_if_idx, + UINT8 checksums, WDFREQUEST request) { - WDFFILEOBJECT object = (WDFFILEOBJECT)WdfWorkItemGetParentObject(item); - context_t context = windivert_context_get(object); + PMDL dst_mdl; + PVOID dst, src; + ULONG dst_len, src_len; + NTSTATUS status; + req_context_t req_context; + windivert_addr_t addr; - if (!windivert_context_verify(context, WINDIVERT_CONTEXT_STATE_OPEN)) + DEBUG("SERVICE: servicing read request (request=%p)", request); + + status = WdfRequestRetrieveOutputWdmMdl(request, &dst_mdl); + dst_len = 0; + if (!NT_SUCCESS(status)) { - return; + DEBUG_ERROR("failed to retrieve output MDL", status); + goto windivert_read_service_request_exit; + } + dst = MmGetSystemAddressForMdlSafe(dst_mdl, NormalPagePriority); + if (dst == NULL) + { + status = STATUS_INSUFFICIENT_RESOURCES; + DEBUG_ERROR("failed to get address of output MDL", status); + goto windivert_read_service_request_exit; } - windivert_read_service(context); + dst_len = MmGetMdlByteCount(dst_mdl); + if (packet != NULL) + { + src_len = packet->data_len; + dst_len = (src_len < dst_len? src_len: dst_len); + src = packet->data; + RtlCopyMemory(dst, src, dst_len); + } + else + { + src_len = NET_BUFFER_DATA_LENGTH(buffer); + dst_len = (src_len < dst_len? src_len: dst_len); + src = NdisGetDataBuffer(buffer, dst_len, NULL, 1, 0); + if (src == NULL) + { + NdisGetDataBuffer(buffer, dst_len, dst, 1, 0); + } + else + { + RtlCopyMemory(dst, src, dst_len); + } + } + + // Write the address information. + req_context = windivert_req_context_get(request); + addr = req_context->addr; + if (addr != NULL) + { + addr->IfIdx = if_idx; + addr->SubIfIdx = sub_if_idx; + addr->Direction = direction; + } + + // Zero the IP/TCP/UDP checksums here (if required). + windivert_zero_checksums(dst, dst_len, checksums); + + status = STATUS_SUCCESS; + +windivert_read_service_request_exit: + if (NT_SUCCESS(status)) + { + WdfRequestCompleteWithInformation(request, status, dst_len); + } + else + { + WdfRequestComplete(request, status); + } } /* @@ -1341,82 +1401,52 @@ VOID windivert_read_service_work_item(IN WDFWORKITEM item) */ static void windivert_read_service(context_t context) { - PNET_BUFFER buffer; KLOCK_QUEUE_HANDLE lock_handle; WDFREQUEST request; PLIST_ENTRY entry; PMDL dst_mdl; PVOID dst, src; ULONG dst_len, src_len; + ULONGLONG timestamp; + BOOL timeout; NTSTATUS status; packet_t packet; req_context_t req_context; windivert_addr_t addr; + timestamp = (ULONGLONG)KeQueryPerformanceCounter(NULL).QuadPart; KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); while (context->state == WINDIVERT_CONTEXT_STATE_OPEN && !IsListEmpty(&context->packet_queue)) { - status = WdfIoQueueRetrieveNextRequest(context->read_queue, &request); - if (!NT_SUCCESS(status)) - { - break; - } entry = RemoveHeadList(&context->packet_queue); packet = CONTAINING_RECORD(entry, struct packet_s, entry); + timeout = (timestamp - packet->timestamp > + context->packet_queue_maxcounts); + request = NULL; + if (!timeout) + { + status = WdfIoQueueRetrieveNextRequest(context->read_queue, + &request); + if (!NT_SUCCESS(status)) + { + InsertHeadList(&context->packet_queue, entry); + break; + } + } context->packet_queue_length--; context->packet_queue_size -= packet->data_len; KeReleaseInStackQueuedSpinLock(&lock_handle); - - DEBUG("SERVICE: servicing read request (context=%p, request=%p, " - "packet=%p)", context, request, packet); - - // We have now have a read request and a packet; service the read. - status = WdfRequestRetrieveOutputWdmMdl(request, &dst_mdl); - dst_len = 0; - if (!NT_SUCCESS(status)) + + if (!timeout) { - DEBUG_ERROR("failed to retrieve output MDL", status); - goto windivert_read_service_complete; - } - dst = MmGetSystemAddressForMdlSafe(dst_mdl, NormalPagePriority); - if (dst == NULL) - { - status = STATUS_INSUFFICIENT_RESOURCES; - DEBUG_ERROR("failed to get address of output MDL", status); - goto windivert_read_service_complete; - } - dst_len = MmGetMdlByteCount(dst_mdl); - src_len = packet->data_len; - dst_len = (src_len < dst_len? src_len: dst_len); - src = packet->data; - RtlCopyMemory(dst, src, dst_len); - - // Write the address information. - req_context = windivert_req_context_get(request); - addr = req_context->addr; - if (addr != NULL) - { - addr->IfIdx = packet->if_idx; - addr->SubIfIdx = packet->sub_if_idx; - addr->Direction = packet->direction; + windivert_read_service_request(packet, NULL, + packet->direction, packet->if_idx, packet->sub_if_idx, + packet->checksums, request); } - // Zero the IP/TCP/UDP checksums here (if required). - windivert_zero_checksums(dst, dst_len, packet->checksums); - - status = STATUS_SUCCESS; - -windivert_read_service_complete: windivert_free_packet(packet); - if (NT_SUCCESS(status)) - { - WdfRequestCompleteWithInformation(request, status, dst_len); - } - else - { - WdfRequestComplete(request, status); - } + timestamp = (ULONGLONG)KeQueryPerformanceCounter(NULL).QuadPart; KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); } KeReleaseInStackQueuedSpinLock(&lock_handle); @@ -1829,10 +1859,6 @@ extern VOID windivert_ioctl(IN WDFQUEUE queue, IN WDFREQUEST request, status = windivert_install_callouts(context, is_inbound, is_outbound, is_ipv4, is_ipv6); - // Start the timer. - WdfTimerStart(context->timer, - WDF_REL_TIMEOUT_IN_MS(context->timer_timeout)); - break; } @@ -1908,7 +1934,9 @@ extern VOID windivert_ioctl(IN WDFQUEUE queue, IN WDFREQUEST request, "value", status); goto windivert_ioctl_exit; } - context->timer_timeout = (UINT)value; + context->packet_queue_maxcounts = + (ULONGLONG)value * counts_per_ms; + context->packet_queue_maxtime = (ULONG)value; break; case WINDIVERT_PARAM_QUEUE_SIZE: @@ -1957,7 +1985,7 @@ extern VOID windivert_ioctl(IN WDFQUEUE queue, IN WDFREQUEST request, *valptr = context->packet_queue_maxlength; break; case WINDIVERT_PARAM_QUEUE_TIME: - *valptr = context->timer_timeout; + *valptr = context->packet_queue_maxtime; break; case WINDIVERT_PARAM_QUEUE_SIZE: *valptr = context->packet_queue_maxsize; @@ -2004,7 +2032,8 @@ static void windivert_classify_outbound_network_v4_callout( const FWPS_FILTER0 *filter, IN UINT64 flow_context, OUT FWPS_CLASSIFY_OUT0 *result) { - windivert_classify_callout(WINDIVERT_DIRECTION_OUTBOUND, + windivert_classify_callout((context_t)filter->context, + WINDIVERT_DIRECTION_OUTBOUND, fixed_vals->incomingValue[ FWPS_FIELD_OUTBOUND_IPPACKET_V4_INTERFACE_INDEX].value.uint32, fixed_vals->incomingValue[ @@ -2013,7 +2042,7 @@ static void windivert_classify_outbound_network_v4_callout( (fixed_vals->incomingValue[ FWPS_FIELD_OUTBOUND_IPPACKET_V4_FLAGS].value.uint32 & FWP_CONDITION_FLAG_IS_LOOPBACK) != 0, - data, filter, flow_context, result); + 0, data, flow_context, result); } /* @@ -2025,7 +2054,8 @@ static void windivert_classify_outbound_network_v6_callout( const FWPS_FILTER0 *filter, IN UINT64 flow_context, OUT FWPS_CLASSIFY_OUT0 *result) { - windivert_classify_callout(WINDIVERT_DIRECTION_OUTBOUND, + windivert_classify_callout((context_t)filter->context, + WINDIVERT_DIRECTION_OUTBOUND, fixed_vals->incomingValue[ FWPS_FIELD_OUTBOUND_IPPACKET_V6_INTERFACE_INDEX].value.uint32, fixed_vals->incomingValue[ @@ -2034,7 +2064,7 @@ static void windivert_classify_outbound_network_v6_callout( (fixed_vals->incomingValue[ FWPS_FIELD_OUTBOUND_IPPACKET_V6_FLAGS].value.uint32 & FWP_CONDITION_FLAG_IS_LOOPBACK) != 0, - data, filter, flow_context, result); + 0, data, flow_context, result); } /* @@ -2046,24 +2076,9 @@ static void windivert_classify_inbound_network_v4_callout( const FWPS_FILTER0 *filter, IN UINT64 flow_context, OUT FWPS_CLASSIFY_OUT0 *result) { - PNET_BUFFER_LIST buffers = (PNET_BUFFER_LIST)data; - PNET_BUFFER buffer; - NTSTATUS status; - - if (!(result->rights & FWPS_RIGHT_ACTION_WRITE) || data == NULL) - { - return; - } - - buffer = NET_BUFFER_LIST_FIRST_NB(buffers); - status = NdisRetreatNetBufferDataStart(buffer, meta_vals->ipHeaderSize, - 0, NULL); - if (!NT_SUCCESS(status)) - { - result->actionType = FWP_ACTION_CONTINUE; - return; - } - windivert_classify_callout(WINDIVERT_DIRECTION_INBOUND, + UINT advance = meta_vals->ipHeaderSize; + windivert_classify_callout((context_t)filter->context, + WINDIVERT_DIRECTION_INBOUND, fixed_vals->incomingValue[ FWPS_FIELD_INBOUND_IPPACKET_V4_INTERFACE_INDEX].value.uint32, fixed_vals->incomingValue[ @@ -2072,9 +2087,7 @@ static void windivert_classify_inbound_network_v4_callout( (fixed_vals->incomingValue[ FWPS_FIELD_INBOUND_IPPACKET_V4_FLAGS].value.uint32 & FWP_CONDITION_FLAG_IS_LOOPBACK) != 0, - data, filter, flow_context, result); - NdisAdvanceNetBufferDataStart(buffer, meta_vals->ipHeaderSize, FALSE, - NULL); + advance, data, flow_context, result); } /* @@ -2086,24 +2099,9 @@ static void windivert_classify_inbound_network_v6_callout( const FWPS_FILTER0 *filter, IN UINT64 flow_context, OUT FWPS_CLASSIFY_OUT0 *result) { - PNET_BUFFER_LIST buffers = (PNET_BUFFER_LIST)data; - PNET_BUFFER buffer; - NTSTATUS status; - - if (!(result->rights & FWPS_RIGHT_ACTION_WRITE) || data == NULL) - { - return; - } - - buffer = NET_BUFFER_LIST_FIRST_NB(buffers); - status = NdisRetreatNetBufferDataStart(buffer, sizeof(struct ipv6hdr), - 0, NULL); - if (!NT_SUCCESS(status)) - { - result->actionType = FWP_ACTION_CONTINUE; - return; - } - windivert_classify_callout(WINDIVERT_DIRECTION_INBOUND, + UINT advance = meta_vals->ipHeaderSize; + windivert_classify_callout((context_t)filter->context, + WINDIVERT_DIRECTION_INBOUND, fixed_vals->incomingValue[ FWPS_FIELD_INBOUND_IPPACKET_V6_INTERFACE_INDEX].value.uint32, fixed_vals->incomingValue[ @@ -2112,9 +2110,7 @@ static void windivert_classify_inbound_network_v6_callout( (fixed_vals->incomingValue[ FWPS_FIELD_INBOUND_IPPACKET_V6_FLAGS].value.uint32 & FWP_CONDITION_FLAG_IS_LOOPBACK) != 0, - data, filter, flow_context, result); - NdisAdvanceNetBufferDataStart(buffer, sizeof(struct ipv6hdr), FALSE, - NULL); + advance, data, flow_context, result); } /* @@ -2126,10 +2122,11 @@ static void windivert_classify_forward_network_v4_callout( const FWPS_FILTER0 *filter, IN UINT64 flow_context, OUT FWPS_CLASSIFY_OUT0 *result) { - windivert_classify_callout(WINDIVERT_DIRECTION_OUTBOUND, + windivert_classify_callout((context_t)filter->context, + WINDIVERT_DIRECTION_OUTBOUND, fixed_vals->incomingValue[ FWPS_FIELD_IPFORWARD_V4_DESTINATION_INTERFACE_INDEX].value.uint32, - 0, TRUE, FALSE, data, filter, flow_context, result); + 0, TRUE, FALSE, 0, data, flow_context, result); } /* @@ -2141,32 +2138,36 @@ static void windivert_classify_forward_network_v6_callout( const FWPS_FILTER0 *filter, IN UINT64 flow_context, OUT FWPS_CLASSIFY_OUT0 *result) { - windivert_classify_callout(WINDIVERT_DIRECTION_OUTBOUND, + windivert_classify_callout((context_t)filter->context, + WINDIVERT_DIRECTION_OUTBOUND, fixed_vals->incomingValue[ FWPS_FIELD_IPFORWARD_V6_DESTINATION_INTERFACE_INDEX].value.uint32, - 0, FALSE, FALSE, data, filter, flow_context, result); + 0, FALSE, FALSE, 0, data, flow_context, result); } /* * WinDivert classify callout. */ -static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx, - IN UINT32 sub_if_idx, IN BOOL isipv4, IN BOOL isloopback, - IN OUT void *data, const FWPS_FILTER0 *filter, IN UINT64 flow_context, +static void windivert_classify_callout(context_t context, IN UINT8 direction, + IN UINT32 if_idx, IN UINT32 sub_if_idx, IN BOOL isipv4, IN BOOL isloopback, + IN UINT advance, IN OUT void *data, IN UINT64 flow_context, OUT FWPS_CLASSIFY_OUT0 *result) { KLOCK_QUEUE_HANDLE lock_handle; FWPS_PACKET_INJECTION_STATE packet_state; HANDLE packet_context; - UINT32 priority; + UINT32 priority, packet_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; + filter_t filter; + WDFOBJECT object; + work_t work; ULONG read_queue_len; + ULONGLONG timestamp; + NTSTATUS status; // Basic checks: if (!(result->rights & FWPS_RIGHT_ACTION_WRITE) || data == NULL) @@ -2174,20 +2175,12 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx, return; } - context = (context_t)filter->context; - if (!windivert_context_verify(context, WINDIVERT_CONTEXT_STATE_OPEN)) - { - result->actionType = FWP_ACTION_CONTINUE; - return; - } buffers = (PNET_BUFFER_LIST)data; buffer = NET_BUFFER_LIST_FIRST_NB(buffers); if (NET_BUFFER_LIST_NEXT_NBL(buffers) != NULL) { - /* - * This is a fragment group. This can be ignored since each - * fragment should already have been indicated. - */ + // This is a fragment group. This can be ignored since each fragment + // should have already been indicated. result->actionType = FWP_ACTION_CONTINUE; return; } @@ -2201,20 +2194,36 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx, packet_state = FwpsQueryPacketInjectionState0(injectv6_handle, buffers, &packet_context); } + + KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); + if (context->state != WINDIVERT_CONTEXT_STATE_OPEN) + { + KeReleaseInStackQueuedSpinLock(&lock_handle); + result->actionType = FWP_ACTION_CONTINUE; + return; + } + priority = context->priority; + filter = context->filter; + object = (WDFOBJECT)context->object; + WdfObjectReference(object); + KeReleaseInStackQueuedSpinLock(&lock_handle); + if (packet_state == FWPS_PACKET_INJECTED_BY_SELF || packet_state == FWPS_PACKET_PREVIOUSLY_INJECTED_BY_SELF) { - priority = (UINT32)packet_context; - if (priority >= context->priority) + packet_priority = (UINT32)packet_context; + if (packet_priority >= priority) { + WdfObjectDereference(object); result->actionType = FWP_ACTION_CONTINUE; return; } } - /* - * Determine which checksum fields are present or not. - */ + timestamp = (ULONGLONG)KeQueryPerformanceCounter(NULL).QuadPart; + timestamp--; + + // Determine which checksum fields are present or not. if (isloopback) { // Loopback packets appear to have bogus checksums, so do not trust. @@ -2234,6 +2243,20 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx, checksums = WINDIVERT_ALL_CHECKSUMS; } + // Retreat the NET_BUFFER to the IP header, if necessary. + // If (advance != 0) then this must be in the inbound path, and the + // NET_BUFFER_LIST must contain exactly one NET_BUFFER. + if (advance != 0) + { + status = NdisRetreatNetBufferDataStart(buffer, advance, 0, NULL); + if (!NT_SUCCESS(status)) + { + WdfObjectDereference(object); + result->actionType = FWP_ACTION_CONTINUE; + return; + } + } + /* * This code is complicated by the fact the a single NET_BUFFER_LIST * may contain several NET_BUFFER structures. Each NET_BUFFER needs to @@ -2241,7 +2264,8 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx, * 1) First check if any NET_BUFFER passes the filter. * 2) If no, then CONTINUE the entire NET_BUFFER_LIST. * 3) Else, split the NET_BUFFER_LIST into individual NET_BUFFERs; and - * either queue or re-inject based on the filter. + * either queue or re-inject based on the filter. This step is done + * out-of-band, see windivert_worker(). */ // Find the first NET_BUFFER we need to queue: @@ -2249,133 +2273,266 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx, outbound = (direction == WINDIVERT_DIRECTION_OUTBOUND); do { - if (windivert_filter(buffer_fst, if_idx, sub_if_idx, outbound, - isipv4, checksums, context->filter)) + BOOL match = windivert_filter(buffer_fst, if_idx, sub_if_idx, outbound, + isipv4, checksums, filter); + if (match) { break; } buffer_fst = NET_BUFFER_NEXT_NB(buffer_fst); } while (buffer_fst != NULL); + if (advance != 0) + { + NdisAdvanceNetBufferDataStart(buffer, advance, FALSE, NULL); + } if (buffer_fst == NULL) { + // No packet matches the filter; continue the entire NET_BUFFER_LIST. + WdfObjectDereference(object); result->actionType = FWP_ACTION_CONTINUE; return; } - - if ((context->flags & WINDIVERT_FLAG_SNIFF) == 0) + + // At least one packet matches the filter. Delay all further processing + // until windivert_worker() at IRQL=PASSIVE_LEVEL. + work = (work_t)ExAllocatePoolWithTag(NonPagedPool, sizeof(struct work_s), + WINDIVERT_TAG); + if (work == NULL) { - // Re-inject all packets up to 'buffer_fst' - buffer_itr = buffer; - while (buffer_itr != buffer_fst) - { - if (!windivert_reinject_packet(context, direction, isipv4, if_idx, - sub_if_idx, priority, buffer_itr)) - { - goto windivert_classify_callout_exit; - } - buffer_itr = NET_BUFFER_NEXT_NB(buffer_itr); - } - } - else - { - buffer_itr = buffer_fst; + goto windivert_classify_callout_exit; } + FwpsReferenceNetBufferList(buffers, TRUE); + work->buffers = buffers; + work->buffer = buffer_fst; + work->advance = advance; + work->is_ipv4 = isipv4; + work->checksums = checksums; + work->direction = direction; + work->if_idx = if_idx; + work->sub_if_idx = sub_if_idx; + work->priority = priority; + work->timestamp = timestamp; queued = FALSE; - if ((context->flags & WINDIVERT_FLAG_DROP) == 0) + KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); + if (context->state == WINDIVERT_CONTEXT_STATE_OPEN) { - if (!windivert_queue_packet(context, buffer_itr, direction, if_idx, - sub_if_idx, checksums)) + InsertTailList(&context->work_queue, &work->entry); + WdfWorkItemEnqueue(context->workers[context->worker_curr]); + context->worker_curr++; + if (context->worker_curr >= WINDIVERT_CONTEXT_MAXWORKERS) { - goto windivert_classify_callout_exit; + context->worker_curr = 0; } queued = TRUE; } - - // Queue or re-inject remaining packets. - buffer_itr = NET_BUFFER_NEXT_NB(buffer_itr); - while (buffer_itr != NULL) + KeReleaseInStackQueuedSpinLock(&lock_handle); + if (!queued) { - if (windivert_filter(buffer_itr, if_idx, sub_if_idx, outbound, - isipv4, checksums, context->filter)) - { - if ((context->flags & WINDIVERT_FLAG_DROP) == 0) - { - if (!windivert_queue_packet(context, buffer_itr, direction, - if_idx, sub_if_idx, checksums)) - { - goto windivert_classify_callout_exit; - } - queued = TRUE; - } - } - else if ((context->flags & WINDIVERT_FLAG_SNIFF) == 0) - { - if (!windivert_reinject_packet(context, direction, isipv4, if_idx, - sub_if_idx, priority, buffer_itr)) - { - goto windivert_classify_callout_exit; - } - } - buffer_itr = NET_BUFFER_NEXT_NB(buffer_itr); - } - - /* - * If the packet was queued, then service any pending read. - */ - if (queued) - { - KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); - if (context->state == WINDIVERT_CONTEXT_STATE_OPEN) - { - WdfIoQueueGetState(context->read_queue, &read_queue_len, NULL); - if (read_queue_len > 0) - { - WdfWorkItemEnqueue(context->workers[context->worker_curr]); - context->worker_curr++; - if (context->worker_curr >= WINDIVERT_CONTEXT_MAXWORKERS) - { - context->worker_curr = 0; - } - } - } - KeReleaseInStackQueuedSpinLock(&lock_handle); + WdfObjectDereference(object); + FwpsDereferenceNetBufferList(buffers, FALSE); + ExFreePoolWithTag(work, WINDIVERT_TAG); + result->actionType = FWP_ACTION_CONTINUE; + return; } windivert_classify_callout_exit: - if ((context->flags & WINDIVERT_FLAG_SNIFF) != 0) + WdfObjectDereference(object); + result->actionType = FWP_ACTION_BLOCK; + result->flags |= FWPS_CLASSIFY_OUT_FLAG_ABSORB; + result->rights &= ~FWPS_RIGHT_ACTION_WRITE; +} + +/* + * WinDivert work item routine for out-of-band filtering. + */ +VOID windivert_worker(IN WDFWORKITEM item) +{ + KLOCK_QUEUE_HANDLE lock_handle; + WDFFILEOBJECT object = (WDFFILEOBJECT)WdfWorkItemGetParentObject(item); + PLIST_ENTRY entry; + work_t work; + context_t context = windivert_context_get(object); + PNET_BUFFER_LIST buffers_clone; + PNET_BUFFER buffer_itr, buffer_fst; + UINT advance; + BOOL match, ok, outbound, sniff_mode, forward; + filter_t filter; + NTSTATUS status; + + KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); + sniff_mode = ((context->flags & WINDIVERT_FLAG_SNIFF) != 0); + forward = (context->layer == WINDIVERT_LAYER_NETWORK_FORWARD); + filter = context->filter; + while (context->state == WINDIVERT_CONTEXT_STATE_OPEN && + !IsListEmpty(&context->work_queue)) { - result->actionType = FWP_ACTION_CONTINUE; - } - else - { - result->actionType = FWP_ACTION_BLOCK; - result->flags |= FWPS_CLASSIFY_OUT_FLAG_ABSORB; - result->rights &= ~FWPS_RIGHT_ACTION_WRITE; + entry = RemoveHeadList(&context->work_queue); + KeReleaseInStackQueuedSpinLock(&lock_handle); + + work = CONTAINING_RECORD(entry, struct work_s, entry); + advance = work->advance; + buffer_fst = work->buffer; + if (advance != 0) + { + status = NdisRetreatNetBufferDataStart(buffer_fst, advance, 0, + NULL); + if (!NT_SUCCESS(status)) + { + advance = 0; + goto windivert_worker_complete; + } + } + + if (!sniff_mode) + { + // In non-SNIFF mode, reinject all non-matching packets up to the + // first matching packet. + buffer_itr = NET_BUFFER_LIST_FIRST_NB(work->buffers); + while (buffer_itr != buffer_fst) + { + ok = windivert_reinject_packet(sniff_mode, forward, + work->direction, work->is_ipv4, work->if_idx, + work->sub_if_idx, work->priority, work->buffers, + buffer_itr); + if (!ok) + { + goto windivert_worker_complete; + } + buffer_itr = NET_BUFFER_NEXT_NB(buffer_itr); + } + } + else + { + buffer_itr = buffer_fst; + } + + // Queue the first matching packet. + ok = windivert_queue_packet(context, buffer_itr, work->direction, + work->if_idx, work->sub_if_idx, work->checksums, work->timestamp); + if (!ok) + { + goto windivert_worker_complete; + } + buffer_itr = NET_BUFFER_NEXT_NB(buffer_itr); + + // Queue or re-inject all remaining packets. + outbound = (work->direction == WINDIVERT_DIRECTION_OUTBOUND); + while (buffer_itr != NULL) + { + match = windivert_filter(buffer_itr, work->if_idx, + work->sub_if_idx, outbound, work->is_ipv4, work->checksums, + filter); + if (match) + { + ok = windivert_queue_packet(context, buffer_itr, + work->direction, work->if_idx, work->sub_if_idx, + work->checksums, work->timestamp); + } + else + { + ok = windivert_reinject_packet(sniff_mode, forward, + work->direction, work->is_ipv4, work->if_idx, + work->sub_if_idx, work->priority, work->buffers, + buffer_itr); + } + if (!ok) + { + goto windivert_worker_complete; + } + buffer_itr = NET_BUFFER_NEXT_NB(buffer_itr); + } + + if (sniff_mode) + { + // In SNIFF mode, reinject the entire NET_BUFFER_LIST. + ok = windivert_reinject_packet(sniff_mode, forward, + work->direction, work->is_ipv4, work->if_idx, + work->sub_if_idx, work->priority, work->buffers, NULL); + if (!ok) + { + goto windivert_worker_complete; + } + } + +windivert_worker_complete: + if (advance != 0) + { + NdisAdvanceNetBufferDataStart(work->buffer, advance, 0, 0); + } + FwpsDereferenceNetBufferList(work->buffers, FALSE); + ExFreePoolWithTag(work, WINDIVERT_TAG); + KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); } + KeReleaseInStackQueuedSpinLock(&lock_handle); } /* * Queue a NET_BUFFER. */ static BOOL windivert_queue_packet(context_t context, PNET_BUFFER buffer, - UINT8 direction, UINT32 if_idx, UINT32 sub_if_idx, UINT8 checksums) + UINT8 direction, UINT32 if_idx, UINT32 sub_if_idx, UINT8 checksums, + ULONGLONG timestamp0) { KLOCK_QUEUE_HANDLE lock_handle; PVOID data; + WDFREQUEST request; PLIST_ENTRY entry, old_entry; packet_t packet, old_packet; UINT data_len; + ULONGLONG timestamp; + BOOL timeout; NTSTATUS status; + // First we attempt to immediately service a read request directly without + // queuing the packet. This helps reduce overhead where possible. + timeout = FALSE; + request = NULL; + timestamp = (ULONGLONG)KeQueryPerformanceCounter(NULL).QuadPart; + KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); + if (context->state != WINDIVERT_CONTEXT_STATE_OPEN) + { + KeReleaseInStackQueuedSpinLock(&lock_handle); + return FALSE; + } + if ((context->flags & WINDIVERT_FLAG_DROP) != 0) + { + KeReleaseInStackQueuedSpinLock(&lock_handle); + return TRUE; + } + timeout = (timestamp - timestamp0 > context->packet_queue_maxcounts); + if (!timeout && IsListEmpty(&context->packet_queue)) + { + status = WdfIoQueueRetrieveNextRequest(context->read_queue, &request); + if (!NT_SUCCESS(status)) + { + request = NULL; + } + } + KeReleaseInStackQueuedSpinLock(&lock_handle); + if (timeout) + { + return TRUE; + } + if (request != NULL) + { + // FAST PATH: Service an I/O request without queueing the packet. + windivert_read_service_request(NULL, buffer, direction, if_idx, + sub_if_idx, checksums, request); + return TRUE; + } + + // SLOW PATH: queue the packet. data_len = NET_BUFFER_DATA_LENGTH(buffer); packet = (packet_t)ExAllocatePoolWithTag(NonPagedPool, WINDIVERT_PACKET_SIZE + data_len, WINDIVERT_TAG); if (packet == NULL) { + status = STATUS_INSUFFICIENT_RESOURCES; + DEBUG_ERROR("failed to allocate queued packet", status); return FALSE; } packet->data_len = data_len; @@ -2392,24 +2549,41 @@ static BOOL windivert_queue_packet(context_t context, PNET_BUFFER buffer, packet->direction = direction; packet->if_idx = if_idx; packet->sub_if_idx = sub_if_idx; - packet->timer_ticktock = context->timer_ticktock; + packet->timestamp = timestamp0; entry = &packet->entry; + timestamp = (ULONGLONG)KeQueryPerformanceCounter(NULL).QuadPart; KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); while (TRUE) { - if (context->state != WINDIVERT_CONTEXT_STATE_OPEN || - data_len > context->packet_queue_size) + if (context->state != WINDIVERT_CONTEXT_STATE_OPEN) { KeReleaseInStackQueuedSpinLock(&lock_handle); windivert_free_packet(packet); return FALSE; } + if (data_len > context->packet_queue_maxsize) + { + // (Corner case) the packet is larger than the max queue size: + KeReleaseInStackQueuedSpinLock(&lock_handle); + windivert_free_packet(packet); + return TRUE; + } + timeout = (timestamp - packet->timestamp > + context->packet_queue_maxcounts); + if (timeout) + { + // (Corner case) the packet has already expired: + KeReleaseInStackQueuedSpinLock(&lock_handle); + windivert_free_packet(packet); + return TRUE; + } + if (context->packet_queue_size + data_len > context->packet_queue_maxsize || context->packet_queue_length + 1 > context->packet_queue_maxlength) { - // The queue is full; drop a packet: + // The queue is full; drop a packet & try again: old_entry = RemoveHeadList(&context->packet_queue); old_packet = CONTAINING_RECORD(old_entry, struct packet_s, entry); context->packet_queue_length--; @@ -2417,6 +2591,7 @@ static BOOL windivert_queue_packet(context_t context, PNET_BUFFER buffer, KeReleaseInStackQueuedSpinLock(&lock_handle); DEBUG("DROP: packet queue is full, dropping packet"); windivert_free_packet(old_packet); + timestamp = (ULONGLONG)KeQueryPerformanceCounter(NULL).QuadPart; KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); continue; } @@ -2433,121 +2608,109 @@ static BOOL windivert_queue_packet(context_t context, PNET_BUFFER buffer, DEBUG("PACKET: diverting packet (packet=%p)", packet); + // Service any pending I/O request. + windivert_read_service(context); + return TRUE; } /* - * Re-inject a NET_BUFFER. + * Re-inject a packet or packets. */ -static BOOL windivert_reinject_packet(context_t context, UINT8 direction, - BOOL isipv4, UINT32 if_idx, UINT32 sub_if_idx, UINT32 priority, - PNET_BUFFER buffer) +static BOOL windivert_reinject_packet(BOOL sniff_mode, BOOL forward, + UINT8 direction, BOOL isipv4, UINT32 if_idx, UINT32 sub_if_idx, + UINT32 priority, PNET_BUFFER_LIST buffers, PNET_BUFFER buffer) { - UINT data_len; - PVOID data, data_copy = NULL; - PNET_BUFFER_LIST buffers = NULL; - PMDL mdl_copy = NULL; + PNET_BUFFER_LIST buffers_cpy; HANDLE handle; - NTSTATUS status = STATUS_SUCCESS; + NTSTATUS status; - data_len = NET_BUFFER_DATA_LENGTH(buffer); - data_copy = ExAllocatePoolWithTag(NonPagedPool, data_len, WINDIVERT_TAG); - if (data_copy == NULL) + if (buffer != NULL) { - status = STATUS_INSUFFICIENT_RESOURCES; - DEBUG_ERROR("failed to allocate memory for (re)injected packet data", - status); - return FALSE; - } - - data = NdisGetDataBuffer(buffer, data_len, NULL, 1, 0); - if (data == NULL) - { - NdisGetDataBuffer(buffer, data_len, data_copy, 1, 0); + // Re-inject a specific packet. Not applicable for SNIFF mode. + if (sniff_mode) + { + return TRUE; + } + status = FwpsAllocateNetBufferAndNetBufferList0(pool_handle, 0, 0, + NET_BUFFER_FIRST_MDL(buffer), NET_BUFFER_DATA_OFFSET(buffer), + NET_BUFFER_DATA_LENGTH(buffer), &buffers_cpy); + if (!NT_SUCCESS(status)) + { + DEBUG_ERROR("failed to create NET_BUFFER_LIST for injected packet", + status); + return FALSE; + } + FwpsReferenceNetBufferList(buffers, TRUE); } else { - RtlCopyMemory(data_copy, data, data_len); - } - - mdl_copy = IoAllocateMdl(data_copy, data_len, FALSE, FALSE, NULL); - if (mdl_copy == NULL) - { - status = STATUS_INSUFFICIENT_RESOURCES; - DEBUG_ERROR("failed to allocate MDL for injected packet", status); - goto windivert_reinject_packet_exit; - } - - MmBuildMdlForNonPagedPool(mdl_copy); - status = FwpsAllocateNetBufferAndNetBufferList0(pool_handle, 0, 0, - mdl_copy, 0, data_len, &buffers); - if (!NT_SUCCESS(status)) - { - DEBUG_ERROR("failed to create NET_BUFFER_LIST for injected packet", - status); - goto windivert_reinject_packet_exit; + // Re-inject all packets for SNIFF mode. + status = FwpsAllocateCloneNetBufferList0(buffers, NULL, pool_handle, + 0, &buffers_cpy); + if (!NT_SUCCESS(status)) + { + DEBUG_ERROR("failed to clone NET_BUFFER_LIST for injected packets", + status); + return FALSE; + } + buffers = NULL; // (buffers == NULL) indicates a clone } handle = (isipv4? inject_handle: injectv6_handle); - if (context->layer == WINDIVERT_LAYER_NETWORK_FORWARD) + if (forward) { status = FwpsInjectForwardAsync0(handle, (HANDLE)priority, 0, (isipv4? AF_INET: AF_INET6), UNSPECIFIED_COMPARTMENT_ID, - if_idx, buffers, windivert_reinject_complete, (HANDLE)NULL); + if_idx, buffers_cpy, windivert_reinject_complete, (HANDLE)buffers); } else if (direction == WINDIVERT_DIRECTION_OUTBOUND) { status = FwpsInjectNetworkSendAsync0(handle, - (HANDLE)priority, 0, UNSPECIFIED_COMPARTMENT_ID, buffers, - windivert_reinject_complete, (HANDLE)NULL); + (HANDLE)priority, 0, UNSPECIFIED_COMPARTMENT_ID, buffers_cpy, + windivert_reinject_complete, (HANDLE)buffers); } else { status = FwpsInjectNetworkReceiveAsync0(handle, (HANDLE)priority, 0, UNSPECIFIED_COMPARTMENT_ID, if_idx, - sub_if_idx, buffers, windivert_reinject_complete, (HANDLE)NULL); + sub_if_idx, buffers_cpy, windivert_reinject_complete, + (HANDLE)buffers); } -windivert_reinject_packet_exit: - if (!NT_SUCCESS(status)) { - DEBUG_ERROR("failed to (re)inject packet", status); + DEBUG_ERROR("failed to (re)inject packet(s)", status); if (buffers != NULL) { - FwpsFreeNetBufferList0(buffers); + FwpsFreeNetBufferList0(buffers_cpy); + FwpsDereferenceNetBufferList(buffers, FALSE); } - if (mdl_copy != NULL) + else { - IoFreeMdl(mdl_copy); - } - if (data_copy != NULL) - { - ExFreePoolWithTag(data_copy, WINDIVERT_TAG); + FwpsFreeCloneNetBufferList0(buffers_cpy, 0); } } - return NT_SUCCESS(status); + return TRUE; } /* * WinDivert (re)inject complete. */ static void NTAPI windivert_reinject_complete(VOID *context, - NET_BUFFER_LIST *buffers, BOOLEAN dispatch_level) + NET_BUFFER_LIST *buffers_cpy, BOOLEAN dispatch_level) { - PMDL mdl; - PVOID data; - PNET_BUFFER buffer = NET_BUFFER_LIST_FIRST_NB(buffers); - - mdl = NET_BUFFER_FIRST_MDL(buffer); - data = MmGetSystemAddressForMdlSafe(mdl, NormalPagePriority); - if (data != NULL) + PNET_BUFFER_LIST buffers = (PNET_BUFFER_LIST)context; + if (buffers != NULL) { - ExFreePoolWithTag(data, WINDIVERT_TAG); + FwpsFreeNetBufferList0(buffers_cpy); + FwpsDereferenceNetBufferList(buffers, dispatch_level); + } + else + { + FwpsFreeCloneNetBufferList0(buffers_cpy, 0); } - IoFreeMdl(mdl); - FwpsFreeNetBufferList0(buffers); } /*