From a205259becb275fe634403362a5ec6fd9366fc7b Mon Sep 17 00:00:00 2001
From: basil00
Date: Tue, 8 Nov 2011 23:54:01 +0800
Subject: [PATCH] - Divert driver and library now fully support multi-threaded
code. This includes making the driver handle I/O requests in parallel, and
changing the library such that it uses overlapped I/O. Overall the package
is now significantly faster. - The divert driver now only adds WFP callouts
when a filter is set, instead of when a handle is created. - The divert
driver now attempts to avoid adding WFP callouts that are superfluous with
respect to the filter. - The divert driver now explicitly deletes filters and
callouts during cleanup. This appears to resolve a bug where network
connectivity is sometimes lost. - Added a new sample program: passthru.exe.
This example doesn't do anything interesting but is useful for speed
testing. - Rename build.sh to mingw-build.sh to avoid confusion.
---
dll/divert.c | 94 ++++----
doc/divert.html | 10 +-
examples/dirs | 1 +
examples/passthru/Makefile | 1 +
examples/passthru/passthru.c | 117 ++++++++++
examples/passthru/sources | 36 +++
build.sh => mingw-build.sh | 2 +-
sys/divert.c | 420 +++++++++++++++++++++++++++--------
8 files changed, 542 insertions(+), 139 deletions(-)
create mode 100644 examples/passthru/Makefile
create mode 100644 examples/passthru/passthru.c
create mode 100644 examples/passthru/sources
rename build.sh => mingw-build.sh (98%)
diff --git a/dll/divert.c b/dll/divert.c
index d3add68..7a652dd 100644
--- a/dll/divert.c
+++ b/dll/divert.c
@@ -212,6 +212,8 @@ static HMODULE DivertLoadCoInstaller(LPWSTR divert_dll);
static BOOLEAN DivertDriverFiles(LPWSTR *divert_dir_ptr,
LPWSTR *divert_sys_ptr, LPWSTR *divert_inf_ptr, LPWSTR *divert_dll_ptr);
static BOOLEAN DivertDriverInstall(VOID);
+static BOOL DivertIoControl(HANDLE handle, DWORD code, UINT64 arg, PVOID buf,
+ UINT len, UINT *iolen);
static BOOL DivertCompileFilter(const char *filter_str,
divert_ioctl_filter_t filter, UINT8 *fp);
static int __cdecl DivertFilterTokenNameCompare(const void *a, const void *b);
@@ -476,6 +478,45 @@ DivertDriverInstallExit:
return installed;
}
+/*
+ * Perform a DeviceIoControl.
+ */
+static BOOL DivertIoControl(HANDLE handle, DWORD code, UINT64 arg, PVOID buf,
+ UINT len, UINT *iolen)
+{
+ struct divert_ioctl_s ioctl;
+ DWORD iolen0;
+ OVERLAPPED overlapped;
+
+ ioctl.version = DIVERT_VERSION;
+ ioctl.magic = DIVERT_MAGIC;
+ ioctl.reserved = 0x0;
+ ioctl.arg = arg;
+ overlapped.Offset = 0;
+ overlapped.OffsetHigh = 0;
+ overlapped.hEvent = CreateEvent(NULL, FALSE, FALSE, NULL);
+ if (overlapped.hEvent == NULL)
+ {
+ return FALSE;
+ }
+ if (!DeviceIoControl(handle, code, &ioctl, sizeof(ioctl), buf, (DWORD)len,
+ &iolen0, &overlapped))
+ {
+ if (GetLastError() != ERROR_IO_PENDING ||
+ !GetOverlappedResult(handle, &overlapped, &iolen0, TRUE))
+ {
+ CloseHandle(overlapped.hEvent);
+ return FALSE;
+ }
+ }
+ CloseHandle(overlapped.hEvent);
+ if (iolen != NULL)
+ {
+ *iolen = (UINT)iolen0;
+ }
+ return TRUE;
+}
+
/*
* Open a handle to the Divert device.
*/
@@ -484,7 +525,7 @@ extern HANDLE DivertOpen(const char *filter)
struct divert_ioctl_s ioctl;
struct divert_ioctl_filter_s ioctl_filter[DIVERT_FILTER_MAXLEN];
UINT8 filter_len;
- DWORD err, iolen;
+ DWORD err;
HANDLE handle;
// Parse the filter:
@@ -500,7 +541,8 @@ extern HANDLE DivertOpen(const char *filter)
// Attempt to open the Divert device:
handle = CreateFile(L"\\\\.\\Divert", GENERIC_READ | GENERIC_WRITE,
- 0, NULL, OPEN_EXISTING, FILE_ATTRIBUTE_NORMAL, INVALID_HANDLE_VALUE);
+ 0, NULL, OPEN_EXISTING, FILE_ATTRIBUTE_NORMAL | FILE_FLAG_OVERLAPPED,
+ INVALID_HANDLE_VALUE);
if (handle == INVALID_HANDLE_VALUE)
{
err = GetLastError();
@@ -524,13 +566,9 @@ extern HANDLE DivertOpen(const char *filter)
}
// Set the filter:
- ioctl.version = DIVERT_VERSION;
- ioctl.magic = DIVERT_MAGIC;
- ioctl.reserved = 0x0;
- ioctl.arg = (UINT64)NULL;
- if (!DeviceIoControl(handle, IOCTL_DIVERT_SET_FILTER, &ioctl,
- sizeof(ioctl), ioctl_filter,
- filter_len*sizeof(struct divert_ioctl_filter_s), &iolen, NULL))
+ if (!DivertIoControl(handle, IOCTL_DIVERT_SET_FILTER, (UINT64)NULL,
+ ioctl_filter, filter_len*sizeof(struct divert_ioctl_filter_s),
+ NULL))
{
CloseHandle(handle);
return INVALID_HANDLE_VALUE;
@@ -546,23 +584,8 @@ extern HANDLE DivertOpen(const char *filter)
extern BOOL DivertRecv(HANDLE handle, PVOID pPacket, UINT packetLen,
PDIVERT_ADDRESS addr, UINT *readlen)
{
- struct divert_ioctl_s ioctl;
- DWORD readlen0;
-
- ioctl.version = DIVERT_VERSION;
- ioctl.magic = DIVERT_MAGIC;
- ioctl.reserved = 0x0;
- ioctl.arg = (UINT64)addr;
- if (!DeviceIoControl(handle, IOCTL_DIVERT_RECV, &ioctl, sizeof(ioctl),
- pPacket, packetLen, &readlen0, NULL))
- {
- return FALSE;
- }
- if (readlen != NULL)
- {
- *readlen = (UINT)readlen0;
- }
- return TRUE;
+ return DivertIoControl(handle, IOCTL_DIVERT_RECV, (UINT64)addr, pPacket,
+ packetLen, readlen);
}
/*
@@ -571,23 +594,8 @@ extern BOOL DivertRecv(HANDLE handle, PVOID pPacket, UINT packetLen,
extern BOOL DivertSend(HANDLE handle, PVOID pPacket, UINT packetLen,
PDIVERT_ADDRESS addr, UINT *writelen)
{
- struct divert_ioctl_s ioctl;
- DWORD writelen0;
-
- ioctl.version = DIVERT_VERSION;
- ioctl.magic = DIVERT_MAGIC;
- ioctl.reserved = 0x0;
- ioctl.arg = (UINT64)addr;
- if (!DeviceIoControl(handle, IOCTL_DIVERT_SEND, &ioctl, sizeof(ioctl),
- pPacket, packetLen, &writelen0, NULL))
- {
- return FALSE;
- }
- if (writelen != NULL)
- {
- *writelen = writelen0;
- }
- return TRUE;
+ return DivertIoControl(handle, IOCTL_DIVERT_SEND, (UINT64)addr, pPacket,
+ packetLen, writelen);
}
/*
diff --git a/doc/divert.html b/doc/divert.html
index cb49714..da07ccb 100644
--- a/doc/divert.html
+++ b/doc/divert.html
@@ -856,6 +856,12 @@ The sample programs are:
other traffic it simply drops.
This is similar to the Linux iptables command with the
-j REJECT option.
+passthru.exe: A simple program that simply re-injects every
+ packet it captures.
+ This example is multi-threaded, where multiple threads are processing
+ packets from a single handle.
+ This example is useful for performance testing, and as a starting point
+ for more interesting applications.
The samples are intended for educational purposes only, and are not
@@ -922,10 +928,6 @@ They are
This is necessary to prevent packet loops and deadlocks.
In the future we intend to implement priorities for divert
handles to allow packets to be seen by multiple divert handles.
-
Speed:
- The divert driver is not re-entrant, and thus is not as
- efficient as it could be.
- In the future we plan to rectify this.
diff --git a/examples/dirs b/examples/dirs
index 92b6118..860b108 100644
--- a/examples/dirs
+++ b/examples/dirs
@@ -1,4 +1,5 @@
DIRS= \
netdump \
netfilter \
+ passthru \
webfilter
diff --git a/examples/passthru/Makefile b/examples/passthru/Makefile
new file mode 100644
index 0000000..53b9a3d
--- /dev/null
+++ b/examples/passthru/Makefile
@@ -0,0 +1 @@
+!INCLUDE $(NTMAKEENV)\makefile.def
diff --git a/examples/passthru/passthru.c b/examples/passthru/passthru.c
new file mode 100644
index 0000000..9b11392
--- /dev/null
+++ b/examples/passthru/passthru.c
@@ -0,0 +1,117 @@
+/*
+ * passthru.c
+ * (C) 2011, all rights reserved,
+ *
+ * This program is free software: you can redistribute it and/or modify
+ * it under the terms of the GNU General Public License as published by
+ * the Free Software Foundation, either version 3 of the License, or
+ * (at your option) any later version.
+ *
+ * This program is distributed in the hope that it will be useful,
+ * but WITHOUT ANY WARRANTY; without even the implied warranty of
+ * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
+ * GNU General Public License for more details.
+ *
+ * You should have received a copy of the GNU General Public License
+ * along with this program. If not, see .
+ */
+
+/*
+ * DESCRIPTION:
+ * This program does nothing except divert packets and re-inject them. This is
+ * useful for performance testing.
+ *
+ * usage: netdump.exe divert-filter num-threads
+ */
+
+#include
+#include
+#include
+#include
+
+#include "divert.h"
+
+#define MAXBUF 2048
+
+static DWORD passthru(LPVOID arg);
+
+/*
+ * Entry.
+ */
+int __cdecl main(int argc, char **argv)
+{
+ int num_threads, i;
+ HANDLE handle, thread;
+
+ if (argc != 3)
+ {
+ fprintf(stderr, "usage: %s filter num-threads\n", argv[0]);
+ exit(EXIT_FAILURE);
+ }
+ num_threads = atoi(argv[2]);
+ if (num_threads < 1 || num_threads > 64)
+ {
+ fprintf(stderr, "error: invalid number of threads\n");
+ exit(EXIT_FAILURE);
+ }
+
+ // Divert traffic matching the filter:
+ handle = DivertOpen(argv[1]);
+ if (handle == INVALID_HANDLE_VALUE)
+ {
+ if (GetLastError() == ERROR_INVALID_PARAMETER)
+ {
+ fprintf(stderr, "error: filter syntax error\n");
+ exit(EXIT_FAILURE);
+ }
+ fprintf(stderr, "error: failed to open Divert device (%d)\n",
+ GetLastError());
+ exit(EXIT_FAILURE);
+ }
+
+ // Start the threads
+ for (i = 1; i < num_threads; i++)
+ {
+ thread = CreateThread(NULL, 1, passthru, (LPVOID)handle, 0, NULL);
+ if (thread == NULL)
+ {
+ fprintf(stderr, "error: failed to start passthru thread (%u)\n",
+ GetLastError());
+ exit(EXIT_FAILURE);
+ }
+ }
+
+ // Main thread:
+ passthru((LPVOID)handle);
+
+ return 0;
+}
+
+// Passthru thread.
+static DWORD passthru(LPVOID arg)
+{
+ char packet[MAXBUF];
+ UINT packet_len;
+ DIVERT_ADDRESS addr;
+ HANDLE handle = (HANDLE)arg;
+
+ // Main loop:
+ while (TRUE)
+ {
+ // Read a matching packet.
+ if (!DivertRecv(handle, packet, sizeof(packet), &addr, &packet_len))
+ {
+ fprintf(stderr, "warning: failed to read packet (%d)\n",
+ GetLastError());
+ continue;
+ }
+
+ // Re-inject the matching packet.
+ if (!DivertSend(handle, packet, packet_len, &addr, NULL))
+ {
+ fprintf(stderr, "warning: failed to reinject packet (%d)\n",
+ GetLastError());
+ }
+ }
+}
+
diff --git a/examples/passthru/sources b/examples/passthru/sources
new file mode 100644
index 0000000..2d39df7
--- /dev/null
+++ b/examples/passthru/sources
@@ -0,0 +1,36 @@
+# sources
+# (C) 2011, all rights reserved,
+#
+# This program is free software: you can redistribute it and/or modify
+# it under the terms of the GNU General Public License as published by
+# the Free Software Foundation, either version 3 of the License, or
+# (at your option) any later version.
+#
+# This program is distributed in the hope that it will be useful,
+# but WITHOUT ANY WARRANTY; without even the implied warranty of
+# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
+# GNU General Public License for more details.
+#
+# You should have received a copy of the GNU General Public License
+# along with this program. If not, see .
+
+!IF "$(_BUILDARCH)" == "x86"
+CPU=i386
+!ELSE
+CPU=$(_BUILDARCH)
+!ENDIF
+
+TARGETNAME=passthru
+TARGETTYPE=PROGRAM
+TARGETPATH=..\..\install
+TARGETLIBS=\
+ $(SDK_LIB_PATH)\setupapi.lib \
+ $(SDK_LIB_PATH)\user32.lib \
+ $(SDK_LIB_PATH)\ws2_32.lib \
+ $(TARGETPATH)\$(CPU)\divert.lib
+UMTYPE=console
+UMENTRY=main
+USE_MSVCRT=1
+INCLUDES=$(DDK_INC_PATH);$(KMDF_INC_PATH)\$(KMDF_VER_PATH);..\..\include
+SOURCES=passthru.c
+
diff --git a/build.sh b/mingw-build.sh
similarity index 98%
rename from build.sh
rename to mingw-build.sh
index 7c677d8..a9185ce 100644
--- a/build.sh
+++ b/mingw-build.sh
@@ -1,6 +1,6 @@
#!/bin/bash
#
-# build.sh
+# mingw-build.sh
# (C) 2011, all rights reserved,
#
# This program is free software: you can redistribute it and/or modify
diff --git a/sys/divert.c b/sys/divert.c
index 13f37e3..d61c0ea 100644
--- a/sys/divert.c
+++ b/sys/divert.c
@@ -124,7 +124,12 @@ struct context_s
// Sublayer GUIDs.
GUID callout_guid[DIVERT_CONTEXT_NUMLAYERS];
// Callout GUIDs.
+ GUID filter_guid[DIVERT_CONTEXT_NUMLAYERS];
+ // Filter GUIDs.
+ BOOL registered[DIVERT_CONTEXT_NUMLAYERS];
+ // What is registered?
HANDLE engine_handle; // WFP engine handle.
+ LONG filter_on; // Is filter on?
struct filter_s filter[DIVERT_FILTER_MAXLEN];
// Packet filter.
};
@@ -286,10 +291,10 @@ static void divert_read_service(context_t context);
static BOOLEAN divert_context_verify(context_t context, context_state_t state);
extern VOID divert_create(IN WDFDEVICE device, IN WDFREQUEST request,
IN WDFFILEOBJECT object);
-extern NTSTATUS divert_register_callout(context_t context, UINT idx,
- wchar_t *sublayer_name, wchar_t *sublayer_desc,
- wchar_t *callout_name, wchar_t *callout_desc,
- wchar_t *filter_name, wchar_t *filter_desc);
+static NTSTATUS divert_register_callouts(context_t context, BOOL is_inbound,
+ BOOL is_outbound, BOOL is_ipv4, BOOL is_ipv6);
+static NTSTATUS divert_register_callout(context_t context, UINT idx,
+ BOOL is_inbound, BOOL is_outbound, BOOL is_ipv4, BOOL is_ipv6);
extern VOID divert_timer(IN WDFTIMER timer);
extern VOID divert_cleanup(IN WDFFILEOBJECT object);
extern VOID divert_close(IN WDFFILEOBJECT object);
@@ -340,6 +345,10 @@ static BOOL divert_filter(PNET_BUFFER buffer, UINT32 if_idx, UINT32 sub_if_idx,
BOOL outbound, filter_t filter);
static BOOL divert_filter_compile(divert_ioctl_filter_t ioctl_filter,
size_t ioctl_filter_len, filter_t filter);
+static void divert_filter_analyze(filter_t filter, BOOL *is_inbound,
+ BOOL *is_outbound, BOOL *ip_ipv4, BOOL *is_ipv6);
+static BOOL divert_filter_test(filter_t filter, UINT8 ip, UINT8 protocol,
+ UINT8 field, UINT32 arg);
/*
* Driver entry routine.
@@ -405,7 +414,7 @@ extern NTSTATUS DriverEntry(IN PDRIVER_OBJECT driver_obj,
return status;
}
WDF_IO_QUEUE_CONFIG_INIT_DEFAULT_QUEUE(&queue_config,
- WdfIoQueueDispatchSequential);
+ WdfIoQueueDispatchParallel);
queue_config.EvtIoRead = NULL;
queue_config.EvtIoWrite = NULL;
queue_config.EvtIoDeviceControl = divert_ioctl;
@@ -487,55 +496,13 @@ static BOOLEAN divert_context_verify(context_t context, context_state_t state)
extern VOID divert_create(IN WDFDEVICE device, IN WDFREQUEST request,
IN WDFFILEOBJECT object)
{
- static wchar_t *sublayer_name[DIVERT_CONTEXT_NUMLAYERS] =
- {
- L"DivertSubLayerOutboundIPv4",
- L"DivertSubLayerInboundIPv4",
- L"DivertSubLayerOutboundIPv6",
- L"DivertSubLayerInboundIPv6"
- };
- static wchar_t *sublayer_desc[DIVERT_CONTEXT_NUMLAYERS] =
- {
- L"Divert sublayer (outbound IPv4)",
- L"Divert sublayer (inbound IPv4)",
- L"Divert sublayer (outbound IPv6)",
- L"Divert sublayer (inbound IPv6)"
- };
- static wchar_t *callout_name[DIVERT_CONTEXT_NUMLAYERS] =
- {
- L"DivertCalloutOutboundIPv4",
- L"DivertCalloutInboundIPv4",
- L"DivertCalloutOutboundIPv6",
- L"DivertCalloutInboundIPv6"
- };
- static wchar_t *callout_desc[DIVERT_CONTEXT_NUMLAYERS] =
- {
- L"Divert callout (outbound IPv4)",
- L"Divert callout (inbound IPv4)",
- L"Divert callout (outbound IPv6)",
- L"Divert callout (inbound IPv6)"
- };
- static wchar_t *filter_name[DIVERT_CONTEXT_NUMLAYERS] =
- {
- L"DivertFilterOutboundIPv4",
- L"DivertFilterInboundIPv4",
- L"DivertFilterOutboundIPv6",
- L"DivertFilterInboundIPv6"
- };
- static wchar_t *filter_desc[DIVERT_CONTEXT_NUMLAYERS] =
- {
- L"Divert filter (outbound IPv4)",
- L"Divert filter (inbound IPv4)",
- L"Divert filter (outbound IPv6)",
- L"Divert filter (inbound IPv6)"
- };
NET_BUFFER_LIST_POOL_PARAMETERS pool_params;
WDF_IO_QUEUE_CONFIG queue_config;
WDF_TIMER_CONFIG timer_config;
WDF_OBJECT_ATTRIBUTES timer_attributes;
FWPM_SESSION0 session;
NTSTATUS status = STATUS_SUCCESS;
- UINT8 i, j = 0;
+ UINT8 i;
context_t context = divert_context_get(object);
DEBUG("CREATE: creating a new divert context (context=%p)", context);
@@ -558,6 +525,11 @@ extern VOID divert_create(IN WDFDEVICE device, IN WDFREQUEST request,
context->filter[i].success = DIVERT_FILTER_RESULT_REJECT;
context->filter[i].failure = DIVERT_FILTER_RESULT_REJECT;
}
+ for (i = 0; i < DIVERT_CONTEXT_NUMLAYERS; i++)
+ {
+ context->registered[i] = FALSE;
+ }
+ context->filter_on = FALSE;
KeInitializeSpinLock(&context->lock);
InitializeListHead(&context->packet_queue);
for (i = 0; i < DIVERT_CONTEXT_NUMLAYERS; i++)
@@ -574,6 +546,12 @@ extern VOID divert_create(IN WDFDEVICE device, IN WDFREQUEST request,
DEBUG_ERROR("failed to create callout GUID", status);
goto divert_create_exit;
}
+ status = ExUuidCreate(&context->filter_guid[i]);
+ if (!NT_SUCCESS(status))
+ {
+ DEBUG_ERROR("failed to create filter GUID", status);
+ goto divert_create_exit;
+ }
}
RtlZeroMemory(&pool_params, sizeof(pool_params));
pool_params.Header.Type = NDIS_OBJECT_TYPE_DEFAULT;
@@ -617,31 +595,6 @@ extern VOID divert_create(IN WDFDEVICE device, IN WDFREQUEST request,
DEBUG_ERROR("failed to create WFP engine handle", status);
goto divert_create_exit;
}
- status = FwpmTransactionBegin0(context->engine_handle, 0);
- if (!NT_SUCCESS(status))
- {
- DEBUG_ERROR("failed to begin WFP transaction", status);
- goto divert_create_exit;
- }
- for (j = 0; j < DIVERT_CONTEXT_NUMLAYERS; j++)
- {
- status = divert_register_callout(context, j, sublayer_name[j],
- sublayer_desc[j], callout_name[j], callout_desc[j], filter_name[j],
- filter_desc[j]);
- if (!NT_SUCCESS(status))
- {
- FwpmTransactionAbort0(context->engine_handle);
- goto divert_create_exit;
- }
- }
- status = FwpmTransactionCommit0(context->engine_handle);
- if (!NT_SUCCESS(status))
- {
- DEBUG_ERROR("failed to commit WFP transaction", status);
- goto divert_create_exit;
- }
-
- // Open for business:
context->state = DIVERT_CONTEXT_STATE_OPEN;
WdfTimerStart(context->timer,
WDF_REL_TIMEOUT_IN_MS(DIVERT_PACKET_TIMEOUT));
@@ -667,10 +620,6 @@ divert_create_exit:
{
FwpmEngineClose0(context->engine_handle);
}
- for (i = 0; i < j; i++)
- {
- FwpsCalloutUnregisterByKey0(&context->callout_guid[i]);
- }
context->state = DIVERT_CONTEXT_STATE_INVALID;
}
@@ -678,37 +627,130 @@ divert_create_exit:
}
/*
- * Add a WFP filter.
+ * Register all WFP callouts.
*/
-extern NTSTATUS divert_register_callout(context_t context, UINT idx,
- wchar_t *sublayer_name, wchar_t *sublayer_desc,
- wchar_t *callout_name, wchar_t *callout_desc,
- wchar_t *filter_name, wchar_t *filter_desc)
+static NTSTATUS divert_register_callouts(context_t context, BOOL is_inbound,
+ BOOL is_outbound, BOOL is_ipv4, BOOL is_ipv6)
{
+ UINT8 i;
+ NTSTATUS status;
+
+ status = FwpmTransactionBegin0(context->engine_handle, 0);
+ if (!NT_SUCCESS(status))
+ {
+ DEBUG_ERROR("failed to begin WFP transaction", status);
+ goto divert_register_callouts_exit;
+ }
+ for (i = 0; i < DIVERT_CONTEXT_NUMLAYERS; i++)
+ {
+ status = divert_register_callout(context, i, is_inbound, is_outbound,
+ is_ipv4, is_ipv6);
+ if (!NT_SUCCESS(status))
+ {
+ FwpmTransactionAbort0(context->engine_handle);
+ goto divert_register_callouts_exit;
+ }
+ }
+ status = FwpmTransactionCommit0(context->engine_handle);
+ if (!NT_SUCCESS(status))
+ {
+ DEBUG_ERROR("failed to commit WFP transaction", status);
+ goto divert_register_callouts_exit;
+ }
+
+divert_register_callouts_exit:
+
+ if (!NT_SUCCESS(status))
+ {
+ for (i = 0; i < DIVERT_CONTEXT_NUMLAYERS; i++)
+ {
+ if (context->registered[i])
+ {
+ FwpsCalloutUnregisterByKey0(&context->callout_guid[i]);
+ context->registered[i] = FALSE;
+ }
+ }
+ }
+
+ return status;
+}
+
+/*
+ * Register a WFP callout.
+ */
+static NTSTATUS divert_register_callout(context_t context, UINT idx,
+ BOOL is_inbound, BOOL is_outbound, BOOL is_ipv4, BOOL is_ipv6)
+{
+ static wchar_t *sublayer_name[DIVERT_CONTEXT_NUMLAYERS] =
+ {
+ L"DivertSubLayerOutboundIPv4",
+ L"DivertSubLayerInboundIPv4",
+ L"DivertSubLayerOutboundIPv6",
+ L"DivertSubLayerInboundIPv6"
+ };
+ static wchar_t *sublayer_desc[DIVERT_CONTEXT_NUMLAYERS] =
+ {
+ L"Divert sublayer (outbound IPv4)",
+ L"Divert sublayer (inbound IPv4)",
+ L"Divert sublayer (outbound IPv6)",
+ L"Divert sublayer (inbound IPv6)"
+ };
+ static wchar_t *callout_name[DIVERT_CONTEXT_NUMLAYERS] =
+ {
+ L"DivertCalloutOutboundIPv4",
+ L"DivertCalloutInboundIPv4",
+ L"DivertCalloutOutboundIPv6",
+ L"DivertCalloutInboundIPv6"
+ };
+ static wchar_t *callout_desc[DIVERT_CONTEXT_NUMLAYERS] =
+ {
+ L"Divert callout (outbound IPv4)",
+ L"Divert callout (inbound IPv4)",
+ L"Divert callout (outbound IPv6)",
+ L"Divert callout (inbound IPv6)"
+ };
+ static wchar_t *filter_name[DIVERT_CONTEXT_NUMLAYERS] =
+ {
+ L"DivertFilterOutboundIPv4",
+ L"DivertFilterInboundIPv4",
+ L"DivertFilterOutboundIPv6",
+ L"DivertFilterInboundIPv6"
+ };
+ static wchar_t *filter_desc[DIVERT_CONTEXT_NUMLAYERS] =
+ {
+ L"Divert filter (outbound IPv4)",
+ L"Divert filter (inbound IPv4)",
+ L"Divert filter (outbound IPv6)",
+ L"Divert filter (inbound IPv6)"
+ };
GUID layer;
FWPM_SUBLAYER0 sublayer;
FWPS_CALLOUT0 scallout;
FWPM_CALLOUT0 mcallout;
FWPM_FILTER0 filter;
- BOOL registered = FALSE;
+ BOOL required, registered = FALSE;
divert_callout_t callout;
NTSTATUS status;
switch (idx)
{
case DIVERT_CONTEXT_OUTBOUND_IPV4_LAYER:
+ required = (is_outbound && is_ipv4);
layer = FWPM_LAYER_OUTBOUND_IPPACKET_V4;
callout = divert_classify_outbound_v4_callout;
break;
case DIVERT_CONTEXT_INBOUND_IPV4_LAYER:
+ required = (is_inbound && is_ipv4);
layer = FWPM_LAYER_INBOUND_IPPACKET_V4;
callout = divert_classify_inbound_v4_callout;
break;
case DIVERT_CONTEXT_OUTBOUND_IPV6_LAYER:
+ required = (is_outbound && is_ipv6);
layer = FWPM_LAYER_OUTBOUND_IPPACKET_V6;
callout = divert_classify_outbound_v6_callout;
break;
case DIVERT_CONTEXT_INBOUND_IPV6_LAYER:
+ required = (is_inbound && is_ipv6);
layer = FWPM_LAYER_INBOUND_IPPACKET_V6;
callout = divert_classify_inbound_v6_callout;
break;
@@ -716,10 +758,15 @@ extern NTSTATUS divert_register_callout(context_t context, UINT idx,
return STATUS_INVALID_PARAMETER;
}
+ if (!required)
+ {
+ return STATUS_SUCCESS;
+ }
+
RtlZeroMemory(&sublayer, sizeof(sublayer));
sublayer.subLayerKey = context->sublayer_guid[idx];
- sublayer.displayData.name = sublayer_name;
- sublayer.displayData.description = sublayer_desc;
+ sublayer.displayData.name = sublayer_name[idx];
+ sublayer.displayData.description = sublayer_desc[idx];
sublayer.weight = FWP_EMPTY;
RtlZeroMemory(&scallout, sizeof(scallout));
scallout.calloutKey = context->callout_guid[idx];
@@ -728,13 +775,14 @@ extern NTSTATUS divert_register_callout(context_t context, UINT idx,
scallout.flowDeleteFn = NULL;
RtlZeroMemory(&mcallout, sizeof(mcallout));
mcallout.calloutKey = context->callout_guid[idx];
- mcallout.displayData.name = callout_name;
- mcallout.displayData.description = callout_desc;
+ mcallout.displayData.name = callout_name[idx];
+ mcallout.displayData.description = callout_desc[idx];
mcallout.applicableLayer = layer;
RtlZeroMemory(&filter, sizeof(filter));
+ filter.filterKey = context->filter_guid[idx];
filter.layerKey = layer;
- filter.displayData.name = filter_name;
- filter.displayData.description = filter_desc;
+ filter.displayData.name = filter_name[idx];
+ filter.displayData.description = filter_desc[idx];
filter.action.type = FWP_ACTION_CALLOUT_TERMINATING;
filter.action.calloutKey = context->callout_guid[idx];
filter.subLayerKey = context->sublayer_guid[idx];
@@ -766,6 +814,7 @@ extern NTSTATUS divert_register_callout(context_t context, UINT idx,
DEBUG_ERROR("failed to add WFP filter", status);
goto divert_register_callout_error;
}
+ context->registered[idx] = TRUE;
return STATUS_SUCCESS;
@@ -833,6 +882,7 @@ extern VOID divert_cleanup(IN WDFFILEOBJECT object)
UINT i;
context_t context = divert_context_get(object);
packet_t packet;
+ NTSTATUS status;
DEBUG("CLEANUP: cleaning up divert context (context=%p)", context);
@@ -856,10 +906,51 @@ extern VOID divert_cleanup(IN WDFFILEOBJECT object)
WdfIoQueuePurge(context->read_queue, NULL, NULL);
WdfObjectDelete(context->read_queue);
WdfObjectDelete(context->timer);
+
+ status = FwpmTransactionBegin0(context->engine_handle, 0);
+ if (!NT_SUCCESS(status))
+ {
+ DEBUG_ERROR("failed to begin WFP transaction", status);
+ goto divert_cleanup_exit;
+ }
+ for (i = 0; i < DIVERT_CONTEXT_NUMLAYERS; i++)
+ {
+ if (!context->registered[i])
+ {
+ continue;
+ }
+ status = FwpmFilterDeleteByKey0(context->engine_handle,
+ context->filter_guid+i);
+ if (!NT_SUCCESS(status))
+ {
+ DEBUG_ERROR("failed delete WFP filter", status);
+ FwpmTransactionAbort0(context->engine_handle);
+ goto divert_cleanup_exit;
+ }
+ status = FwpmSubLayerDeleteByKey0(context->engine_handle,
+ context->sublayer_guid+i);
+ if (!NT_SUCCESS(status))
+ {
+ DEBUG_ERROR("failed delete WFP sub-layer", status);
+ FwpmTransactionAbort0(context->engine_handle);
+ goto divert_cleanup_exit;
+ }
+ }
+ status = FwpmTransactionCommit0(context->engine_handle);
+ if (!NT_SUCCESS(status))
+ {
+ DEBUG_ERROR("failed to commit WFP transaction", status);
+ goto divert_cleanup_exit;
+ }
+
+divert_cleanup_exit:
FwpmEngineClose0(context->engine_handle);
for (i = 0; i < DIVERT_CONTEXT_NUMLAYERS; i++)
{
- FwpsCalloutUnregisterByKey0(&context->callout_guid[i]);
+ if (context->registered[i])
+ {
+ FwpsCalloutUnregisterByKey0(&context->callout_guid[i]);
+ }
}
NdisFreeNetBufferPool(context->pool_handle);
}
@@ -1327,6 +1418,16 @@ extern VOID divert_ioctl(IN WDFQUEUE queue, IN WDFREQUEST request,
break;
case IOCTL_DIVERT_SET_FILTER:
+ {
+ BOOL is_inbound, is_outbound, is_ipv4, is_ipv6;
+
+ if (InterlockedExchange(&context->filter_on, TRUE) == TRUE)
+ {
+ status = STATUS_INVALID_DEVICE_REQUEST;
+ DEBUG_ERROR("duplicate SET_FILTER ioctl", status);
+ goto divert_ioctl_exit;
+ }
+
filter = (divert_ioctl_filter_t)outbuf;
filter_len = outbuflen;
if (!divert_filter_compile(filter, filter_len, context->filter))
@@ -1335,8 +1436,14 @@ extern VOID divert_ioctl(IN WDFQUEUE queue, IN WDFREQUEST request,
DEBUG_ERROR("failed to compile filter", status);
goto divert_ioctl_exit;
}
- break;
+ divert_filter_analyze(context->filter, &is_inbound, &is_outbound,
+ &is_ipv4, &is_ipv6);
+ status = divert_register_callouts(context, is_inbound,
+ is_outbound, is_ipv4, is_ipv6);
+
+ break;
+ }
default:
status = STATUS_INVALID_DEVICE_REQUEST;
DEBUG_ERROR("failed to complete I/O control; invalid request",
@@ -2303,6 +2410,137 @@ static BOOL divert_filter(PNET_BUFFER buffer, UINT32 if_idx, UINT32 sub_if_idx,
return FALSE;
}
+/*
+ * Analyze the given filter.
+ */
+static void divert_filter_analyze(filter_t filter, BOOL *is_inbound,
+ BOOL *is_outbound, BOOL *is_ipv4, BOOL *is_ipv6)
+{
+ BOOL result;
+
+ // False filter?
+ result = divert_filter_test(filter, 0, DIVERT_FILTER_PROTOCOL_NONE,
+ DIVERT_FILTER_FIELD_ZERO, 0);
+ if (!result)
+ {
+ *is_inbound = FALSE;
+ *is_outbound = FALSE;
+ *is_ipv4 = FALSE;
+ *is_ipv6 = FALSE;
+ return;
+ }
+
+ // Inbound?
+ result = divert_filter_test(filter, 0, DIVERT_FILTER_PROTOCOL_NONE,
+ DIVERT_FILTER_FIELD_INBOUND, 1);
+ if (result)
+ {
+ result = divert_filter_test(filter, 0, DIVERT_FILTER_PROTOCOL_NONE,
+ DIVERT_FILTER_FIELD_OUTBOUND, 0);
+ }
+ *is_inbound = result;
+
+ // Outbound?
+ result = divert_filter_test(filter, 0, DIVERT_FILTER_PROTOCOL_NONE,
+ DIVERT_FILTER_FIELD_OUTBOUND, 1);
+ if (result)
+ {
+ result = divert_filter_test(filter, 0, DIVERT_FILTER_PROTOCOL_NONE,
+ DIVERT_FILTER_FIELD_INBOUND, 0);
+ }
+ *is_outbound = result;
+
+ // IPv4?
+ result = divert_filter_test(filter, 0, DIVERT_FILTER_PROTOCOL_NONE,
+ DIVERT_FILTER_FIELD_IP, 1);
+ if (result)
+ {
+ result = divert_filter_test(filter, 0, DIVERT_FILTER_PROTOCOL_NONE,
+ DIVERT_FILTER_FIELD_IPV6, 0);
+ }
+ *is_ipv4 = result;
+
+ // Ipv6?
+ result = divert_filter_test(filter, 0, DIVERT_FILTER_PROTOCOL_NONE,
+ DIVERT_FILTER_FIELD_IPV6, 1);
+ if (result)
+ {
+ result = divert_filter_test(filter, 0, DIVERT_FILTER_PROTOCOL_NONE,
+ DIVERT_FILTER_FIELD_IP, 0);
+ }
+ *is_ipv6 = result;
+}
+
+/*
+ * Test a filter for any packet where field = arg.
+ */
+static BOOL divert_filter_test(filter_t filter, UINT8 ip, UINT8 protocol,
+ UINT8 field, UINT32 arg)
+{
+ BOOL known = FALSE;
+ BOOL result = FALSE;
+
+ if (ip == DIVERT_FILTER_RESULT_ACCEPT)
+ {
+ return TRUE;
+ }
+ if (ip == DIVERT_FILTER_RESULT_REJECT)
+ {
+ return FALSE;
+ }
+ if (ip > DIVERT_FILTER_MAXLEN)
+ {
+ return FALSE;
+ }
+
+ if (filter[ip].protocol == protocol &&
+ filter[ip].field == field)
+ {
+ known = TRUE;
+ switch (filter[ip].test)
+ {
+ case DIVERT_FILTER_TEST_EQ:
+ result = (arg == filter[ip].arg[0]);
+ break;
+ case DIVERT_FILTER_TEST_NEQ:
+ result = (arg != filter[ip].arg[0]);
+ break;
+ case DIVERT_FILTER_TEST_LT:
+ result = (arg < filter[ip].arg[0]);
+ break;
+ case DIVERT_FILTER_TEST_LEQ:
+ result = (arg <= filter[ip].arg[0]);
+ break;
+ case DIVERT_FILTER_TEST_GT:
+ result = (arg > filter[ip].arg[0]);
+ break;
+ case DIVERT_FILTER_TEST_GEQ:
+ result = (arg >= filter[ip].arg[0]);
+ break;
+ default:
+ result = FALSE;
+ break;
+ }
+ }
+
+ if (!known)
+ {
+ result = divert_filter_test(filter, filter[ip].success, protocol,
+ field, arg);
+ if (result)
+ {
+ return TRUE;
+ }
+ return divert_filter_test(filter, filter[ip].failure, protocol, field,
+ arg);
+ }
+ else
+ {
+ ip = (result? filter[ip].success: filter[ip].failure);
+ return divert_filter_test(filter, ip, protocol, field, arg);
+ }
+}
+
/*
* Compile a divert filter from an IOCTL.
*/