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); } /*