From 76459ac3fe622e8576a4b97e7f9f62d814002af0 Mon Sep 17 00:00:00 2001 From: basil00 Date: Wed, 18 Dec 2013 09:54:01 +0800 Subject: [PATCH] - Use WDF work items instead of threads. Also use multiple work items so that read_service() can be run in parallel. --- sys/windivert.c | 132 +++++++++++++++++++++++++----------------------- 1 file changed, 70 insertions(+), 62 deletions(-) diff --git a/sys/windivert.c b/sys/windivert.c index 53f2389..bd14997 100644 --- a/sys/windivert.c +++ b/sys/windivert.c @@ -36,7 +36,7 @@ EVT_WDF_DEVICE_FILE_CREATE windivert_create; EVT_WDF_TIMER windivert_timer; EVT_WDF_FILE_CLEANUP windivert_cleanup; EVT_WDF_FILE_CLOSE windivert_close; -KSTART_ROUTINE windivert_read_service_worker; +EVT_WDF_WORKITEM windivert_read_service_work_item; /* * Debugging macros. @@ -105,6 +105,7 @@ typedef struct filter_s *filter_t; #define WINDIVERT_CONTEXT_MAGIC 0xAA5D1C5BC439AA72ull #define WINDIVERT_CONTEXT_SIZE (sizeof(struct context_s)) #define WINDIVERT_CONTEXT_MAXLAYERS 4 +#define WINDIVERT_CONTEXT_MAXWORKERS 4 #define WINDIVERT_CONTEXT_OUTBOUND_IPV4_LAYER 0 #define WINDIVERT_CONTEXT_INBOUND_IPV4_LAYER 1 #define WINDIVERT_CONTEXT_OUTBOUND_IPV6_LAYER 2 @@ -130,8 +131,9 @@ struct context_s UINT timer_timeout; // Packet timeout (in ms). BOOL timer_ticktock; // Packet timer ticktock. WDFQUEUE read_queue; // Read queue. - KEVENT read_event; // Read event. - void *read_thread; // Read thread. + WDFWORKITEM workers[WINDIVERT_CONTEXT_MAXWORKERS]; + // Read workers. + UINT8 worker_curr; // Current read worker. UINT8 layer_0; // Context's layer (initial). UINT8 layer; // Context's layer. UINT64 flags_0; // Context's flags (initial). @@ -313,8 +315,8 @@ NDIS_HANDLE pool_handle; */ extern VOID windivert_ioctl(IN WDFQUEUE queue, IN WDFREQUEST request, IN size_t in_length, IN size_t out_len, IN ULONG code); -extern NTSTATUS windivert_read(context_t context, WDFREQUEST request); -static void windivert_read_service_worker(PVOID context_0); +static NTSTATUS windivert_read(context_t context, WDFREQUEST request); +extern VOID windivert_read_service_work_item(IN WDFWORKITEM item); static void windivert_read_service(context_t context); static BOOLEAN windivert_context_verify(context_t context, context_state_t state); @@ -546,6 +548,8 @@ extern NTSTATUS DriverEntry(IN PDRIVER_OBJECT driver_obj, queue_config.EvtIoWrite = NULL; queue_config.EvtIoDeviceControl = windivert_ioctl; WDF_OBJECT_ATTRIBUTES_INIT(&obj_attrs); + obj_attrs.ExecutionLevel = WdfExecutionLevelPassive; + obj_attrs.SynchronizationScope = WdfSynchronizationScopeNone; status = WdfIoQueueCreate(device, &queue_config, &obj_attrs, &queue); if (!NT_SUCCESS(status)) { @@ -648,7 +652,8 @@ extern VOID windivert_create(IN WDFDEVICE device, IN WDFREQUEST request, { WDF_IO_QUEUE_CONFIG queue_config; WDF_TIMER_CONFIG timer_config; - WDF_OBJECT_ATTRIBUTES timer_attributes; + WDF_WORKITEM_CONFIG item_config; + WDF_OBJECT_ATTRIBUTES obj_attrs; FWPM_SESSION0 session; HANDLE thread; NTSTATUS status = STATUS_SUCCESS; @@ -663,6 +668,7 @@ extern VOID windivert_create(IN WDFDEVICE device, IN WDFREQUEST request, context->device = device; context->packet_queue_length = 0; context->packet_queue_maxlength = WINDIVERT_PARAM_QUEUE_LEN_DEFAULT; + context->timer = NULL; context->timer_timeout = WINDIVERT_PARAM_QUEUE_TIME_DEFAULT; context->layer_0 = WINDIVERT_LAYER_DEFAULT; context->layer = WINDIVERT_LAYER_DEFAULT; @@ -670,8 +676,12 @@ extern VOID windivert_create(IN WDFDEVICE device, IN WDFREQUEST request, context->flags = 0; context->priority_0 = WINDIVERT_PRIORITY_DEFAULT; context->priority = WINDIVERT_PRIORITY_DEFAULT; - context->read_thread = NULL; context->filter = NULL; + for (i = 0; i < WINDIVERT_CONTEXT_MAXWORKERS; i++) + { + context->workers[i] = NULL; + } + context->worker_curr = 0; for (i = 0; i < WINDIVERT_CONTEXT_MAXLAYERS; i++) { context->registered[i] = FALSE; @@ -708,31 +718,30 @@ extern VOID windivert_create(IN WDFDEVICE device, IN WDFREQUEST request, DEBUG_ERROR("failed to create I/O read queue", status); goto windivert_create_exit; } - KeInitializeEvent(&context->read_event, NotificationEvent, FALSE); - status = PsCreateSystemThread(&thread, THREAD_ALL_ACCESS, NULL, NULL, - NULL, windivert_read_service_worker, (PVOID)context); - if (!NT_SUCCESS(status)) - { - DEBUG_ERROR("failed to create read service thread", status); - goto windivert_create_exit; - } - status = ObReferenceObjectByHandle(thread, 0, NULL, KernelMode, - &context->read_thread, NULL); - if (!NT_SUCCESS(status)) - { - DEBUG_ERROR("failed to create read service thread object", status); - goto windivert_create_exit; - } WDF_TIMER_CONFIG_INIT(&timer_config, windivert_timer); timer_config.AutomaticSerialization = TRUE; - WDF_OBJECT_ATTRIBUTES_INIT(&timer_attributes); - timer_attributes.ParentObject = (WDFOBJECT)object; - status = WdfTimerCreate(&timer_config, &timer_attributes, &context->timer); + 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); + item_config.AutomaticSerialization = FALSE; + WDF_OBJECT_ATTRIBUTES_INIT(&obj_attrs); + obj_attrs.ParentObject = (WDFOBJECT)object; + for (i = 0; i < WINDIVERT_CONTEXT_MAXWORKERS; i++) + { + status = WdfWorkItemCreate(&item_config, &obj_attrs, + context->workers + i); + if (!NT_SUCCESS(status)) + { + DEBUG_ERROR("failed to create read service work item", status); + goto windivert_create_exit; + } + } RtlZeroMemory(&session, sizeof(session)); session.flags |= FWPM_SESSION_FLAG_DYNAMIC; status = FwpmEngineOpen0(NULL, RPC_C_AUTHN_DEFAULT, NULL, &session, @@ -758,17 +767,17 @@ windivert_create_exit: { WdfObjectDelete(context->timer); } + for (i = 0; i < WINDIVERT_CONTEXT_MAXWORKERS; i++) + { + if (context->workers[i] != NULL) + { + WdfObjectDelete(context->workers[i]); + } + } if (context->engine_handle != NULL) { FwpmEngineClose0(context->engine_handle); } - if (context->read_thread != NULL) - { - KeSetEvent(&context->read_event, IO_NO_INCREMENT, FALSE); - KeWaitForSingleObject(context->read_thread, Executive, - KernelMode, FALSE, NULL); - ObDereferenceObject(context->read_thread); - } } WdfRequestComplete(request, status); @@ -1005,7 +1014,6 @@ extern VOID windivert_cleanup(IN WDFFILEOBJECT object) WdfTimerStop(context->timer, TRUE); KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); context->state = WINDIVERT_CONTEXT_STATE_CLOSING; - KeSetEvent(&context->read_event, IO_NO_INCREMENT, FALSE); while (!IsListEmpty(&context->packet_queue)) { entry = RemoveHeadList(&context->packet_queue); @@ -1018,6 +1026,11 @@ extern VOID windivert_cleanup(IN WDFFILEOBJECT object) 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]); + WdfObjectDelete(context->workers[i]); + } status = FwpmTransactionBegin0(context->engine_handle, 0); if (!NT_SUCCESS(status)) @@ -1069,9 +1082,6 @@ windivert_cleanup_exit: ExFreePoolWithTag(context->filter, WINDIVERT_TAG); context->filter = NULL; } - KeWaitForSingleObject(context->read_thread, Executive, KernelMode, FALSE, - NULL); - ObDereferenceObject(context->read_thread); } /* @@ -1115,37 +1125,19 @@ static NTSTATUS windivert_read(context_t context, WDFREQUEST request) } /* - * Handle windivert_read_service worker. + * WinDivert read service worker. */ -static void windivert_read_service_worker(PVOID context_0) +VOID windivert_read_service_work_item(IN WDFWORKITEM item) { - KLOCK_QUEUE_HANDLE lock_handle; - context_t context = (context_t)context_0; + WDFFILEOBJECT object = (WDFFILEOBJECT)WdfWorkItemGetParentObject(item); + context_t context = windivert_context_get(object); - /* - * NOTE: We cannot verify the context because we do not know what state - * it is in. - */ - - while (TRUE) + if (!windivert_context_verify(context, WINDIVERT_CONTEXT_STATE_OPEN)) { - KeWaitForSingleObject(&context->read_event, Executive, KernelMode, - FALSE, NULL); - - KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle); - KeClearEvent(&context->read_event); - if (context->state != WINDIVERT_CONTEXT_STATE_OPEN) - { - break; - } - KeReleaseInStackQueuedSpinLock(&lock_handle); - - // Service reads: - windivert_read_service(context); + return; } - KeReleaseInStackQueuedSpinLock(&lock_handle); - PsTerminateSystemThread(STATUS_SUCCESS); + windivert_read_service(context); } /* @@ -1930,6 +1922,7 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx, const FWPS_FILTER0 *filter, 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; @@ -1938,6 +1931,7 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx, BOOL outbound, queued; context_t context; packet_t packet; + ULONG read_queue_len; // Basic checks: if (!(result->rights & FWPS_RIGHT_ACTION_WRITE) || data == NULL) @@ -2013,7 +2007,21 @@ static void windivert_classify_callout(IN UINT8 direction, IN UINT32 if_idx, */ if (queued) { - KeSetEvent(&context->read_event, IO_NO_INCREMENT, FALSE); + 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); } windivert_classify_callout_exit: