diff --git a/dll/windivert.c b/dll/windivert.c
index e593e31..62208c2 100644
--- a/dll/windivert.c
+++ b/dll/windivert.c
@@ -119,13 +119,10 @@ static BOOL WinDivertIoControlEx(HANDLE handle, DWORD code, UINT8 arg8,
UINT64 arg, PVOID buf, UINT len, UINT *iolen, LPOVERLAPPED overlapped);
static UINT8 WinDivertSkipExtHeaders(UINT8 proto, UINT8 **header, UINT *len);
-#ifdef WINDIVERT_DEBUG
-static void WinDivertFilterDump(windivert_ioctl_filter_t filter, UINT16 len);
-#endif
-
/*
* Include the helper API implementation.
*/
+#include "windivert_shared.c"
#include "windivert_helper.c"
/*
@@ -379,7 +376,7 @@ static BOOL WinDivertIoControl(HANDLE handle, DWORD code, UINT8 arg8,
static BOOL WinDivertIoControlEx(HANDLE handle, DWORD code, UINT8 arg8,
UINT64 arg, PVOID buf, UINT len, UINT *iolen, LPOVERLAPPED overlapped)
{
- struct windivert_ioctl_s ioctl;
+ WINDIVERT_IOCTL ioctl;
BOOL result;
DWORD iolen0;
@@ -402,24 +399,21 @@ static BOOL WinDivertIoControlEx(HANDLE handle, DWORD code, UINT8 arg8,
extern HANDLE WinDivertOpen(const char *filter, WINDIVERT_LAYER layer,
INT16 priority, UINT64 flags)
{
- struct windivert_ioctl_filter_s object[WINDIVERT_FILTER_MAXLEN];
+ WINDIVERT_FILTER object[WINDIVERT_FILTER_MAXLEN];
UINT obj_len;
ERROR comp_err;
DWORD err;
HANDLE handle;
SC_HANDLE service;
- UINT32 priority32;
-
+ UINT64 priority64, filter_flags;
+
// Parameter checking.
- if (layer == 0)
- {
- layer = WINDIVERT_LAYER_NETWORK;
- }
switch (layer)
{
case WINDIVERT_LAYER_NETWORK:
case WINDIVERT_LAYER_NETWORK_FORWARD:
case WINDIVERT_LAYER_FLOW:
+ case WINDIVERT_LAYER_REFLECT:
break;
default:
SetLastError(ERROR_INVALID_PARAMETER);
@@ -431,25 +425,21 @@ extern HANDLE WinDivertOpen(const char *filter, WINDIVERT_LAYER layer,
return INVALID_HANDLE_VALUE;
}
- priority32 = WINDIVERT_PRIORITY(priority);
- if (priority32 < WINDIVERT_PRIORITY_MIN ||
- priority32 > WINDIVERT_PRIORITY_MAX)
+ if (priority < WINDIVERT_PRIORITY_MIN ||
+ priority > WINDIVERT_PRIORITY_MAX)
{
SetLastError(ERROR_INVALID_PARAMETER);
return INVALID_HANDLE_VALUE;
}
- // Compile the filter:
+ // Compile & analyze the filter:
comp_err = WinDivertCompileFilter(filter, layer, object, &obj_len);
if (IS_ERROR(comp_err))
{
SetLastError(ERROR_INVALID_PARAMETER);
return INVALID_HANDLE_VALUE;
}
-
-#ifdef WINDIVERT_DEBUG
- WinDivertFilterDump(object, obj_len);
-#endif
+ filter_flags = WinDivertAnalyzeFilter(object, obj_len);
// Attempt to open the WinDivert device:
handle = CreateFile(L"\\\\.\\" WINDIVERT_DEVICE_NAME,
@@ -464,6 +454,11 @@ extern HANDLE WinDivertOpen(const char *filter, WINDIVERT_LAYER layer,
}
// Open failed because the device isn't installed; install it now.
+ if ((flags & WINDIVERT_FLAG_NO_INSTALL) != 0)
+ {
+ SetLastError(ERROR_SERVICE_DOES_NOT_EXIST);
+ return INVALID_HANDLE_VALUE;
+ }
SetLastError(0);
service = WinDivertDriverInstall();
if (service == NULL)
@@ -503,8 +498,8 @@ extern HANDLE WinDivertOpen(const char *filter, WINDIVERT_LAYER layer,
// Set the flags:
if (flags != 0)
{
- if (!WinDivertIoControl(handle, IOCTL_WINDIVERT_SET_FLAGS, 0,
- (UINT64)flags, NULL, 0, NULL))
+ if (!WinDivertIoControl(handle, IOCTL_WINDIVERT_SET_FLAGS, 0, flags,
+ NULL, 0, NULL))
{
CloseHandle(handle);
return INVALID_HANDLE_VALUE;
@@ -512,10 +507,12 @@ extern HANDLE WinDivertOpen(const char *filter, WINDIVERT_LAYER layer,
}
// Set the priority:
- if (priority32 != WINDIVERT_PRIORITY_DEFAULT)
+ if (priority != WINDIVERT_PRIORITY_DEFAULT)
{
+ // Make positive:
+ priority64 = (UINT64)((INT64)priority + WINDIVERT_PRIORITY_MAX);
if (!WinDivertIoControl(handle, IOCTL_WINDIVERT_SET_PRIORITY, 0,
- (UINT64)priority32, NULL, 0, NULL))
+ priority64, NULL, 0, NULL))
{
CloseHandle(handle);
return INVALID_HANDLE_VALUE;
@@ -523,8 +520,8 @@ extern HANDLE WinDivertOpen(const char *filter, WINDIVERT_LAYER layer,
}
// Start the filter:
- if (!WinDivertIoControl(handle, IOCTL_WINDIVERT_START_FILTER, 0, 0,
- object, obj_len*sizeof(struct windivert_ioctl_filter_s), NULL))
+ if (!WinDivertIoControl(handle, IOCTL_WINDIVERT_START_FILTER, 0,
+ filter_flags, object, obj_len * sizeof(WINDIVERT_FILTER), NULL))
{
CloseHandle(handle);
return INVALID_HANDLE_VALUE;
@@ -856,253 +853,3 @@ static BOOLEAN WinDivertAToX(const char *str, char **endptr, UINT32 *intptr)
return TRUE;
}
-/***************************************************************************/
-/* DEBUGGING */
-/***************************************************************************/
-
-#ifdef WINDIVERT_DEBUG
-/*
- * Print a filter (debugging).
- */
-static void WinDivertFilterDump(windivert_ioctl_filter_t filter, UINT16 len)
-{
- UINT16 i;
-
- for (i = 0; i < len; i++)
- {
- printf("label_%u:\n\tif (", i);
- switch (filter[i].field)
- {
- case WINDIVERT_FILTER_FIELD_ZERO:
- printf("zero ");
- break;
- case WINDIVERT_FILTER_FIELD_INBOUND:
- printf("inbound ");
- break;
- case WINDIVERT_FILTER_FIELD_OUTBOUND:
- printf("outbound ");
- break;
- case WINDIVERT_FILTER_FIELD_IFIDX:
- printf("ifIdx ");
- break;
- case WINDIVERT_FILTER_FIELD_SUBIFIDX:
- printf("subIfIdx ");
- break;
- case WINDIVERT_FILTER_FIELD_IP:
- printf("ip ");
- break;
- case WINDIVERT_FILTER_FIELD_IPV6:
- printf("ipv6 ");
- break;
- case WINDIVERT_FILTER_FIELD_ICMP:
- printf("icmp ");
- break;
- case WINDIVERT_FILTER_FIELD_ICMPV6:
- printf("icmpv6 ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP:
- printf("tcp ");
- break;
- case WINDIVERT_FILTER_FIELD_UDP:
- printf("udp ");
- break;
- case WINDIVERT_FILTER_FIELD_IP_HDRLENGTH:
- printf("ip.HdrLength ");
- break;
- case WINDIVERT_FILTER_FIELD_IP_TOS:
- printf("ip.TOS ");
- break;
- case WINDIVERT_FILTER_FIELD_IP_LENGTH:
- printf("ip.Length ");
- break;
- case WINDIVERT_FILTER_FIELD_IP_ID:
- printf("ip.Id ");
- break;
- case WINDIVERT_FILTER_FIELD_IP_DF:
- printf("ip.DF ");
- break;
- case WINDIVERT_FILTER_FIELD_IP_MF:
- printf("ip.MF ");
- break;
- case WINDIVERT_FILTER_FIELD_IP_FRAGOFF:
- printf("ip.FragOff ");
- break;
- case WINDIVERT_FILTER_FIELD_IP_TTL:
- printf("ip.TTL ");
- break;
- case WINDIVERT_FILTER_FIELD_IP_PROTOCOL:
- printf("ip.Protocol ");
- break;
- case WINDIVERT_FILTER_FIELD_IP_CHECKSUM:
- printf("ip.Checksum ");
- break;
- case WINDIVERT_FILTER_FIELD_IP_SRCADDR:
- printf("ip.SrcAddr ");
- break;
- case WINDIVERT_FILTER_FIELD_IP_DSTADDR:
- printf("ip.DstAddr ");
- break;
- case WINDIVERT_FILTER_FIELD_IPV6_TRAFFICCLASS:
- printf("ipv6.TrafficClass ");
- break;
- case WINDIVERT_FILTER_FIELD_IPV6_FLOWLABEL:
- printf("ipv6.FlowLabel ");
- break;
- case WINDIVERT_FILTER_FIELD_IPV6_LENGTH:
- printf("ipv6.Length ");
- break;
- case WINDIVERT_FILTER_FIELD_IPV6_NEXTHDR:
- printf("ipv6.NextHdr ");
- break;
- case WINDIVERT_FILTER_FIELD_IPV6_HOPLIMIT:
- printf("ipv6.HopLimit ");
- break;
- case WINDIVERT_FILTER_FIELD_IPV6_SRCADDR:
- printf("ipv6.SrcAddr ");
- break;
- case WINDIVERT_FILTER_FIELD_IPV6_DSTADDR:
- printf("ipv6.DstAddr ");
- break;
- case WINDIVERT_FILTER_FIELD_ICMP_TYPE:
- printf("icmp.Type ");
- break;
- case WINDIVERT_FILTER_FIELD_ICMP_CODE:
- printf("icmp.Code ");
- break;
- case WINDIVERT_FILTER_FIELD_ICMP_CHECKSUM:
- printf("icmp.Checksum ");
- break;
- case WINDIVERT_FILTER_FIELD_ICMP_BODY:
- printf("icmp.Body ");
- break;
- case WINDIVERT_FILTER_FIELD_ICMPV6_TYPE:
- printf("icmpv6.Type ");
- break;
- case WINDIVERT_FILTER_FIELD_ICMPV6_CODE:
- printf("icmpv6.Code ");
- break;
- case WINDIVERT_FILTER_FIELD_ICMPV6_CHECKSUM:
- printf("icmpv6.Checksum ");
- break;
- case WINDIVERT_FILTER_FIELD_ICMPV6_BODY:
- printf("icmpv6.Body ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_SRCPORT:
- printf("tcp.SrcPort ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_DSTPORT:
- printf("tcp.DstPort ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_SEQNUM:
- printf("tcp.SeqNum ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_ACKNUM:
- printf("tcp.AckNum ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_HDRLENGTH:
- printf("tcp.HdrLength ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_URG:
- printf("tcp.Urg ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_ACK:
- printf("tcp.Ack ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_PSH:
- printf("tcp.Psh ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_RST:
- printf("tcp.Rst ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_SYN:
- printf("tcp.Syn ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_FIN:
- printf("tcp.Fin ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_WINDOW:
- printf("tcp.Window ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_CHECKSUM:
- printf("tcp.Checksum ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_URGPTR:
- printf("tcp.UrgPtr ");
- break;
- case WINDIVERT_FILTER_FIELD_TCP_PAYLOADLENGTH:
- printf("tcp.PayloadLength " );
- break;
- case WINDIVERT_FILTER_FIELD_UDP_SRCPORT:
- printf("udp.SrcPort ");
- break;
- case WINDIVERT_FILTER_FIELD_UDP_DSTPORT:
- printf("udp.DstPort ");
- break;
- case WINDIVERT_FILTER_FIELD_UDP_LENGTH:
- printf("udp.Length ");
- break;
- case WINDIVERT_FILTER_FIELD_UDP_CHECKSUM:
- printf("udp.Checksum ");
- break;
- case WINDIVERT_FILTER_FIELD_UDP_PAYLOADLENGTH:
- printf("udp.PayloadLength ");
- break;
- default:
- printf("unknown.Field ");
- break;
- }
- switch (filter[i].test)
- {
- case WINDIVERT_FILTER_TEST_EQ:
- printf("== ");
- break;
- case WINDIVERT_FILTER_TEST_NEQ:
- printf("!= ");
- break;
- case WINDIVERT_FILTER_TEST_LT:
- printf("< ");
- break;
- case WINDIVERT_FILTER_TEST_LEQ:
- printf("<= ");
- break;
- case WINDIVERT_FILTER_TEST_GT:
- printf("> ");
- break;
- case WINDIVERT_FILTER_TEST_GEQ:
- printf(">= ");
- break;
- default:
- printf("?? ");
- break;
- }
- printf("%u)\n", filter[i].arg[0]);
- switch (filter[i].success)
- {
- case WINDIVERT_FILTER_RESULT_ACCEPT:
- printf("\t\treturn ACCEPT;\n");
- break;
- case WINDIVERT_FILTER_RESULT_REJECT:
- printf("\t\treturn REJECT;\n");
- break;
- default:
- printf("\t\tgoto label_%u;\n", filter[i].success);
- break;
- }
- printf("\telse\n");
- switch (filter[i].failure)
- {
- case WINDIVERT_FILTER_RESULT_ACCEPT:
- printf("\t\treturn ACCEPT;\n");
- break;
- case WINDIVERT_FILTER_RESULT_REJECT:
- printf("\t\treturn REJECT;\n");
- break;
- default:
- printf("\t\tgoto label_%u;\n", filter[i].failure);
- break;
- }
- }
-}
-
-#endif /* WINDIVERT_DEBUG */
-
diff --git a/dll/windivert.def b/dll/windivert.def
index 2c20d03..a36cd4a 100644
--- a/dll/windivert.def
+++ b/dll/windivert.def
@@ -13,5 +13,6 @@ EXPORTS
WinDivertHelperParsePacket
WinDivertHelperParseIPv4Address
WinDivertHelperParseIPv6Address
- WinDivertHelperCheckFilter
+ WinDivertHelperCompileFilter
WinDivertHelperEvalFilter
+ WinDivertHelperFormatFilter
diff --git a/dll/windivert_helper.c b/dll/windivert_helper.c
index 2b0c40a..86b39b6 100644
--- a/dll/windivert_helper.c
+++ b/dll/windivert_helper.c
@@ -108,6 +108,7 @@ typedef enum
TOKEN_UDP_LENGTH,
TOKEN_UDP_PAYLOAD_LENGTH,
TOKEN_UDP_SRC_PORT,
+ TOKEN_ZERO,
TOKEN_TRUE,
TOKEN_FALSE,
TOKEN_INBOUND,
@@ -122,6 +123,11 @@ typedef enum
TOKEN_LOCAL_PORT,
TOKEN_REMOTE_PORT,
TOKEN_PROTOCOL,
+ TOKEN_LAYER,
+ TOKEN_FLOW,
+ TOKEN_NETWORK,
+ TOKEN_NETWORK_FORWARD,
+ TOKEN_REFLECT,
TOKEN_OPEN,
TOKEN_CLOSE,
TOKEN_EQ,
@@ -166,6 +172,7 @@ struct EXPR
PEXPR arg[3];
};
UINT8 kind;
+ UINT8 count;
UINT16 succ;
UINT16 fail;
};
@@ -174,7 +181,7 @@ struct EXPR
* Error handling.
*/
#undef ERROR
-typedef UINT64 ERROR;
+typedef UINT64 ERROR, *PERROR;
#define WINDIVERT_ERROR_NONE 0
#define WINDIVERT_ERROR_NO_MEMORY 1
@@ -184,7 +191,8 @@ typedef UINT64 ERROR;
#define WINDIVERT_ERROR_BAD_TOKEN_FOR_LAYER 5
#define WINDIVERT_ERROR_UNEXPECTED_TOKEN 6
#define WINDIVERT_ERROR_OUTPUT_TOO_SHORT 7
-#define WINDIVERT_ERROR_ASSERTION_FAILED 8
+#define WINDIVERT_ERROR_BAD_OBJECT 8
+#define WINDIVERT_ERROR_ASSERTION_FAILED 9
#define MAKE_ERROR(code, pos) \
(((ERROR)(code) << 32) | (ERROR)(pos));
@@ -198,26 +206,22 @@ typedef UINT64 ERROR;
#define MAX(a, b) ((a) > (b)? (a): (b))
-/*
- * Compiler memory pool:
- */
-typedef struct POOL
-{
- unsigned offset;
- ERROR error;
- char memory[3 * 4096 - 32];
-} POOL, *PPOOL;
-
/*
* Prototypes.
*/
-static PEXPR WinDivertParseFilter(PPOOL pool, TOKEN *toks, UINT *i, INT depth,
- BOOL and);
+static PEXPR WinDivertParseFilter(HANDLE pool, TOKEN *toks, UINT *i,
+ INT depth, BOOL and, PERROR error);
+static BOOL WinDivertCondExecFilter(PWINDIVERT_FILTER filter, UINT length,
+ UINT8 field, UINT32 arg);
static UINT16 WinDivertInitPseudoHeader(PWINDIVERT_IPHDR ip_header,
PWINDIVERT_IPV6HDR ipv6_header, UINT8 protocol, UINT len,
void *pseudo_header);
static UINT16 WinDivertHelperCalcChecksum(PVOID pseudo_header,
UINT16 pseudo_header_len, PVOID data, UINT len);
+static BOOL WinDivertDeserializeFilter(PWINDIVERT_STREAM stream,
+ PWINDIVERT_FILTER filter, UINT *length);
+static void WinDivertFormatExpr(PWINDIVERT_STREAM stream, PEXPR expr,
+ BOOL top_level, BOOL and);
/*
* Skip well-known IPv6 extension headers.
@@ -810,6 +814,12 @@ static BOOL WinDivertCheckTokenKindForLayer(WINDIVERT_LAYER layer, KIND kind)
case TOKEN_REMOTE_ADDR:
case TOKEN_LOCAL_PORT:
case TOKEN_REMOTE_PORT:
+ case TOKEN_PROTOCOL:
+ case TOKEN_LAYER:
+ case TOKEN_FLOW:
+ case TOKEN_NETWORK:
+ case TOKEN_NETWORK_FORWARD:
+ case TOKEN_REFLECT:
return FALSE;
default:
return TRUE;
@@ -867,10 +877,110 @@ static BOOL WinDivertCheckTokenKindForLayer(WINDIVERT_LAYER layer, KIND kind)
case TOKEN_IF_IDX:
case TOKEN_SUB_IF_IDX:
case TOKEN_IMPOSTOR:
+ case TOKEN_LAYER:
+ case TOKEN_FLOW:
+ case TOKEN_NETWORK:
+ case TOKEN_NETWORK_FORWARD:
+ case TOKEN_REFLECT:
return FALSE;
default:
return TRUE;
}
+ case WINDIVERT_LAYER_REFLECT:
+ switch (kind)
+ {
+ case TOKEN_ICMP_BODY:
+ case TOKEN_ICMP_CHECKSUM:
+ case TOKEN_ICMP_CODE:
+ case TOKEN_ICMP_TYPE:
+ case TOKEN_ICMPV6_BODY:
+ case TOKEN_ICMPV6_CHECKSUM:
+ case TOKEN_ICMPV6_CODE:
+ case TOKEN_ICMPV6_TYPE:
+ case TOKEN_IP_CHECKSUM:
+ case TOKEN_IP_DF:
+ case TOKEN_IP_DST_ADDR:
+ case TOKEN_IP_FRAG_OFF:
+ case TOKEN_IP_HDR_LENGTH:
+ case TOKEN_IP_ID:
+ case TOKEN_IP_LENGTH:
+ case TOKEN_IP_MF:
+ case TOKEN_IP_PROTOCOL:
+ case TOKEN_IP_SRC_ADDR:
+ case TOKEN_IP_TOS:
+ case TOKEN_IP_TTL:
+ case TOKEN_IPV6_DST_ADDR:
+ case TOKEN_IPV6_FLOW_LABEL:
+ case TOKEN_IPV6_HOP_LIMIT:
+ case TOKEN_IPV6_LENGTH:
+ case TOKEN_IPV6_NEXT_HDR:
+ case TOKEN_IPV6_SRC_ADDR:
+ case TOKEN_IPV6_TRAFFIC_CLASS:
+ case TOKEN_TCP_ACK:
+ case TOKEN_TCP_ACK_NUM:
+ case TOKEN_TCP_CHECKSUM:
+ case TOKEN_TCP_DST_PORT:
+ case TOKEN_TCP_FIN:
+ case TOKEN_TCP_HDR_LENGTH:
+ case TOKEN_TCP_PAYLOAD_LENGTH:
+ case TOKEN_TCP_PSH:
+ case TOKEN_TCP_RST:
+ case TOKEN_TCP_SEQ_NUM:
+ case TOKEN_TCP_SRC_PORT:
+ case TOKEN_TCP_SYN:
+ case TOKEN_TCP_URG:
+ case TOKEN_TCP_URG_PTR:
+ case TOKEN_TCP_WINDOW:
+ case TOKEN_UDP_CHECKSUM:
+ case TOKEN_UDP_DST_PORT:
+ case TOKEN_UDP_LENGTH:
+ case TOKEN_UDP_PAYLOAD_LENGTH:
+ case TOKEN_UDP_SRC_PORT:
+ case TOKEN_IP:
+ case TOKEN_IPV6:
+ case TOKEN_ICMP:
+ case TOKEN_ICMPV6:
+ case TOKEN_TCP:
+ case TOKEN_UDP:
+ case TOKEN_LOOPBACK:
+ case TOKEN_IF_IDX:
+ case TOKEN_SUB_IF_IDX:
+ case TOKEN_IMPOSTOR:
+ case TOKEN_INBOUND:
+ case TOKEN_OUTBOUND:
+ case TOKEN_LOCAL_ADDR:
+ case TOKEN_REMOTE_ADDR:
+ case TOKEN_LOCAL_PORT:
+ case TOKEN_REMOTE_PORT:
+ case TOKEN_PROTOCOL:
+ return FALSE;
+ default:
+ return TRUE;
+ }
+ default:
+ return FALSE;
+ }
+}
+
+/*
+ * Expand a "macro" value.
+ */
+static BOOL WinDivertExpandMacro(KIND kind, UINT32 *val)
+{
+ switch (kind)
+ {
+ case TOKEN_NETWORK:
+ *val = WINDIVERT_LAYER_NETWORK;
+ return TRUE;
+ case TOKEN_NETWORK_FORWARD:
+ *val = WINDIVERT_LAYER_NETWORK_FORWARD;
+ return TRUE;
+ case TOKEN_FLOW:
+ *val = WINDIVERT_LAYER_FLOW;
+ return TRUE;
+ case TOKEN_REFLECT:
+ *val = WINDIVERT_LAYER_REFLECT;
+ return TRUE;
default:
return FALSE;
}
@@ -884,6 +994,10 @@ static ERROR WinDivertTokenizeFilter(const char *filter, WINDIVERT_LAYER layer,
{
static const TOKEN_NAME token_names[] =
{
+ {"FLOW", TOKEN_FLOW},
+ {"NETWORK", TOKEN_NETWORK},
+ {"NETWORK_FORWARD", TOKEN_NETWORK_FORWARD},
+ {"REFLECT", TOKEN_REFLECT},
{"and", TOKEN_AND},
{"false", TOKEN_FALSE},
{"icmp", TOKEN_ICMP},
@@ -920,6 +1034,7 @@ static ERROR WinDivertTokenizeFilter(const char *filter, WINDIVERT_LAYER layer,
{"ipv6.NextHdr", TOKEN_IPV6_NEXT_HDR},
{"ipv6.SrcAddr", TOKEN_IPV6_SRC_ADDR},
{"ipv6.TrafficClass", TOKEN_IPV6_TRAFFIC_CLASS},
+ {"layer", TOKEN_LAYER},
{"localAddr", TOKEN_LOCAL_ADDR},
{"localPort", TOKEN_LOCAL_PORT},
{"loopback", TOKEN_LOOPBACK},
@@ -954,6 +1069,7 @@ static ERROR WinDivertTokenizeFilter(const char *filter, WINDIVERT_LAYER layer,
{"udp.Length", TOKEN_UDP_LENGTH},
{"udp.PayloadLength", TOKEN_UDP_PAYLOAD_LENGTH},
{"udp.SrcPort", TOKEN_UDP_SRC_PORT},
+ {"zero", TOKEN_ZERO},
};
TOKEN_NAME *result;
char c;
@@ -1053,12 +1169,13 @@ static ERROR WinDivertTokenizeFilter(const char *filter, WINDIVERT_LAYER layer,
break;
}
token[0] = c;
- if (WinDivertIsAlNum(c) || c == '.' || c == ':')
+ if (WinDivertIsAlNum(c) || c == '.' || c == ':' || c == '_')
{
UINT32 num;
char *end;
for (j = 1; j < TOKEN_MAXLEN && (WinDivertIsAlNum(filter[i]) ||
- filter[i] == '.' || filter[i] == ':'); j++, i++)
+ filter[i] == '.' || filter[i] == ':' || filter[i] == '_');
+ j++, i++)
{
token[j] = filter[i];
}
@@ -1087,7 +1204,15 @@ static ERROR WinDivertTokenizeFilter(const char *filter, WINDIVERT_LAYER layer,
{
return MAKE_ERROR(WINDIVERT_ERROR_BAD_TOKEN_FOR_LAYER, i-j);
}
- tokens[tp++].kind = result->kind;
+ if (WinDivertExpandMacro(result->kind, &tokens[tp].val[0]))
+ {
+ tokens[tp].kind = TOKEN_NUMBER;
+ }
+ else
+ {
+ tokens[tp].kind = result->kind;
+ }
+ tp++;
continue;
}
@@ -1135,26 +1260,10 @@ static ERROR WinDivertTokenizeFilter(const char *filter, WINDIVERT_LAYER layer,
}
}
-/*
- * Pool allocation.
- */
-static void *WinDivertAlloc(PPOOL pool, UINT size)
-{
- void *ptr;
- if (pool->offset + size >= sizeof(pool->memory))
- {
- pool->error = MAKE_ERROR(WINDIVERT_ERROR_NO_MEMORY, 0);
- return NULL;
- }
- ptr = pool->memory + pool->offset;
- pool->offset += size;
- return ptr;
-};
-
/*
* Construct a variable/field.
*/
-static PEXPR WinDivertMakeVar(PPOOL pool, KIND kind)
+static PEXPR WinDivertMakeVar(KIND kind, PERROR error)
{
// NOTE: must be in order of kind.
static const EXPR vars[] =
@@ -1212,6 +1321,7 @@ static PEXPR WinDivertMakeVar(PPOOL pool, KIND kind)
{{{0}}, TOKEN_UDP_LENGTH},
{{{0}}, TOKEN_UDP_PAYLOAD_LENGTH},
{{{0}}, TOKEN_UDP_SRC_PORT},
+ {{{0}}, TOKEN_ZERO},
{{{0}}, TOKEN_TRUE},
{{{0}}, TOKEN_FALSE},
{{{0}}, TOKEN_INBOUND},
@@ -1226,6 +1336,7 @@ static PEXPR WinDivertMakeVar(PPOOL pool, KIND kind)
{{{0}}, TOKEN_LOCAL_PORT},
{{{0}}, TOKEN_REMOTE_PORT},
{{{0}}, TOKEN_PROTOCOL},
+ {{{0}}, TOKEN_LAYER},
};
// Binary search:
@@ -1245,14 +1356,14 @@ static PEXPR WinDivertMakeVar(PPOOL pool, KIND kind)
}
return (PEXPR)(vars + mid);
}
- pool->error = MAKE_ERROR(WINDIVERT_ERROR_ASSERTION_FAILED, 0);
+ *error = MAKE_ERROR(WINDIVERT_ERROR_ASSERTION_FAILED, 0);
return NULL;
}
/*
* Construct zero.
*/
-static PEXPR WinDivertMakeZero(PPOOL pool)
+static PEXPR WinDivertMakeZero(void)
{
static const EXPR zero = {{{0, 0, 0, 0}}, TOKEN_NUMBER};
return (PEXPR)&zero;
@@ -1261,44 +1372,39 @@ static PEXPR WinDivertMakeZero(PPOOL pool)
/*
* Construct a number.
*/
-static PEXPR WinDivertMakeNumber(PPOOL pool, TOKEN *tok)
+static PEXPR WinDivertMakeNumber(HANDLE pool, UINT32 *val, PERROR error)
{
- PEXPR expr;
- if (tok->kind != TOKEN_NUMBER)
- {
- pool->error = MAKE_ERROR(WINDIVERT_ERROR_ASSERTION_FAILED, 0);
- return NULL;
- }
- expr = (PEXPR)WinDivertAlloc(pool, sizeof(EXPR));
+ PEXPR expr = (PEXPR)HeapAlloc(pool, HEAP_ZERO_MEMORY, sizeof(EXPR));
if (expr == NULL)
{
+ *error = MAKE_ERROR(WINDIVERT_ERROR_NO_MEMORY, 0);
return NULL;
}
- memset(expr, 0, sizeof(EXPR));
expr->kind = TOKEN_NUMBER;
- expr->val[0] = tok->val[0];
- expr->val[1] = tok->val[1];
- expr->val[2] = tok->val[2];
- expr->val[3] = tok->val[3];
+ expr->val[0] = val[0];
+ expr->val[1] = val[1];
+ expr->val[2] = val[2];
+ expr->val[3] = val[3];
return expr;
}
/*
* Construct a binary operator.
*/
-static PEXPR WinDivertMakeBinOp(PPOOL pool, KIND kind, PEXPR arg0, PEXPR arg1)
+static PEXPR WinDivertMakeBinOp(HANDLE pool, KIND kind, PEXPR arg0, PEXPR arg1,
+ PERROR error)
{
PEXPR expr;
if (arg0 == NULL || arg1 == NULL)
{
return NULL;
}
- expr = (PEXPR)WinDivertAlloc(pool, sizeof(EXPR));
+ expr = (PEXPR)HeapAlloc(pool, HEAP_ZERO_MEMORY, sizeof(EXPR));
if (expr == NULL)
{
+ *error = MAKE_ERROR(WINDIVERT_ERROR_NO_MEMORY, 0);
return NULL;
}
- memset(expr, 0, sizeof(EXPR));
expr->kind = kind;
expr->arg[0] = arg0;
expr->arg[1] = arg1;
@@ -1308,15 +1414,15 @@ static PEXPR WinDivertMakeBinOp(PPOOL pool, KIND kind, PEXPR arg0, PEXPR arg1)
/*
* Construct an if-then-else.
*/
-static PEXPR WinDivertMakeIfThenElse(PPOOL pool, PEXPR cond, PEXPR th,
- PEXPR el)
+static PEXPR WinDivertMakeIfThenElse(HANDLE pool, PEXPR cond, PEXPR th,
+ PEXPR el, PERROR error)
{
- PEXPR expr = (PEXPR)WinDivertAlloc(pool, sizeof(EXPR));
+ PEXPR expr = (PEXPR)HeapAlloc(pool, HEAP_ZERO_MEMORY, sizeof(EXPR));
if (expr == NULL)
{
+ *error = MAKE_ERROR(WINDIVERT_ERROR_NO_MEMORY, 0);
return NULL;
}
- memset(expr, 0, sizeof(EXPR));
expr->kind = TOKEN_QUESTION;
expr->arg[0] = cond;
expr->arg[1] = th;
@@ -1327,7 +1433,7 @@ static PEXPR WinDivertMakeIfThenElse(PPOOL pool, PEXPR cond, PEXPR th,
/*
* Parse a filter test.
*/
-static PEXPR WinDivertParseTest(PPOOL pool, TOKEN *toks, UINT *i)
+static PEXPR WinDivertParseTest(HANDLE pool, TOKEN *toks, UINT *i, PERROR error)
{
PEXPR var, val;
KIND kind;
@@ -1339,6 +1445,7 @@ static PEXPR WinDivertParseTest(PPOOL pool, TOKEN *toks, UINT *i)
}
switch (toks[*i].kind)
{
+ case TOKEN_ZERO:
case TOKEN_TRUE:
case TOKEN_FALSE:
case TOKEN_OUTBOUND:
@@ -1359,6 +1466,7 @@ static PEXPR WinDivertParseTest(PPOOL pool, TOKEN *toks, UINT *i)
case TOKEN_LOCAL_PORT:
case TOKEN_REMOTE_PORT:
case TOKEN_PROTOCOL:
+ case TOKEN_LAYER:
case TOKEN_IP_HDR_LENGTH:
case TOKEN_IP_TOS:
case TOKEN_IP_LENGTH:
@@ -1408,11 +1516,10 @@ static PEXPR WinDivertParseTest(PPOOL pool, TOKEN *toks, UINT *i)
case TOKEN_UDP_PAYLOAD_LENGTH:
break;
default:
- pool->error = MAKE_ERROR(WINDIVERT_ERROR_UNEXPECTED_TOKEN,
- toks[*i].pos);
+ *error = MAKE_ERROR(WINDIVERT_ERROR_UNEXPECTED_TOKEN, toks[*i].pos);
return NULL;
}
- var = WinDivertMakeVar(pool, toks[*i].kind);
+ var = WinDivertMakeVar(toks[*i].kind, error);
*i = *i + 1;
switch (toks[*i].kind)
{
@@ -1426,7 +1533,7 @@ static PEXPR WinDivertParseTest(PPOOL pool, TOKEN *toks, UINT *i)
break;
default:
return WinDivertMakeBinOp(pool, (not? TOKEN_EQ: TOKEN_NEQ), var,
- WinDivertMakeZero(pool));
+ WinDivertMakeZero(), error);
}
if (not)
{
@@ -1457,31 +1564,31 @@ static PEXPR WinDivertParseTest(PPOOL pool, TOKEN *toks, UINT *i)
*i = *i + 1;
if (toks[*i].kind != TOKEN_NUMBER)
{
- pool->error = MAKE_ERROR(WINDIVERT_ERROR_UNEXPECTED_TOKEN,
- toks[*i].pos);
+ *error = MAKE_ERROR(WINDIVERT_ERROR_UNEXPECTED_TOKEN, toks[*i].pos);
return NULL;
}
- val = WinDivertMakeNumber(pool, toks + *i);
+ val = WinDivertMakeNumber(pool, toks[*i].val, error);
*i = *i + 1;
- return WinDivertMakeBinOp(pool, kind, var, val);
+ return WinDivertMakeBinOp(pool, kind, var, val, error);
}
/*
* Parse a filter argument to an (and) (or) operator.
*/
-static PEXPR WinDivertParseArg(PPOOL pool, TOKEN *toks, UINT *i, INT depth)
+static PEXPR WinDivertParseArg(HANDLE pool, TOKEN *toks, UINT *i, INT depth,
+ PERROR error)
{
PEXPR arg, th, el;
if (depth-- < 0)
{
- pool->error = MAKE_ERROR(WINDIVERT_ERROR_TOO_DEEP, toks[*i].pos);
+ *error = MAKE_ERROR(WINDIVERT_ERROR_TOO_DEEP, toks[*i].pos);
return NULL;
}
switch (toks[*i].kind)
{
case TOKEN_OPEN:
*i = *i + 1;
- arg = WinDivertParseFilter(pool, toks, i, depth, FALSE);
+ arg = WinDivertParseFilter(pool, toks, i, depth, FALSE, error);
if (toks[*i].kind == TOKEN_CLOSE)
{
*i = *i + 1;
@@ -1490,57 +1597,56 @@ static PEXPR WinDivertParseArg(PPOOL pool, TOKEN *toks, UINT *i, INT depth)
if (toks[*i].kind == TOKEN_QUESTION)
{
*i = *i + 1;
- th = WinDivertParseFilter(pool, toks, i, depth, FALSE);
+ th = WinDivertParseFilter(pool, toks, i, depth, FALSE, error);
if (th == NULL)
{
return NULL;
}
if (toks[*i].kind != TOKEN_COLON)
{
- pool->error = MAKE_ERROR(WINDIVERT_ERROR_UNEXPECTED_TOKEN,
+ *error = MAKE_ERROR(WINDIVERT_ERROR_UNEXPECTED_TOKEN,
toks[*i].pos);
return NULL;
}
*i = *i + 1;
- el = WinDivertParseFilter(pool, toks, i, depth, FALSE);
+ el = WinDivertParseFilter(pool, toks, i, depth, FALSE, error);
if (el == NULL)
{
return NULL;
}
if (toks[*i].kind != TOKEN_CLOSE)
{
- pool->error = MAKE_ERROR(WINDIVERT_ERROR_UNEXPECTED_TOKEN,
+ *error = MAKE_ERROR(WINDIVERT_ERROR_UNEXPECTED_TOKEN,
toks[*i].pos);
return NULL;
}
*i = *i + 1;
- arg = WinDivertMakeIfThenElse(pool, arg, th, el);
+ arg = WinDivertMakeIfThenElse(pool, arg, th, el, error);
return arg;
}
- pool->error = MAKE_ERROR(WINDIVERT_ERROR_UNEXPECTED_TOKEN,
- toks[*i].pos);
+ *error = MAKE_ERROR(WINDIVERT_ERROR_UNEXPECTED_TOKEN, toks[*i].pos);
return NULL;
default:
- return WinDivertParseTest(pool, toks, i);
+ return WinDivertParseTest(pool, toks, i, error);
}
}
/*
* Parse the filter into an expression object.
*/
-static PEXPR WinDivertParseFilter(PPOOL pool, TOKEN *toks, UINT *i, INT depth,
- BOOL and)
+static PEXPR WinDivertParseFilter(HANDLE pool, TOKEN *toks, UINT *i, INT depth,
+ BOOL and, PERROR error)
{
PEXPR expr, arg;
if (depth-- < 0)
{
- pool->error = MAKE_ERROR(WINDIVERT_ERROR_TOO_DEEP, toks[*i].pos);
+ *error = MAKE_ERROR(WINDIVERT_ERROR_TOO_DEEP, toks[*i].pos);
return NULL;
}
if (and)
- expr = WinDivertParseArg(pool, toks, i, depth);
+ expr = WinDivertParseArg(pool, toks, i, depth, error);
else
- expr = WinDivertParseFilter(pool, toks, i, depth, TRUE);
+ expr = WinDivertParseFilter(pool, toks, i, depth, TRUE, error);
do
{
if (expr == NULL)
@@ -1551,13 +1657,13 @@ static PEXPR WinDivertParseFilter(PPOOL pool, TOKEN *toks, UINT *i, INT depth,
{
case TOKEN_AND:
*i = *i + 1;
- arg = WinDivertParseArg(pool, toks, i, depth);
- expr = WinDivertMakeBinOp(pool, TOKEN_AND, expr, arg);
+ arg = WinDivertParseArg(pool, toks, i, depth, error);
+ expr = WinDivertMakeBinOp(pool, TOKEN_AND, expr, arg, error);
continue;
case TOKEN_OR:
*i = *i + 1;
- arg = WinDivertParseFilter(pool, toks, i, depth, TRUE);
- expr = WinDivertMakeBinOp(pool, TOKEN_OR, expr, arg);
+ arg = WinDivertParseFilter(pool, toks, i, depth, TRUE, error);
+ expr = WinDivertMakeBinOp(pool, TOKEN_OR, expr, arg, error);
continue;
default:
return expr;
@@ -1578,12 +1684,18 @@ static BOOL WinDivertEvalTest(PEXPR test, BOOL *res)
UINT32 lb, ub;
switch (var->kind)
{
+ case TOKEN_ZERO:
+ lb = ub = 0;
+ break;
case TOKEN_TRUE:
lb = ub = 1;
break;
case TOKEN_FALSE:
lb = ub = 0;
break;
+ case TOKEN_LAYER:
+ lb = 0; ub = WINDIVERT_LAYER_MAX;
+ break;
case TOKEN_INBOUND:
case TOKEN_OUTBOUND:
case TOKEN_IP:
@@ -1615,7 +1727,7 @@ static BOOL WinDivertEvalTest(PEXPR test, BOOL *res)
case TOKEN_ICMP_CODE:
case TOKEN_ICMPV6_TYPE:
case TOKEN_ICMPV6_CODE:
- case TOKEN_PROCESS_ID:
+ case TOKEN_PROTOCOL:
lb = 0; ub = 0xFF;
break;
case TOKEN_IP_FRAG_OFF:
@@ -1787,7 +1899,7 @@ static INT16 WinDivertFlattenExpr(PEXPR expr, INT16 *label, INT16 succ,
* Emit a test.
*/
static void WinDivertEmitTest(PEXPR test, UINT16 offset,
- windivert_ioctl_filter_t object)
+ PWINDIVERT_FILTER object)
{
PEXPR var = test->arg[0], val = test->arg[1];
switch (test->kind)
@@ -1815,6 +1927,9 @@ static void WinDivertEmitTest(PEXPR test, UINT16 offset,
}
switch (var->kind)
{
+ case TOKEN_ZERO:
+ object->field = WINDIVERT_FILTER_FIELD_ZERO;
+ break;
case TOKEN_OUTBOUND:
object->field = WINDIVERT_FILTER_FIELD_OUTBOUND;
break;
@@ -1848,6 +1963,12 @@ static void WinDivertEmitTest(PEXPR test, UINT16 offset,
case TOKEN_REMOTE_PORT:
object->field = WINDIVERT_FILTER_FIELD_REMOTEPORT;
break;
+ case TOKEN_PROTOCOL:
+ object->field = WINDIVERT_FILTER_FIELD_PROTOCOL;
+ break;
+ case TOKEN_LAYER:
+ object->field = WINDIVERT_FILTER_FIELD_LAYER;
+ break;
case TOKEN_IP:
object->field = WINDIVERT_FILTER_FIELD_IP;
break;
@@ -2041,7 +2162,7 @@ static void WinDivertEmitTest(PEXPR test, UINT16 offset,
* Emit a filter object.
*/
static void WinDivertEmitFilter(PEXPR *stack, UINT len, UINT16 label,
- windivert_ioctl_filter_t object, UINT *obj_len)
+ PWINDIVERT_FILTER object, UINT *obj_len)
{
UINT i;
switch (label)
@@ -2049,12 +2170,11 @@ static void WinDivertEmitFilter(PEXPR *stack, UINT len, UINT16 label,
case WINDIVERT_FILTER_RESULT_ACCEPT:
case WINDIVERT_FILTER_RESULT_REJECT:
object[0].field = WINDIVERT_FILTER_FIELD_ZERO;
- object[0].test = (label == WINDIVERT_FILTER_RESULT_ACCEPT?
- WINDIVERT_FILTER_TEST_EQ: WINDIVERT_FILTER_TEST_NEQ);
+ object[0].test = WINDIVERT_FILTER_TEST_EQ;
object[0].arg[0] = object[0].arg[1] = object[0].arg[2] =
object[0].arg[3] = 0;
- object[0].success = WINDIVERT_FILTER_RESULT_ACCEPT;
- object[0].failure = WINDIVERT_FILTER_RESULT_REJECT;
+ object[0].success = label;
+ object[0].failure = label;
*obj_len = 1;
return;
default:
@@ -2067,50 +2187,225 @@ static void WinDivertEmitFilter(PEXPR *stack, UINT len, UINT16 label,
}
}
+/*
+ * Analyze a filter object.
+ */
+static UINT64 WinDivertAnalyzeFilter(PWINDIVERT_FILTER filter, UINT length)
+{
+ BOOL result;
+ UINT64 flags = 0;
+
+ // False filter?
+ result = WinDivertCondExecFilter(filter, length,
+ WINDIVERT_FILTER_FIELD_ZERO, 0);
+ if (!result)
+ {
+ return 0;
+ }
+
+ // Inbound?
+ result = WinDivertCondExecFilter(filter, length,
+ WINDIVERT_FILTER_FIELD_INBOUND, 1);
+ if (result)
+ {
+ result = WinDivertCondExecFilter(filter, length,
+ WINDIVERT_FILTER_FIELD_OUTBOUND, 0);
+ }
+ flags |= (result? WINDIVERT_FILTER_FLAG_INBOUND: 0);
+
+ // Outbound?
+ result = WinDivertCondExecFilter(filter, length,
+ WINDIVERT_FILTER_FIELD_OUTBOUND, 1);
+ if (result)
+ {
+ result = WinDivertCondExecFilter(filter, length,
+ WINDIVERT_FILTER_FIELD_INBOUND, 0);
+ }
+ flags |= (result? WINDIVERT_FILTER_FLAG_OUTBOUND: 0);
+
+ // IPv4?
+ result = WinDivertCondExecFilter(filter, length,
+ WINDIVERT_FILTER_FIELD_IP, 1);
+ if (result)
+ {
+ result = WinDivertCondExecFilter(filter, length,
+ WINDIVERT_FILTER_FIELD_IPV6, 0);
+ }
+ flags |= (result? WINDIVERT_FILTER_FLAG_IP: 0);
+
+ // Ipv6?
+ result = WinDivertCondExecFilter(filter, length,
+ WINDIVERT_FILTER_FIELD_IPV6, 1);
+ if (result)
+ {
+ result = WinDivertCondExecFilter(filter, length,
+ WINDIVERT_FILTER_FIELD_IP, 0);
+ }
+ flags |= (result? WINDIVERT_FILTER_FLAG_IPV6: 0);
+
+ return flags;
+}
+
+/*
+ * Execute a filter object with respect to an assumption/condition.
+ * FALSE = definite reject; TRUE = maybe accept.
+ */
+static BOOL WinDivertCondExecFilter(PWINDIVERT_FILTER filter, UINT length,
+ UINT8 field, UINT32 arg)
+{
+ INT16 ip;
+ UINT8 succ, fail;
+ BOOL result[WINDIVERT_FILTER_MAXLEN];
+ BOOL result_succ, result_fail, result_test;
+
+ if (length == 0)
+ {
+ return TRUE;
+ }
+
+ for (ip = (INT16)(length-1); ip >= 0; ip--)
+ {
+ succ = filter[ip].success;
+ if (succ == WINDIVERT_FILTER_RESULT_ACCEPT || succ <= ip ||
+ succ >= length)
+ {
+ result_succ = TRUE;
+ }
+ else if (succ == WINDIVERT_FILTER_RESULT_REJECT)
+ {
+ result_succ = FALSE;
+ }
+ else
+ {
+ result_succ = result[succ];
+ }
+
+ fail = filter[ip].failure;
+ if (fail == WINDIVERT_FILTER_RESULT_ACCEPT || fail <= ip ||
+ fail >= length)
+ {
+ result_fail = TRUE;
+ }
+ else if (fail == WINDIVERT_FILTER_RESULT_REJECT)
+ {
+ result_fail = FALSE;
+ }
+ else
+ {
+ result_fail = result[fail];
+ }
+
+ if (result_succ && result_fail)
+ {
+ result[ip] = TRUE;
+ }
+ else if (!result_succ && !result_fail)
+ {
+ result[ip] = FALSE;
+ }
+ else if (filter[ip].field == field)
+ {
+ switch (filter[ip].test)
+ {
+ case WINDIVERT_FILTER_TEST_EQ:
+ result_test = (arg == filter[ip].arg[0]);
+ break;
+ case WINDIVERT_FILTER_TEST_NEQ:
+ result_test = (arg != filter[ip].arg[0]);
+ break;
+ case WINDIVERT_FILTER_TEST_LT:
+ result_test = (arg < filter[ip].arg[0]);
+ break;
+ case WINDIVERT_FILTER_TEST_LEQ:
+ result_test = (arg <= filter[ip].arg[0]);
+ break;
+ case WINDIVERT_FILTER_TEST_GT:
+ result_test = (arg > filter[ip].arg[0]);
+ break;
+ case WINDIVERT_FILTER_TEST_GEQ:
+ result_test = (arg >= filter[ip].arg[0]);
+ break;
+ default:
+ return TRUE; // abort.
+ }
+ result[ip] = (result_test? result_succ: result_fail);
+ }
+ else
+ {
+ result[ip] = TRUE;
+ }
+ }
+
+ return result[0];
+}
+
/*
* Compile a filter string into an executable filter object.
*/
static ERROR WinDivertCompileFilter(const char *filter,
- WINDIVERT_LAYER layer, windivert_ioctl_filter_t object, UINT *obj_len)
+ WINDIVERT_LAYER layer, PWINDIVERT_FILTER object, UINT *obj_len)
{
- TOKEN tokens[WINDIVERT_FILTER_MAXLEN*3];
- PEXPR stack[WINDIVERT_FILTER_MAXLEN];
- PPOOL pool;
+ TOKEN *tokens;
+ PEXPR *stack;
+ HANDLE pool;
PEXPR expr;
UINT i, max_depth;
INT16 label;
+ const SIZE_T min_pool_size = 8192;
+ const SIZE_T tokens_size = 5 * WINDIVERT_FILTER_MAXLEN;
ERROR error;
- // Tokenize the filter string:
- error = WinDivertTokenizeFilter(filter, layer, tokens,
- sizeof(tokens) / sizeof(tokens[0]) - 1);
- if (IS_ERROR(error))
+ // Check for pre-compiled filter object:
+ if (filter[0] == '@')
{
- return error;
+ WINDIVERT_STREAM stream;
+ stream.data = (char *)filter;
+ stream.pos = 0;
+ stream.max = UINT_MAX;
+ stream.overflow = FALSE;
+
+ if (!WinDivertDeserializeFilter(&stream, object, obj_len))
+ {
+ return MAKE_ERROR(WINDIVERT_ERROR_BAD_OBJECT, 0);
+ }
+ return MAKE_ERROR(WINDIVERT_ERROR_NONE, 0);
}
- // Allocate memory pool for the compiler:
- pool = (PPOOL)HeapAlloc(GetProcessHeap(), 0, sizeof(POOL));
+ // Allocate memory for the compiler:
+ pool = HeapCreate(HEAP_NO_SERIALIZE, min_pool_size, 16 * min_pool_size);
if (pool == NULL)
{
return MAKE_ERROR(WINDIVERT_ERROR_NO_MEMORY, 0);
}
- pool->offset = 0;
- pool->error = MAKE_ERROR(WINDIVERT_ERROR_NONE, 0);
+ tokens = (TOKEN *)HeapAlloc(pool, 0, tokens_size * sizeof(TOKEN));
+ stack = (PEXPR *)HeapAlloc(pool, 0,
+ WINDIVERT_FILTER_MAXLEN * sizeof(PEXPR));
+ if (tokens == NULL || stack == NULL)
+ {
+ HeapDestroy(pool);
+ return MAKE_ERROR(WINDIVERT_ERROR_NO_MEMORY, 0);
+ }
+
+ // Tokenize the filter string:
+ error = WinDivertTokenizeFilter(filter, layer, tokens, tokens_size-1);
+ if (IS_ERROR(error))
+ {
+ HeapDestroy(pool);
+ return error;
+ }
// Parse the filter into an expression:
i = 0;
max_depth = 1024;
- expr = WinDivertParseFilter(pool, tokens, &i, max_depth, FALSE);
+ expr = WinDivertParseFilter(pool, tokens, &i, max_depth, FALSE, &error);
if (expr == NULL)
{
- error = pool->error;
- HeapFree(GetProcessHeap(), 0, pool);
+ HeapDestroy(pool);
return error;
}
if (tokens[i].kind != TOKEN_END)
{
- HeapFree(GetProcessHeap(), 0, pool);
+ HeapDestroy(pool);
return MAKE_ERROR(WINDIVERT_ERROR_UNEXPECTED_TOKEN, tokens[i].pos);
}
@@ -2120,7 +2415,7 @@ static ERROR WinDivertCompileFilter(const char *filter,
WINDIVERT_FILTER_RESULT_REJECT, stack);
if (label < 0)
{
- HeapFree(GetProcessHeap(), 0, pool);
+ HeapDestroy(pool);
return MAKE_ERROR(WINDIVERT_ERROR_TOO_LONG, 0);
}
@@ -2129,7 +2424,7 @@ static ERROR WinDivertCompileFilter(const char *filter,
{
WinDivertEmitFilter(stack, label, label, object, obj_len);
}
- HeapFree(GetProcessHeap(), 0, pool);
+ HeapDestroy(pool);
return MAKE_ERROR(WINDIVERT_ERROR_NONE, 0);
}
@@ -2157,6 +2452,8 @@ static const char *WinDivertErrorString(UINT code)
return "Filter expression parse error";
case WINDIVERT_ERROR_OUTPUT_TOO_SHORT:
return "Filter object buffer is too short";
+ case WINDIVERT_ERROR_BAD_OBJECT:
+ return "Filter object is invalid";
case WINDIVERT_ERROR_ASSERTION_FAILED:
return "Internal assertion failed";
default:
@@ -2165,10 +2462,11 @@ static const char *WinDivertErrorString(UINT code)
}
/*
- * Check the given filter string.
+ * Compile the given filter string.
*/
-extern BOOL WinDivertHelperCheckFilter(const char *filter_str,
- WINDIVERT_LAYER layer, const char **error, UINT *error_pos)
+extern BOOL WinDivertHelperCompileFilter(const char *filter_str,
+ WINDIVERT_LAYER layer, char *object, UINT obj_len, const char **error,
+ UINT *error_pos)
{
ERROR err;
if (filter_str == NULL)
@@ -2176,7 +2474,33 @@ extern BOOL WinDivertHelperCheckFilter(const char *filter_str,
SetLastError(ERROR_INVALID_PARAMETER);
return FALSE;
}
- err = WinDivertCompileFilter(filter_str, layer, NULL, NULL);
+
+ SetLastError(ERROR_SUCCESS);
+ if (object == NULL)
+ {
+ err = WinDivertCompileFilter(filter_str, layer, NULL, NULL);
+ }
+ else
+ {
+ WINDIVERT_FILTER object0[WINDIVERT_FILTER_MAXLEN];
+ UINT obj0_len;
+ err = WinDivertCompileFilter(filter_str, layer, object0, &obj0_len);
+ if (!IS_ERROR(err))
+ {
+ WINDIVERT_STREAM stream;
+ stream.data = object;
+ stream.pos = 0;
+ stream.max = obj_len;
+ stream.overflow = FALSE;
+
+ WinDivertSerializeFilter(&stream, object0, obj0_len);
+ if (stream.overflow)
+ {
+ SetLastError(ERROR_INSUFFICIENT_BUFFER);
+ err = MAKE_ERROR(WINDIVERT_ERROR_OUTPUT_TOO_SHORT, 0);
+ }
+ }
+ }
if (error != NULL)
{
*error = WinDivertErrorString(GET_CODE(err));
@@ -2246,7 +2570,7 @@ extern BOOL WinDivertHelperEvalFilter(const char *filter, PVOID packet,
UINT32 val[4];
BOOL pass;
int cmp;
- struct windivert_ioctl_filter_s object[WINDIVERT_FILTER_MAXLEN];
+ WINDIVERT_FILTER object[WINDIVERT_FILTER_MAXLEN];
UINT obj_len;
if (filter == NULL || addr == NULL)
@@ -2279,6 +2603,8 @@ extern BOOL WinDivertHelperEvalFilter(const char *filter, PVOID packet,
return FALSE;
}
break;
+ case WINDIVERT_LAYER_REFLECT:
+ break;
default:
SetLastError(ERROR_INVALID_PARAMETER);
return FALSE;
@@ -2648,3 +2974,1134 @@ extern BOOL WinDivertHelperEvalFilter(const char *filter, PVOID packet,
}
}
+/*
+ * Get a char from a stream.
+ */
+static char WinDivertGetChar(PWINDIVERT_STREAM stream)
+{
+ char c;
+ if (stream->pos >= stream->max)
+ {
+ stream->overflow = TRUE;
+ return EOF;
+ }
+ c = stream->data[stream->pos];
+ stream->pos++;
+ return c;
+}
+
+/*
+ * Deserialize a number.
+ */
+static BOOL WinDivertDeserializeNumber(PWINDIVERT_STREAM stream, UINT max_len,
+ UINT32 *result)
+{
+ UINT32 i, val = 0;
+ char c;
+ for (i = 0; i < max_len; i++)
+ {
+ if ((val & 0xF8000000) != 0)
+ {
+ return FALSE; // Overflow
+ }
+ val <<= 5;
+ c = WinDivertGetChar(stream);
+ if (c >= '!' && c <= '!' + 31)
+ {
+ val += (UINT32)(c - '!');
+ }
+ else if (c >= '!' + 32 && c <= '!' + 64)
+ {
+ val += (UINT32)(c - '!' - 32);
+ *result = val;
+ return TRUE;
+ }
+ else
+ {
+ return FALSE;
+ }
+ }
+ return FALSE;
+}
+
+/*
+ * Deserialize a test.
+ */
+static BOOL WinDivertDeserializeTest(PWINDIVERT_STREAM stream,
+ PWINDIVERT_FILTER filter)
+{
+ UINT32 val;
+ UINT i;
+
+ if (WinDivertGetChar(stream) != '_')
+ {
+ return FALSE;
+ }
+
+ if (!WinDivertDeserializeNumber(stream, 2, &val) ||
+ val > WINDIVERT_FILTER_FIELD_MAX)
+ {
+ return FALSE;
+ }
+ filter->field = (UINT8)val;
+
+ if (!WinDivertDeserializeNumber(stream, 2, &val) ||
+ val > WINDIVERT_FILTER_TEST_MAX)
+ {
+ return FALSE;
+ }
+ filter->test = (UINT8)val;
+
+ if (!WinDivertDeserializeNumber(stream, 7, &filter->arg[0]))
+ {
+ return FALSE;
+ }
+
+ switch (filter->field)
+ {
+ case WINDIVERT_FILTER_FIELD_IPV6_SRCADDR:
+ case WINDIVERT_FILTER_FIELD_IPV6_DSTADDR:
+ case WINDIVERT_FILTER_FIELD_LOCALADDR:
+ case WINDIVERT_FILTER_FIELD_REMOTEADDR:
+ for (i = 1; i < 4; i++)
+ {
+ if (!WinDivertDeserializeNumber(stream, 7, &filter->arg[i]))
+ {
+ return FALSE;
+ }
+ }
+ break;
+ case WINDIVERT_FILTER_FIELD_IP_SRCADDR:
+ case WINDIVERT_FILTER_FIELD_IP_DSTADDR:
+ filter->arg[1] = 0x0000FFFF;
+ filter->arg[2] = filter->arg[3] = 0;
+ break;
+ default:
+ filter->arg[1] = filter->arg[2] = filter->arg[3] = 0;
+ break;
+ }
+
+ if (!WinDivertDeserializeNumber(stream, 2, &val) || val > UINT8_MAX)
+ {
+ return FALSE;
+ }
+ filter->success = (UINT8)val - 2;
+
+ if (!WinDivertDeserializeNumber(stream, 2, &val) || val > UINT8_MAX)
+ {
+ return FALSE;
+ }
+ filter->failure = (UINT8)val - 2;
+
+ return TRUE;
+}
+
+/*
+ * Deserialize a filter header.
+ */
+static BOOL WinDivertDeserializeFilterHeader(PWINDIVERT_STREAM stream,
+ UINT *length)
+{
+ UINT32 version, length32;
+
+ if (WinDivertGetChar(stream) != '@' ||
+ WinDivertGetChar(stream) != 'W' ||
+ WinDivertGetChar(stream) != 'i' ||
+ WinDivertGetChar(stream) != 'n' ||
+ WinDivertGetChar(stream) != 'D' ||
+ WinDivertGetChar(stream) != 'i' ||
+ WinDivertGetChar(stream) != 'v' ||
+ WinDivertGetChar(stream) != '_')
+ {
+ return FALSE;
+ }
+
+ if (!WinDivertDeserializeNumber(stream, 4, &version) || (version != 0))
+ {
+ return FALSE;
+ }
+
+ if (!WinDivertDeserializeNumber(stream, 2, &length32) ||
+ length32 == 0 || length32 > WINDIVERT_FILTER_MAXLEN)
+ {
+ return FALSE;
+ }
+ *length = length32;
+
+ return TRUE;
+}
+
+/*
+ * Deserialize a filter.
+ */
+static BOOL WinDivertDeserializeFilter(PWINDIVERT_STREAM stream,
+ PWINDIVERT_FILTER filter, UINT *length)
+{
+ UINT i;
+
+ if (!WinDivertDeserializeFilterHeader(stream, length))
+ {
+ return FALSE;
+ }
+
+ for (i = 0; i < *length; i++)
+ {
+ if (!WinDivertDeserializeTest(stream, filter + i))
+ {
+ return FALSE;
+ }
+ }
+
+ if (WinDivertGetChar(stream) != '\0')
+ {
+ return FALSE;
+ }
+
+ return TRUE;
+}
+
+/*
+ * Decompile a test into an expression.
+ */
+static PEXPR WinDivertDecompileTest(HANDLE pool, PWINDIVERT_FILTER test)
+{
+ KIND kind;
+ PEXPR var, val, expr;
+ ERROR error;
+
+ switch (test->field)
+ {
+ case WINDIVERT_FILTER_FIELD_ZERO:
+ kind = TOKEN_ZERO; break;
+ case WINDIVERT_FILTER_FIELD_INBOUND:
+ kind = TOKEN_INBOUND; break;
+ case WINDIVERT_FILTER_FIELD_OUTBOUND:
+ kind = TOKEN_OUTBOUND; break;
+ case WINDIVERT_FILTER_FIELD_IFIDX:
+ kind = TOKEN_IF_IDX; break;
+ case WINDIVERT_FILTER_FIELD_SUBIFIDX:
+ kind = TOKEN_SUB_IF_IDX; break;
+ case WINDIVERT_FILTER_FIELD_IP:
+ kind = TOKEN_IP; break;
+ case WINDIVERT_FILTER_FIELD_IPV6:
+ kind = TOKEN_IPV6; break;
+ case WINDIVERT_FILTER_FIELD_ICMP:
+ kind = TOKEN_ICMP; break;
+ case WINDIVERT_FILTER_FIELD_TCP:
+ kind = TOKEN_TCP; break;
+ case WINDIVERT_FILTER_FIELD_UDP:
+ kind = TOKEN_UDP; break;
+ case WINDIVERT_FILTER_FIELD_ICMPV6:
+ kind = TOKEN_ICMPV6; break;
+ case WINDIVERT_FILTER_FIELD_IP_HDRLENGTH:
+ kind = TOKEN_IP_HDR_LENGTH; break;
+ case WINDIVERT_FILTER_FIELD_IP_TOS:
+ kind = TOKEN_IP_TOS; break;
+ case WINDIVERT_FILTER_FIELD_IP_LENGTH:
+ kind = TOKEN_IP_LENGTH; break;
+ case WINDIVERT_FILTER_FIELD_IP_ID:
+ kind = TOKEN_IP_ID; break;
+ case WINDIVERT_FILTER_FIELD_IP_DF:
+ kind = TOKEN_IP_DF; break;
+ case WINDIVERT_FILTER_FIELD_IP_MF:
+ kind = TOKEN_IP_MF; break;
+ case WINDIVERT_FILTER_FIELD_IP_FRAGOFF:
+ kind = TOKEN_IP_FRAG_OFF; break;
+ case WINDIVERT_FILTER_FIELD_IP_TTL:
+ kind = TOKEN_IP_TTL; break;
+ case WINDIVERT_FILTER_FIELD_IP_PROTOCOL:
+ kind = TOKEN_IP_PROTOCOL; break;
+ case WINDIVERT_FILTER_FIELD_IP_CHECKSUM:
+ kind = TOKEN_IP_CHECKSUM; break;
+ case WINDIVERT_FILTER_FIELD_IP_SRCADDR:
+ kind = TOKEN_IP_SRC_ADDR; break;
+ case WINDIVERT_FILTER_FIELD_IP_DSTADDR:
+ kind = TOKEN_IP_DST_ADDR; break;
+ case WINDIVERT_FILTER_FIELD_IPV6_TRAFFICCLASS:
+ kind = TOKEN_IPV6_TRAFFIC_CLASS; break;
+ case WINDIVERT_FILTER_FIELD_IPV6_FLOWLABEL:
+ kind = TOKEN_IPV6_FLOW_LABEL; break;
+ case WINDIVERT_FILTER_FIELD_IPV6_LENGTH:
+ kind = TOKEN_IPV6_LENGTH; break;
+ case WINDIVERT_FILTER_FIELD_IPV6_NEXTHDR:
+ kind = TOKEN_IPV6_NEXT_HDR; break;
+ case WINDIVERT_FILTER_FIELD_IPV6_HOPLIMIT:
+ kind = TOKEN_IPV6_HOP_LIMIT; break;
+ case WINDIVERT_FILTER_FIELD_IPV6_SRCADDR:
+ kind = TOKEN_IPV6_SRC_ADDR; break;
+ case WINDIVERT_FILTER_FIELD_IPV6_DSTADDR:
+ kind = TOKEN_IPV6_DST_ADDR; break;
+ case WINDIVERT_FILTER_FIELD_ICMP_TYPE:
+ kind = TOKEN_ICMP_TYPE; break;
+ case WINDIVERT_FILTER_FIELD_ICMP_CODE:
+ kind = TOKEN_ICMP_CODE; break;
+ case WINDIVERT_FILTER_FIELD_ICMP_CHECKSUM:
+ kind = TOKEN_ICMP_CHECKSUM; break;
+ case WINDIVERT_FILTER_FIELD_ICMP_BODY:
+ kind = TOKEN_ICMP_BODY; break;
+ case WINDIVERT_FILTER_FIELD_ICMPV6_TYPE:
+ kind = TOKEN_ICMPV6_TYPE; break;
+ case WINDIVERT_FILTER_FIELD_ICMPV6_CODE:
+ kind = TOKEN_ICMPV6_CODE; break;
+ case WINDIVERT_FILTER_FIELD_ICMPV6_CHECKSUM:
+ kind = TOKEN_ICMPV6_CHECKSUM; break;
+ case WINDIVERT_FILTER_FIELD_ICMPV6_BODY:
+ kind = TOKEN_ICMPV6_BODY; break;
+ case WINDIVERT_FILTER_FIELD_TCP_SRCPORT:
+ kind = TOKEN_TCP_SRC_PORT; break;
+ case WINDIVERT_FILTER_FIELD_TCP_DSTPORT:
+ kind = TOKEN_TCP_DST_PORT; break;
+ case WINDIVERT_FILTER_FIELD_TCP_SEQNUM:
+ kind = TOKEN_TCP_SEQ_NUM; break;
+ case WINDIVERT_FILTER_FIELD_TCP_ACKNUM:
+ kind = TOKEN_TCP_ACK_NUM; break;
+ case WINDIVERT_FILTER_FIELD_TCP_HDRLENGTH:
+ kind = TOKEN_TCP_HDR_LENGTH; break;
+ case WINDIVERT_FILTER_FIELD_TCP_URG:
+ kind = TOKEN_TCP_URG; break;
+ case WINDIVERT_FILTER_FIELD_TCP_ACK:
+ kind = TOKEN_TCP_ACK; break;
+ case WINDIVERT_FILTER_FIELD_TCP_PSH:
+ kind = TOKEN_TCP_PSH; break;
+ case WINDIVERT_FILTER_FIELD_TCP_RST:
+ kind = TOKEN_TCP_RST; break;
+ case WINDIVERT_FILTER_FIELD_TCP_SYN:
+ kind = TOKEN_TCP_SYN; break;
+ case WINDIVERT_FILTER_FIELD_TCP_FIN:
+ kind = TOKEN_TCP_FIN; break;
+ case WINDIVERT_FILTER_FIELD_TCP_WINDOW:
+ kind = TOKEN_TCP_WINDOW; break;
+ case WINDIVERT_FILTER_FIELD_TCP_CHECKSUM:
+ kind = TOKEN_TCP_CHECKSUM; break;
+ case WINDIVERT_FILTER_FIELD_TCP_URGPTR:
+ kind = TOKEN_TCP_URG_PTR; break;
+ case WINDIVERT_FILTER_FIELD_TCP_PAYLOADLENGTH:
+ kind = TOKEN_TCP_PAYLOAD_LENGTH; break;
+ case WINDIVERT_FILTER_FIELD_UDP_SRCPORT:
+ kind = TOKEN_UDP_SRC_PORT; break;
+ case WINDIVERT_FILTER_FIELD_UDP_DSTPORT:
+ kind = TOKEN_UDP_DST_PORT; break;
+ case WINDIVERT_FILTER_FIELD_UDP_LENGTH:
+ kind = TOKEN_UDP_LENGTH; break;
+ case WINDIVERT_FILTER_FIELD_UDP_CHECKSUM:
+ kind = TOKEN_UDP_CHECKSUM; break;
+ case WINDIVERT_FILTER_FIELD_UDP_PAYLOADLENGTH:
+ kind = TOKEN_UDP_PAYLOAD_LENGTH; break;
+ case WINDIVERT_FILTER_FIELD_LOOPBACK:
+ kind = TOKEN_LOOPBACK; break;
+ case WINDIVERT_FILTER_FIELD_IMPOSTOR:
+ kind = TOKEN_IMPOSTOR; break;
+ case WINDIVERT_FILTER_FIELD_PROCESSID:
+ kind = TOKEN_PROCESS_ID; break;
+ case WINDIVERT_FILTER_FIELD_LOCALADDR:
+ kind = TOKEN_LOCAL_ADDR; break;
+ case WINDIVERT_FILTER_FIELD_REMOTEADDR:
+ kind = TOKEN_REMOTE_ADDR; break;
+ case WINDIVERT_FILTER_FIELD_LOCALPORT:
+ kind = TOKEN_LOCAL_PORT; break;
+ case WINDIVERT_FILTER_FIELD_REMOTEPORT:
+ kind = TOKEN_REMOTE_PORT; break;
+ case WINDIVERT_FILTER_FIELD_PROTOCOL:
+ kind = TOKEN_PROTOCOL; break;
+ case WINDIVERT_FILTER_FIELD_LAYER:
+ kind = TOKEN_LAYER; break;
+ default:
+ return NULL;
+ }
+
+ var = WinDivertMakeVar(kind, &error);
+ if (var == NULL)
+ {
+ return NULL;
+ }
+ val = WinDivertMakeNumber(pool, test->arg, &error);
+ if (val == NULL)
+ {
+ return NULL;
+ }
+
+ switch (test->test)
+ {
+ case WINDIVERT_FILTER_TEST_EQ:
+ kind = TOKEN_EQ; break;
+ case WINDIVERT_FILTER_TEST_NEQ:
+ kind = TOKEN_NEQ; break;
+ case WINDIVERT_FILTER_TEST_LT:
+ kind = TOKEN_LT; break;
+ case WINDIVERT_FILTER_TEST_LEQ:
+ kind = TOKEN_LEQ; break;
+ case WINDIVERT_FILTER_TEST_GT:
+ kind = TOKEN_GT; break;
+ case WINDIVERT_FILTER_TEST_GEQ:
+ kind = TOKEN_GEQ; break;
+ default:
+ return NULL;
+ }
+
+ expr = WinDivertMakeBinOp(pool, kind, var, val, &error);
+ if (expr == NULL)
+ {
+ return NULL;
+ }
+ expr->succ = test->success;
+ expr->fail = test->failure;
+ return expr;
+}
+
+/*
+ * Dereference an expression.
+ */
+static void WinDivertDerefExpr(PEXPR *exprs, UINT8 i)
+{
+ switch (i)
+ {
+ case WINDIVERT_FILTER_RESULT_ACCEPT:
+ case WINDIVERT_FILTER_RESULT_REJECT:
+ return;
+ default:
+ exprs[i]->count--;
+ if (exprs[i]->count == 0)
+ {
+ exprs[i] = NULL;
+ }
+ return;
+ }
+}
+
+/*
+ * Apply an and/or simplification for WinDivertCoalesceAndOr().
+ */
+static PEXPR WinDivertSimplifyAndOr(HANDLE pool, PEXPR *exprs, PEXPR expr,
+ BOOL and, UINT8 next, UINT8 other)
+{
+ PEXPR next_expr = exprs[next], new_expr;
+ ERROR error;
+
+ new_expr = WinDivertMakeBinOp(pool, (and? TOKEN_AND: TOKEN_OR), expr,
+ next_expr, &error);
+ if (new_expr == NULL)
+ {
+ return NULL;
+ }
+ new_expr->succ = next_expr->succ;
+ new_expr->fail = next_expr->fail;
+ new_expr->count = expr->count;
+ WinDivertDerefExpr(exprs, next);
+ WinDivertDerefExpr(exprs, other);
+ return new_expr;
+}
+
+/*
+ * Detect and coalesce and/or (& (?:)) expression patterns.
+ */
+static PEXPR WinDivertCoalesceAndOr(HANDLE pool, PEXPR *exprs, UINT8 i,
+ ERROR *error)
+{
+ PEXPR expr, next_expr, new_expr;
+ BOOL singleton;
+ static const EXPR true_expr = {{{0}}, TOKEN_TRUE};
+
+ expr = exprs[i];
+ while (TRUE)
+ {
+ if (expr == NULL || expr->count == 0)
+ {
+ return NULL;
+ }
+
+ singleton = FALSE;
+ switch (expr->succ)
+ {
+ case WINDIVERT_FILTER_RESULT_ACCEPT:
+ case WINDIVERT_FILTER_RESULT_REJECT:
+ break;
+ default:
+ next_expr = exprs[expr->succ];
+ if (next_expr->count != 1)
+ {
+ break;
+ }
+ singleton = TRUE;
+ if (next_expr->fail == expr->fail)
+ {
+ expr = WinDivertSimplifyAndOr(pool, exprs, expr,
+ /*and=*/TRUE, expr->succ, expr->fail);
+ continue;
+ }
+ else if (next_expr->succ == expr->fail)
+ {
+ new_expr = (PEXPR)HeapAlloc(pool, HEAP_ZERO_MEMORY,
+ sizeof(EXPR));
+ if (new_expr == NULL)
+ {
+ return NULL;
+ }
+ new_expr->kind = TOKEN_QUESTION;
+ new_expr->arg[0] = expr;
+ new_expr->arg[1] = next_expr;
+ new_expr->arg[2] = (PEXPR)&true_expr;
+ new_expr->succ = next_expr->succ;
+ new_expr->fail = next_expr->fail;
+ new_expr->count = expr->count;
+ WinDivertDerefExpr(exprs, expr->succ);
+ WinDivertDerefExpr(exprs, expr->fail);
+ expr = new_expr;
+ continue;
+ }
+ break;
+ }
+ switch (expr->fail)
+ {
+ case WINDIVERT_FILTER_RESULT_ACCEPT:
+ case WINDIVERT_FILTER_RESULT_REJECT:
+ singleton = FALSE;
+ break;
+ default:
+ next_expr = exprs[expr->fail];
+ if (next_expr->count != 1)
+ {
+ singleton = FALSE;
+ break;
+ }
+ if (next_expr->succ == expr->succ)
+ {
+ expr = WinDivertSimplifyAndOr(pool, exprs, expr,
+ /*and=*/FALSE, expr->fail, expr->succ);
+ continue;
+ }
+ else if (next_expr->fail == expr->succ)
+ {
+ expr = WinDivertSimplifyAndOr(pool, exprs, expr,
+ /*and=*/TRUE, expr->fail, expr->succ);
+ continue;
+ }
+ break;
+ }
+
+ if (singleton)
+ {
+ // Both branches have count==1; simplify into a (?:) expression:
+ PEXPR succ_expr, fail_expr;
+ succ_expr = exprs[expr->succ];
+ fail_expr = exprs[expr->fail];
+ if (succ_expr->succ != fail_expr->succ ||
+ succ_expr->fail != fail_expr->fail)
+ {
+ break;
+ }
+ new_expr = (PEXPR)HeapAlloc(pool, HEAP_ZERO_MEMORY, sizeof(EXPR));
+ if (new_expr == NULL)
+ {
+ return NULL;
+ }
+ new_expr->kind = TOKEN_QUESTION;
+ new_expr->arg[0] = expr;
+ new_expr->arg[1] = succ_expr;
+ new_expr->arg[2] = fail_expr;
+ new_expr->succ = succ_expr->succ;
+ new_expr->fail = fail_expr->fail;
+ new_expr->count = expr->count;
+ WinDivertDerefExpr(exprs, expr->succ);
+ WinDivertDerefExpr(exprs, expr->fail);
+ WinDivertDerefExpr(exprs, new_expr->succ);
+ WinDivertDerefExpr(exprs, new_expr->fail);
+ expr = new_expr;
+ continue;
+ }
+
+ // No simplifications, so we are done.
+ break;
+ }
+
+ exprs[i] = expr;
+ return expr;
+}
+
+/*
+ * Coalesce all remaining expressions.
+ */
+static PEXPR WinDivertCoalesceExpr(HANDLE pool, PEXPR *exprs, UINT8 i)
+{
+ PEXPR expr, succ_expr, fail_expr, new_expr;
+ static const EXPR true_expr = {{{0}}, TOKEN_TRUE};
+ static const EXPR false_expr = {{{0}}, TOKEN_FALSE};
+
+ switch (i)
+ {
+ case WINDIVERT_FILTER_RESULT_ACCEPT:
+ return (PEXPR)&true_expr;
+ case WINDIVERT_FILTER_RESULT_REJECT:
+ return (PEXPR)&false_expr;
+ default:
+ break;
+ }
+
+ expr = exprs[i];
+ if (expr == NULL)
+ {
+ return NULL;
+ }
+
+ if (expr->succ == expr->fail)
+ {
+ return WinDivertCoalesceExpr(pool, exprs, expr->succ);
+ }
+
+ succ_expr = WinDivertCoalesceExpr(pool, exprs, expr->succ);
+ fail_expr = WinDivertCoalesceExpr(pool, exprs, expr->fail);
+ if (succ_expr == NULL || fail_expr == NULL)
+ {
+ return NULL;
+ }
+ if (succ_expr->kind == TOKEN_TRUE && fail_expr->kind == TOKEN_FALSE)
+ {
+ return expr;
+ }
+
+ new_expr = (PEXPR)HeapAlloc(pool, HEAP_ZERO_MEMORY, sizeof(EXPR));
+ if (new_expr == NULL)
+ {
+ return NULL;
+ }
+
+ new_expr->kind = TOKEN_QUESTION;
+ new_expr->arg[0] = expr;
+ new_expr->arg[1] = succ_expr;
+ new_expr->arg[2] = fail_expr;
+ return new_expr;
+}
+
+/*
+ * Format a decimal number.
+ */
+static void WinDivertFormatNumber(PWINDIVERT_STREAM stream, UINT32 val)
+{
+ UINT64 r = 1000000000, dig;
+ BOOL zeroes = FALSE;
+
+ while (r != 0)
+ {
+ dig = val / r;
+ val = val % r;
+ r = r / 10;
+ if (dig == 0 && !zeroes && r != 0)
+ {
+ continue;
+ }
+ WinDivertPutChar(stream, '0' + dig);
+ zeroes = TRUE;
+ }
+}
+
+/*
+ * Format a hexidecimal number.
+ */
+static void WinDivertFormatHexNumber(PWINDIVERT_STREAM stream, UINT32 val)
+{
+ INT s = 28;
+ UINT32 dig;
+ BOOL zeroes = FALSE;
+
+ while (s >= 0)
+ {
+ dig = (val & ((UINT32)0xF << s)) >> s;
+ s -= 4;
+ if (dig == 0 && !zeroes && s >= 0)
+ {
+ continue;
+ }
+ WinDivertPutChar(stream, (dig <= 9? '0' + dig: 'a' + (dig - 10)));
+ zeroes = TRUE;
+ }
+}
+
+/*
+ * Format an IPv4 address.
+ */
+static void WinDivertFormatIPv4Addr(PWINDIVERT_STREAM stream, UINT32 addr)
+{
+ WinDivertFormatNumber(stream, (addr & 0xFF000000) >> 24);
+ WinDivertPutChar(stream, '.');
+ WinDivertFormatNumber(stream, (addr & 0x00FF0000) >> 16);
+ WinDivertPutChar(stream, '.');
+ WinDivertFormatNumber(stream, (addr & 0x0000FF00) >> 8);
+ WinDivertPutChar(stream, '.');
+ WinDivertFormatNumber(stream, (addr & 0x000000FF) >> 0);
+}
+
+/*
+ * Format an IPv6 address.
+ */
+static void WinDivertFormatIPv6Addr(PWINDIVERT_STREAM stream,
+ const UINT32 *addr32)
+{
+ INT i, z_curr, z_count, z_start, z_max;
+ UINT16 addr[8];
+
+ // IPv4 special case:
+ if (addr32[3] == 0 && addr32[2] == 0 && addr32[1] == 0x0000FFFF)
+ {
+ WinDivertFormatIPv4Addr(stream, addr32[0]);
+ return;
+ }
+
+ // Find zeroes:
+ memcpy(addr, addr32, sizeof(addr));
+ z_curr = 7;
+ z_count = 0;
+ z_start = z_max = -1;
+ for (i = 7; i >= 0; i--)
+ {
+ if (addr[i] == 0)
+ {
+ z_count++;
+ z_start = (z_count > z_max? z_curr: z_start);
+ z_max = (z_count > z_max? z_count: z_max);
+ }
+ else
+ {
+ z_curr = i-1;
+ z_count = 0;
+ }
+ }
+
+ // Format address:
+ for (i = 7; i >= 0; i--)
+ {
+ if (i == z_start)
+ {
+ WinDivertPutString(stream, (i == 7? "::": ":"));
+ i -= (z_max-1);
+ continue;
+ }
+ WinDivertFormatHexNumber(stream, addr[i]);
+ WinDivertPutString(stream, (i != 0? ":": ""));
+ }
+}
+
+/*
+ * Format a test expression.
+ */
+static void WinDivertFormatTestExpr(PWINDIVERT_STREAM stream, PEXPR expr)
+{
+ PEXPR field = expr->arg[0], val = expr->arg[1];
+ BOOL ipv4_addr = FALSE, ipv6_addr = FALSE, layer = FALSE;
+
+ switch (field->kind)
+ {
+ case TOKEN_ZERO:
+ case TOKEN_INBOUND:
+ case TOKEN_OUTBOUND:
+ case TOKEN_IP:
+ case TOKEN_IPV6:
+ case TOKEN_ICMP:
+ case TOKEN_TCP:
+ case TOKEN_UDP:
+ case TOKEN_ICMPV6:
+ case TOKEN_IP_DF:
+ case TOKEN_IP_MF:
+ case TOKEN_TCP_URG:
+ case TOKEN_TCP_ACK:
+ case TOKEN_TCP_PSH:
+ case TOKEN_TCP_RST:
+ case TOKEN_TCP_SYN:
+ case TOKEN_TCP_FIN:
+ case TOKEN_LOOPBACK:
+ case TOKEN_IMPOSTOR:
+ if (val->val[1] != 0 || val->val[2] != 0 || val->val[3] != 0 ||
+ val->val[0] > 1)
+ {
+ break;
+ }
+ switch (expr->kind)
+ {
+ case TOKEN_EQ:
+ WinDivertPutString(stream, (val->val[0] == 0? "not ": ""));
+ WinDivertFormatExpr(stream, field, /*top_level=*/FALSE,
+ /*and=*/FALSE);
+ return;
+ case TOKEN_NEQ:
+ WinDivertPutString(stream, (val->val[0] != 0? "not ": ""));
+ WinDivertFormatExpr(stream, field, /*top_level=*/FALSE,
+ /*and=*/FALSE);
+ return;
+ default:
+ break;
+ }
+ break;
+ case TOKEN_IP_SRC_ADDR:
+ case TOKEN_IP_DST_ADDR:
+ ipv4_addr = TRUE;
+ break;
+ case TOKEN_IPV6_SRC_ADDR:
+ case TOKEN_IPV6_DST_ADDR:
+ case TOKEN_LOCAL_ADDR:
+ case TOKEN_REMOTE_ADDR:
+ ipv6_addr = TRUE;
+ break;
+ case TOKEN_LAYER:
+ layer = TRUE;
+ break;
+ default:
+ break;
+ }
+
+ WinDivertFormatExpr(stream, field, /*top_level=*/FALSE, /*and=*/FALSE);
+ switch (expr->kind)
+ {
+ case TOKEN_EQ:
+ WinDivertPutString(stream, " = "); break;
+ case TOKEN_NEQ:
+ WinDivertPutString(stream, " != "); break;
+ case TOKEN_LT:
+ WinDivertPutString(stream, " < "); break;
+ case TOKEN_LEQ:
+ WinDivertPutString(stream, " <= "); break;
+ case TOKEN_GT:
+ WinDivertPutString(stream, " > "); break;
+ case TOKEN_GEQ:
+ WinDivertPutString(stream, " >= "); break;
+ }
+ if (ipv4_addr)
+ {
+ WinDivertFormatIPv4Addr(stream, val->val[0]);
+ }
+ else if (ipv6_addr)
+ {
+ WinDivertFormatIPv6Addr(stream, val->val);
+ }
+ else if (layer)
+ {
+ switch (val->val[0])
+ {
+ case WINDIVERT_LAYER_NETWORK:
+ WinDivertPutString(stream, "NETWORK"); break;
+ case WINDIVERT_LAYER_NETWORK_FORWARD:
+ WinDivertPutString(stream, "NETWORK_FORWARD"); break;
+ case WINDIVERT_LAYER_FLOW:
+ WinDivertPutString(stream, "FLOW"); break;
+ case WINDIVERT_LAYER_REFLECT:
+ WinDivertPutString(stream, "REFLECT"); break;
+ default:
+ WinDivertFormatNumber(stream, val->val[0]); break;
+ }
+ }
+ else
+ {
+ WinDivertFormatNumber(stream, val->val[0]);
+ }
+}
+
+/*
+ * Format an expression.
+ */
+static void WinDivertFormatExpr(PWINDIVERT_STREAM stream, PEXPR expr,
+ BOOL top_level, BOOL and)
+{
+ if (stream->pos >= stream->max)
+ {
+ return;
+ }
+
+ switch (expr->kind)
+ {
+ case TOKEN_AND:
+ if (!top_level && !and)
+ {
+ WinDivertPutChar(stream, '(');
+ }
+ WinDivertFormatExpr(stream, expr->arg[0], /*top_level=*/FALSE,
+ /*and=*/TRUE);
+ WinDivertPutString(stream, " and ");
+ WinDivertFormatExpr(stream, expr->arg[1], /*top_level=*/FALSE,
+ /*and=*/TRUE);
+ if (!top_level && !and)
+ {
+ WinDivertPutChar(stream, ')');
+ }
+ return;
+ case TOKEN_OR:
+ if (!top_level && and)
+ {
+ WinDivertPutChar(stream, '(');
+ }
+ WinDivertFormatExpr(stream, expr->arg[0], /*top_level=*/FALSE,
+ /*and=*/FALSE);
+ WinDivertPutString(stream, " or ");
+ WinDivertFormatExpr(stream, expr->arg[1], /*top_level=*/FALSE,
+ /*and=*/FALSE);
+ if (!top_level && and)
+ {
+ WinDivertPutChar(stream, ')');
+ }
+ return;
+ case TOKEN_QUESTION:
+ WinDivertPutChar(stream, '(');
+ WinDivertFormatExpr(stream, expr->arg[0], /*top_level=*/TRUE,
+ /*and=*/FALSE);
+ WinDivertPutString(stream, "? ");
+ WinDivertFormatExpr(stream, expr->arg[1], /*top_level=*/TRUE,
+ /*and=*/FALSE);
+ WinDivertPutString(stream, ": ");
+ WinDivertFormatExpr(stream, expr->arg[2], /*top_level=*/TRUE,
+ /*and=*/FALSE);
+ WinDivertPutChar(stream, ')');
+ return;
+ case TOKEN_TRUE:
+ WinDivertPutString(stream, "true");
+ return;
+ case TOKEN_FALSE:
+ WinDivertPutString(stream, "false");
+ return;
+ case TOKEN_EQ:
+ case TOKEN_NEQ:
+ case TOKEN_LT:
+ case TOKEN_LEQ:
+ case TOKEN_GT:
+ case TOKEN_GEQ:
+ WinDivertFormatTestExpr(stream, expr);
+ return;
+ case TOKEN_ZERO:
+ WinDivertPutString(stream, "zero"); return;
+ case TOKEN_INBOUND:
+ WinDivertPutString(stream, "inbound"); return;
+ case TOKEN_OUTBOUND:
+ WinDivertPutString(stream, "outbound"); return;
+ case TOKEN_IF_IDX:
+ WinDivertPutString(stream, "ifIdx"); return;
+ case TOKEN_SUB_IF_IDX:
+ WinDivertPutString(stream, "subIfIdx"); return;
+ case TOKEN_IP:
+ WinDivertPutString(stream, "ip"); return;
+ case TOKEN_IPV6:
+ WinDivertPutString(stream, "ipv6"); return;
+ case TOKEN_ICMP:
+ WinDivertPutString(stream, "icmp"); return;
+ case TOKEN_TCP:
+ WinDivertPutString(stream, "tcp"); return;
+ case TOKEN_UDP:
+ WinDivertPutString(stream, "udp"); return;
+ case TOKEN_ICMPV6:
+ WinDivertPutString(stream, "icmpv6"); return;
+ case TOKEN_IP_HDR_LENGTH:
+ WinDivertPutString(stream, "ip.HdrLength"); return;
+ case TOKEN_IP_TOS:
+ WinDivertPutString(stream, "ip.TOS"); return;
+ case TOKEN_IP_LENGTH:
+ WinDivertPutString(stream, "ip.Length"); return;
+ case TOKEN_IP_ID:
+ WinDivertPutString(stream, "ip.Id"); return;
+ case TOKEN_IP_DF:
+ WinDivertPutString(stream, "ip.DF"); return;
+ case TOKEN_IP_MF:
+ WinDivertPutString(stream, "ip.MF"); return;
+ case TOKEN_IP_FRAG_OFF:
+ WinDivertPutString(stream, "ip.FragOff"); return;
+ case TOKEN_IP_TTL:
+ WinDivertPutString(stream, "ip.TTL"); return;
+ case TOKEN_IP_PROTOCOL:
+ WinDivertPutString(stream, "ip.Protocol"); return;
+ case TOKEN_IP_CHECKSUM:
+ WinDivertPutString(stream, "ip.Checksum"); return;
+ case TOKEN_IP_SRC_ADDR:
+ WinDivertPutString(stream, "ip.SrcAddr"); return;
+ case TOKEN_IP_DST_ADDR:
+ WinDivertPutString(stream, "ip.DstAddr"); return;
+ case TOKEN_IPV6_TRAFFIC_CLASS:
+ WinDivertPutString(stream, "ipv6.TrafficClass"); return;
+ case TOKEN_IPV6_FLOW_LABEL:
+ WinDivertPutString(stream, "ipv6.FlowLabel"); return;
+ case TOKEN_IPV6_LENGTH:
+ WinDivertPutString(stream, "ipv6.Length"); return;
+ case TOKEN_IPV6_NEXT_HDR:
+ WinDivertPutString(stream, "ipv6.NextHdr"); return;
+ case TOKEN_IPV6_HOP_LIMIT:
+ WinDivertPutString(stream, "ipv6.HopLimit"); return;
+ case TOKEN_IPV6_SRC_ADDR:
+ WinDivertPutString(stream, "ipv6.SrcAddr"); return;
+ case TOKEN_IPV6_DST_ADDR:
+ WinDivertPutString(stream, "ipv6.DstAddr"); return;
+ case TOKEN_ICMP_TYPE:
+ WinDivertPutString(stream, "icmp.Type"); return;
+ case TOKEN_ICMP_CODE:
+ WinDivertPutString(stream, "icmp.Code"); return;
+ case TOKEN_ICMP_CHECKSUM:
+ WinDivertPutString(stream, "icmp.Checksum"); return;
+ case TOKEN_ICMP_BODY:
+ WinDivertPutString(stream, "icmp.Body"); return;
+ case TOKEN_ICMPV6_TYPE:
+ WinDivertPutString(stream, "icmpv6.Type"); return;
+ case TOKEN_ICMPV6_CODE:
+ WinDivertPutString(stream, "icmpv6.Code"); return;
+ case TOKEN_ICMPV6_CHECKSUM:
+ WinDivertPutString(stream, "icmpv6.Checksum"); return;
+ case TOKEN_ICMPV6_BODY:
+ WinDivertPutString(stream, "icmpv6.Body"); return;
+ case TOKEN_TCP_SRC_PORT:
+ WinDivertPutString(stream, "tcp.SrcPort"); return;
+ case TOKEN_TCP_DST_PORT:
+ WinDivertPutString(stream, "tcp.DstPort"); return;
+ case TOKEN_TCP_SEQ_NUM:
+ WinDivertPutString(stream, "tcp.SeqNum"); return;
+ case TOKEN_TCP_ACK_NUM:
+ WinDivertPutString(stream, "tcp.AckNum"); return;
+ case TOKEN_TCP_HDR_LENGTH:
+ WinDivertPutString(stream, "tcp.HdrLength"); return;
+ case TOKEN_TCP_URG:
+ WinDivertPutString(stream, "tcp.Urg"); return;
+ case TOKEN_TCP_ACK:
+ WinDivertPutString(stream, "tcp.Ack"); return;
+ case TOKEN_TCP_PSH:
+ WinDivertPutString(stream, "tcp.Psh"); return;
+ case TOKEN_TCP_RST:
+ WinDivertPutString(stream, "tcp.Rst"); return;
+ case TOKEN_TCP_SYN:
+ WinDivertPutString(stream, "tcp.Syn"); return;
+ case TOKEN_TCP_FIN:
+ WinDivertPutString(stream, "tcp.Fin"); return;
+ case TOKEN_TCP_WINDOW:
+ WinDivertPutString(stream, "tcp.Window"); return;
+ case TOKEN_TCP_CHECKSUM:
+ WinDivertPutString(stream, "tcp.Checksum"); return;
+ case TOKEN_TCP_URG_PTR:
+ WinDivertPutString(stream, "tcp.UrgPtr"); return;
+ case TOKEN_TCP_PAYLOAD_LENGTH:
+ WinDivertPutString(stream, "tcp.PayloadLength"); return;
+ case TOKEN_UDP_SRC_PORT:
+ WinDivertPutString(stream, "udp.SrcPort"); return;
+ case TOKEN_UDP_DST_PORT:
+ WinDivertPutString(stream, "udp.DstPort"); return;
+ case TOKEN_UDP_LENGTH:
+ WinDivertPutString(stream, "udp.Length"); return;
+ case TOKEN_UDP_CHECKSUM:
+ WinDivertPutString(stream, "udp.Checksum"); return;
+ case TOKEN_UDP_PAYLOAD_LENGTH:
+ WinDivertPutString(stream, "udp.PayloadLength"); return;
+ case TOKEN_LOOPBACK:
+ WinDivertPutString(stream, "loopback"); return;
+ case TOKEN_IMPOSTOR:
+ WinDivertPutString(stream, "impostor"); return;
+ case TOKEN_PROCESS_ID:
+ WinDivertPutString(stream, "processId"); return;
+ case TOKEN_LOCAL_ADDR:
+ WinDivertPutString(stream, "localAddr"); return;
+ case TOKEN_REMOTE_ADDR:
+ WinDivertPutString(stream, "remoteAddr"); return;
+ case TOKEN_LOCAL_PORT:
+ WinDivertPutString(stream, "localPort"); return;
+ case TOKEN_REMOTE_PORT:
+ WinDivertPutString(stream, "remotePort"); return;
+ case TOKEN_PROTOCOL:
+ WinDivertPutString(stream, "protocol"); return;
+ case TOKEN_LAYER:
+ WinDivertPutString(stream, "layer"); return;
+ case TOKEN_NUMBER:
+ WinDivertFormatNumber(stream, expr->val[0]);
+ return;
+ }
+}
+
+/*
+ * Format a filter string.
+ */
+BOOL WinDivertHelperFormatFilter(const char *filter, WINDIVERT_LAYER layer,
+ char *buffer, UINT buflen)
+{
+ PEXPR exprs[WINDIVERT_FILTER_MAXLEN], expr;
+ ERROR err;
+ WINDIVERT_FILTER object[WINDIVERT_FILTER_MAXLEN];
+ UINT obj_len;
+ INT i;
+ HANDLE pool;
+ WINDIVERT_STREAM stream;
+ ERROR error;
+ const SIZE_T min_pool_size = 8192;
+
+ if (filter == NULL || buffer == NULL)
+ {
+ SetLastError(ERROR_INVALID_PARAMETER);
+ return FALSE;
+ }
+
+ err = WinDivertCompileFilter(filter, layer, object, &obj_len);
+ if (IS_ERROR(err))
+ {
+ SetLastError(ERROR_INVALID_PARAMETER);
+ return FALSE;
+ }
+
+ pool = HeapCreate(HEAP_NO_SERIALIZE, min_pool_size, 16 * min_pool_size);
+ if (pool == NULL)
+ {
+ return FALSE;
+ }
+
+ // Decompile all tests:
+ for (i = (INT)obj_len-1; i >= 0; i--)
+ {
+ expr = WinDivertDecompileTest(pool, object + i);
+ if (expr == NULL)
+ {
+ SetLastError(ERROR_INVALID_PARAMETER);
+ return FALSE;
+ }
+ exprs[i] = expr;
+ switch (expr->succ)
+ {
+ case WINDIVERT_FILTER_RESULT_ACCEPT:
+ case WINDIVERT_FILTER_RESULT_REJECT:
+ break;
+ default:
+ exprs[expr->succ]->count++;
+ break;
+ }
+ switch (expr->fail)
+ {
+ case WINDIVERT_FILTER_RESULT_ACCEPT:
+ case WINDIVERT_FILTER_RESULT_REJECT:
+ break;
+ default:
+ exprs[expr->fail]->count++;
+ break;
+ }
+ }
+ exprs[0]->count++;
+
+ // Coalesce (unflatten) tests into and/or expressions:
+ for (i = (INT)obj_len-1; i >= 0; i--)
+ {
+ error = MAKE_ERROR(WINDIVERT_ERROR_NONE, 0);
+ (PVOID)WinDivertCoalesceAndOr(pool, exprs, i, &error);
+ if (IS_ERROR(error))
+ {
+ HeapDestroy(pool);
+ return FALSE;
+ }
+ }
+
+ // Coalesce remaining expressions:
+ expr = WinDivertCoalesceExpr(pool, exprs, 0);
+ if (expr == NULL)
+ {
+ HeapDestroy(pool);
+ return FALSE;
+ }
+
+ // Format the final expression:
+ stream.data = buffer;
+ stream.pos = 0;
+ stream.max = buflen;
+ stream.overflow = FALSE;
+ WinDivertFormatExpr(&stream, expr, /*top_level=*/TRUE, /*and=*/FALSE);
+ WinDivertPutChar(&stream, '\0');
+
+ // Clean-up:
+ HeapDestroy(pool);
+ if (!stream.overflow)
+ {
+ return TRUE;
+ }
+ SetLastError(ERROR_INSUFFICIENT_BUFFER);
+ return FALSE;
+}
+
diff --git a/examples/flowtrack/flowtrack.c b/examples/flowtrack/flowtrack.c
index f705cae..83a1774 100644
--- a/examples/flowtrack/flowtrack.c
+++ b/examples/flowtrack/flowtrack.c
@@ -69,7 +69,7 @@ static void print_address(const UINT32 *addr)
if (addr[3] == 0 && addr[2] == 0 && addr[1] == 0x0000FFFF)
{
// IPv4 address:
- UINT32 a, b, c, d;
+ UINT32 a, b, c, d;
a = (addr[0] >> 24) & 0xFF;
b = (addr[0] >> 16) & 0xFF;
c = (addr[0] >> 8) & 0xFF;
@@ -82,9 +82,9 @@ static void print_address(const UINT32 *addr)
int i;
for (i = 3; i >= 0; i--)
{
- UINT32 a, b;
- a = (addr[i] >> 16) & 0xFFFF;
- b = (addr[i] >> 0) & 0xFFFF;
+ UINT32 a, b;
+ a = (addr[i] >> 16) & 0xFFFF;
+ b = (addr[i] >> 0) & 0xFFFF;
printf("%x:%x", a, b);
if (i != 0)
{
@@ -114,8 +114,8 @@ static DWORD draw(LPVOID arg)
while (TRUE)
{
- GetConsoleScreenBufferInfo(console, &screen);
- SetConsoleCursorPosition(console, top_left);
+ GetConsoleScreenBufferInfo(console, &screen);
+ SetConsoleCursorPosition(console, top_left);
rows = screen.srWindow.Bottom - screen.srWindow.Top + 1;
columns = screen.srWindow.Right - screen.srWindow.Left + 1;
@@ -132,7 +132,7 @@ static DWORD draw(LPVOID arg)
}
ReleaseMutex(lock);
- // Print the flows:
+ // Print the flows:
SetConsoleTextAttribute(console, BACKGROUND_RED | BACKGROUND_GREEN |
BACKGROUND_BLUE);
WriteConsole(console, header, sizeof(header)-1, &written, NULL);
@@ -142,21 +142,21 @@ static DWORD draw(LPVOID arg)
COORD pos = {sizeof(header)-1, 0};
FillConsoleOutputCharacterA(console, ' ', fill_len, pos,
&written);
- FillConsoleOutputAttribute(console,
+ FillConsoleOutputAttribute(console,
BACKGROUND_RED | BACKGROUND_GREEN | BACKGROUND_BLUE,
- fill_len, pos, &written);
+ fill_len, pos, &written);
}
putchar('\n');
SetConsoleTextAttribute(console,
FOREGROUND_RED | FOREGROUND_GREEN | FOREGROUND_BLUE);
- for (i = 0; i < num_addrs && i < rows-1; i++)
+ for (i = 0; i < num_addrs && i < rows-1; i++)
{
COORD pos = {0, i+1};
addr = &addrs[i];
FillConsoleOutputCharacterA(console, ' ', columns, pos, &written);
- FillConsoleOutputAttribute(console,
- FOREGROUND_GREEN | FOREGROUND_RED | FOREGROUND_BLUE,
- columns, pos, &written);
+ FillConsoleOutputAttribute(console,
+ FOREGROUND_GREEN | FOREGROUND_RED | FOREGROUND_BLUE,
+ columns, pos, &written);
SetConsoleCursorPosition(console, pos);
if (i == rows-2 && (i+1) < num_addrs)
{
@@ -191,7 +191,7 @@ static DWORD draw(LPVOID arg)
}
SetConsoleTextAttribute(console,
FOREGROUND_RED | FOREGROUND_GREEN | FOREGROUND_BLUE);
- switch (addr->Flow.Protocol)
+ switch (addr->Flow.Protocol)
{
case IPPROTO_TCP:
SetConsoleTextAttribute(console, FOREGROUND_GREEN);
@@ -227,9 +227,9 @@ static DWORD draw(LPVOID arg)
{
COORD pos = {0, i+1};
FillConsoleOutputCharacterA(console, ' ', columns, pos, &written);
- FillConsoleOutputAttribute(console,
- FOREGROUND_GREEN | FOREGROUND_RED | FOREGROUND_BLUE,
- columns, pos, &written);
+ FillConsoleOutputAttribute(console,
+ FOREGROUND_GREEN | FOREGROUND_RED | FOREGROUND_BLUE,
+ columns, pos, &written);
}
Sleep(1000);
@@ -260,7 +260,24 @@ int __cdecl main(int argc, char **argv)
exit(EXIT_FAILURE);
}
- // Spawn the draw() thread.
+ // Open WinDivert FLOW handle:
+ handle = WinDivertOpen(filter, WINDIVERT_LAYER_FLOW, priority,
+ WINDIVERT_FLAG_SNIFF | WINDIVERT_FLAG_RECV_ONLY);
+ if (handle == INVALID_HANDLE_VALUE)
+ {
+ if (GetLastError() == ERROR_INVALID_PARAMETER &&
+ !WinDivertHelperCompileFilter(filter, WINDIVERT_LAYER_FLOW,
+ NULL, 0, &err_str, NULL))
+ {
+ fprintf(stderr, "error: invalid filter \"%s\"\n", err_str);
+ exit(EXIT_FAILURE);
+ }
+ fprintf(stderr, "error: failed to open the WinDivert device (%d)\n",
+ GetLastError());
+ return EXIT_FAILURE;
+ }
+
+ // Spawn the draw() thread.
lock = CreateMutex(NULL, FALSE, NULL);
thread = CreateThread(NULL, 1, (LPTHREAD_START_ROUTINE)draw, NULL, 0,
NULL);
@@ -272,23 +289,6 @@ int __cdecl main(int argc, char **argv)
}
CloseHandle(thread);
- // Open WinDivert FLOW handle:
- handle = WinDivertOpen(filter, WINDIVERT_LAYER_FLOW, priority,
- WINDIVERT_FLAGS_LAYER_FLOW);
- if (handle == INVALID_HANDLE_VALUE)
- {
- if (GetLastError() == ERROR_INVALID_PARAMETER &&
- !WinDivertHelperCheckFilter(filter, WINDIVERT_LAYER_FLOW,
- &err_str, NULL))
- {
- fprintf(stderr, "error: invalid filter \"%s\"\n", err_str);
- exit(EXIT_FAILURE);
- }
- fprintf(stderr, "error: failed to open the WinDivert device (%d)\n",
- GetLastError());
- return EXIT_FAILURE;
- }
-
// Main loop:
while (TRUE)
{
@@ -302,7 +302,7 @@ int __cdecl main(int argc, char **argv)
{
case WINDIVERT_EVENT_FLOW_ESTABLISHED:
- // Flow established:
+ // Flow established:
flow = (PFLOW)malloc(sizeof(FLOW));
if (flow == NULL)
{
@@ -318,7 +318,7 @@ int __cdecl main(int argc, char **argv)
case WINDIVERT_EVENT_FLOW_DELETED:
- // Flow deleted:
+ // Flow deleted:
prev = NULL;
WaitForSingleObject(lock, INFINITE);
flow = flows;
diff --git a/examples/netdump/netdump.c b/examples/netdump/netdump.c
index ee1f46d..1785c93 100644
--- a/examples/netdump/netdump.c
+++ b/examples/netdump/netdump.c
@@ -100,8 +100,8 @@ int __cdecl main(int argc, char **argv)
if (handle == INVALID_HANDLE_VALUE)
{
if (GetLastError() == ERROR_INVALID_PARAMETER &&
- !WinDivertHelperCheckFilter(argv[1], WINDIVERT_LAYER_NETWORK,
- &err_str, NULL))
+ !WinDivertHelperCompileFilter(argv[1], WINDIVERT_LAYER_NETWORK,
+ NULL, 0, &err_str, NULL))
{
fprintf(stderr, "error: invalid filter \"%s\"\n", err_str);
exit(EXIT_FAILURE);
diff --git a/examples/netfilter/netfilter.c b/examples/netfilter/netfilter.c
index b8e3fc1..f898191 100644
--- a/examples/netfilter/netfilter.c
+++ b/examples/netfilter/netfilter.c
@@ -170,8 +170,8 @@ int __cdecl main(int argc, char **argv)
if (handle == INVALID_HANDLE_VALUE)
{
if (GetLastError() == ERROR_INVALID_PARAMETER &&
- !WinDivertHelperCheckFilter(argv[1], WINDIVERT_LAYER_NETWORK,
- &err_str, NULL))
+ !WinDivertHelperCompileFilter(argv[1], WINDIVERT_LAYER_NETWORK,
+ NULL, 0, &err_str, NULL))
{
fprintf(stderr, "error: invalid filter \"%s\"\n", err_str);
exit(EXIT_FAILURE);
diff --git a/examples/windivertctl/windivertctl.c b/examples/windivertctl/windivertctl.c
new file mode 100644
index 0000000..f01111d
--- /dev/null
+++ b/examples/windivertctl/windivertctl.c
@@ -0,0 +1,408 @@
+/*
+ * streamdump.c
+ * (C) 2018, all rights reserved,
+ *
+ * This file is part of WinDivert.
+ *
+ * WinDivert is free software: you can redistribute it and/or modify it under
+ * the terms of the GNU Lesser 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 Lesser General Public
+ * License for more details.
+ *
+ * You should have received a copy of the GNU Lesser General Public License
+ * along with this program. If not, see .
+ *
+ * WinDivert 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 2 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, write to the Free Software Foundation, Inc., 51
+ * Franklin Street, Fifth Floor, Boston, MA 02110-1301, USA.
+ */
+
+/*
+ * DESCRIPTION:
+ *
+ * usage: windivertctl.exe list
+ */
+
+#include
+#include
+#include
+#include
+#include
+#include
+
+#include "windivert.h"
+
+#define MAX_PACKET 0xFFFF
+#define MAX_FILTER_LEN 30000
+
+/*
+ * Process info.
+ */
+typedef struct INFO
+{
+ UINT32 process_id;
+ UINT32 ref_count;
+ HANDLE process;
+ struct INFO *next;
+} INFO, *PINFO;
+
+static INFO *open = NULL; // All open handles
+
+/*
+ * Modes.
+ */
+typedef enum
+{
+ LIST,
+ WATCH,
+ KILLALL
+} MODE;
+
+/*
+ * Months.
+ */
+static const char *months[12] =
+{
+ "Jan", "Feb", "Mar", "Apr", "May", "Jun", "Jul", "Aug", "Sep", "Oct",
+ "Nov", "Dec"
+};
+
+/*
+ * Add a new process.
+ */
+static HANDLE add_process(UINT32 process_id)
+{
+ PINFO info = open;
+ HANDLE process;
+
+ while (info != NULL)
+ {
+ if (info->process_id == process_id)
+ {
+ info->ref_count++;
+ return info->process;
+ }
+ info = info->next;
+ }
+
+ process = OpenProcess(PROCESS_QUERY_LIMITED_INFORMATION | PROCESS_TERMINATE,
+ FALSE, process_id);
+ info = (INFO *)malloc(sizeof(INFO));
+ if (info == NULL)
+ {
+ fprintf(stderr, "error: failed to allocate memory (%d)\n",
+ GetLastError());
+ exit(EXIT_FAILURE);
+ }
+ info->process_id = process_id;
+ info->process = process;
+ info->ref_count = 1;
+ info->next = open;
+ open = info;
+ return process;
+}
+
+/*
+ * Lookup a process.
+ */
+static HANDLE lookup_process(UINT32 process_id)
+{
+ PINFO info = open;
+
+ while (info != NULL)
+ {
+ if (info->process_id == process_id)
+ {
+ return info->process;
+ }
+ info = info->next;
+ }
+}
+
+/*
+ * Remove an old process.
+ */
+static void remove_process(UINT32 process_id)
+{
+ PINFO info = open, prev = NULL;
+
+ while (info != NULL)
+ {
+ if (info->process_id == process_id)
+ {
+ info->ref_count--;
+ if (info->ref_count > 0)
+ {
+ return;
+ }
+ break;
+ }
+ prev = info;
+ info = info->next;
+ }
+
+ if (info->process != NULL)
+ {
+ CloseHandle(info->process);
+ }
+ if (prev != NULL)
+ {
+ prev->next = info->next;
+ }
+ else
+ {
+ open = info->next;
+ }
+ free(info);
+}
+
+/*
+ * Entry.
+ */
+int __cdecl main(int argc, char **argv)
+{
+ HANDLE handle, process, console;
+ INT16 priority = -333; // Arbitrary.
+ UINT packet_len;
+ static UINT8 packet[MAX_PACKET];
+ static char path[MAX_PATH+1];
+ static char filter_str[MAX_FILTER_LEN];
+ PVOID object;
+ DWORD path_len;
+ BOOL or;
+ WINDIVERT_ADDRESS addr;
+ ULONGLONG freq, start_count;
+ LARGE_INTEGER li;
+ MODE mode;
+ const char *filter = "true";
+ const char *err_str = NULL;
+
+ if (argc != 2 && argc != 3)
+ {
+usage:
+ fprintf(stderr, "usage: %s (list|watch|killall) [filter]\n", argv[0]);
+ exit(EXIT_FAILURE);
+ }
+ if (strcmp(argv[1], "list") == 0)
+ {
+ mode = LIST;
+ }
+ else if (strcmp(argv[1], "watch") == 0)
+ {
+ mode = WATCH;
+ }
+ else if (strcmp(argv[1], "killall") == 0)
+ {
+ mode = KILLALL;
+ }
+ else
+ {
+ goto usage;
+ }
+ if (argc == 3)
+ {
+ filter = argv[2];
+ }
+
+ // Time management
+ QueryPerformanceFrequency(&li);
+ freq = li.QuadPart;
+ QueryPerformanceCounter(&li);
+ start_count = li.QuadPart;
+
+ // Open WinDivert REFLECT handle:
+ handle = WinDivertOpen(filter, WINDIVERT_LAYER_REFLECT, priority,
+ WINDIVERT_FLAG_SNIFF | WINDIVERT_FLAG_RECV_ONLY |
+ (mode == WATCH? 0: WINDIVERT_FLAG_NO_INSTALL));
+ if (handle == INVALID_HANDLE_VALUE)
+ {
+ if (mode != WATCH && GetLastError() == ERROR_SERVICE_DOES_NOT_EXIST)
+ {
+ // WinDivert driver is not running, so no open handles.
+ return 0;
+ }
+ if (GetLastError() == ERROR_INVALID_PARAMETER &&
+ !WinDivertHelperCompileFilter(filter, WINDIVERT_LAYER_FLOW,
+ NULL, 0, &err_str, NULL))
+ {
+ fprintf(stderr, "error: invalid filter \"%s\"\n", err_str);
+ exit(EXIT_FAILURE);
+ }
+ fprintf(stderr, "error: failed to open the WinDivert device (%d)\n",
+ GetLastError());
+ return EXIT_FAILURE;
+ }
+
+ // Main loop:
+ console = GetStdHandle(STD_OUTPUT_HANDLE);
+ while (TRUE)
+ {
+ if (!WinDivertRecv(handle, packet, sizeof(packet), &addr, &packet_len))
+ {
+ fprintf(stderr, "failed to event (%d)\n", GetLastError());
+ continue;
+ }
+
+ switch (addr.Event)
+ {
+ case WINDIVERT_EVENT_REFLECT_ESTABLISHED:
+ case WINDIVERT_EVENT_REFLECT_OPEN:
+ // Open handle:
+ process = add_process(addr.Reflect.ProcessId);
+ if (mode == KILLALL)
+ {
+ SetConsoleTextAttribute(console, FOREGROUND_RED);
+ fputs("KILL", stdout);
+ TerminateProcess(process, 0);
+ }
+ else
+ {
+ SetConsoleTextAttribute(console, FOREGROUND_GREEN);
+ fputs("OPEN", stdout);
+ }
+ break;
+
+ case WINDIVERT_EVENT_REFLECT_CLOSE:
+ // Close handle:
+ if (mode != WATCH)
+ {
+ continue;
+ }
+ process = lookup_process(addr.Reflect.ProcessId);
+ SetConsoleTextAttribute(console, FOREGROUND_RED);
+ fputs("CLOSE", stdout);
+ break;
+ }
+ SetConsoleTextAttribute(console,
+ FOREGROUND_RED | FOREGROUND_GREEN | FOREGROUND_BLUE);
+ fputs(" time=", stdout);
+ SetConsoleTextAttribute(console, FOREGROUND_RED | FOREGROUND_GREEN);
+ printf("%.3fs", (double)(addr.Reflect.Timestamp - (INT64)start_count) /
+ (double)freq);
+ SetConsoleTextAttribute(console,
+ FOREGROUND_RED | FOREGROUND_GREEN | FOREGROUND_BLUE);
+ fputs(" pid=", stdout);
+ SetConsoleTextAttribute(console, FOREGROUND_RED | FOREGROUND_GREEN);
+ printf("%u", addr.Reflect.ProcessId);
+ SetConsoleTextAttribute(console,
+ FOREGROUND_RED | FOREGROUND_GREEN | FOREGROUND_BLUE);
+ fputs(" exe=", stdout);
+ path_len = 0;
+ if (process != NULL)
+ {
+ path_len = GetProcessImageFileName(process, path, sizeof(path));
+ }
+ SetConsoleTextAttribute(console, FOREGROUND_RED | FOREGROUND_GREEN);
+ printf("%s", (path_len != 0? path: "???"));
+ SetConsoleTextAttribute(console,
+ FOREGROUND_RED | FOREGROUND_GREEN | FOREGROUND_BLUE);
+ fputs(" layer=", stdout);
+ SetConsoleTextAttribute(console, FOREGROUND_RED | FOREGROUND_GREEN);
+ switch (addr.Reflect.Layer)
+ {
+ case WINDIVERT_LAYER_NETWORK:
+ fputs("NETWORK", stdout);
+ break;
+ case WINDIVERT_LAYER_NETWORK_FORWARD:
+ fputs("NETWORK_FORWARD", stdout);
+ break;
+ case WINDIVERT_LAYER_FLOW:
+ fputs("FLOW", stdout);
+ break;
+ case WINDIVERT_LAYER_REFLECT:
+ fputs("REFLECT", stdout);
+ break;
+ default:
+ fputs("???", stdout);
+ break;
+ }
+ SetConsoleTextAttribute(console,
+ FOREGROUND_RED | FOREGROUND_GREEN | FOREGROUND_BLUE);
+ fputs(" flags=", stdout);
+ SetConsoleTextAttribute(console, FOREGROUND_RED | FOREGROUND_GREEN);
+ if (addr.Reflect.Flags == 0)
+ {
+ fputs("0", stdout);
+ }
+ else
+ {
+ or = FALSE;
+ if ((addr.Reflect.Flags & WINDIVERT_FLAG_SNIFF) != 0)
+ {
+ fputs("SNIFF", stdout);
+ or = TRUE;
+ }
+ if ((addr.Reflect.Flags & WINDIVERT_FLAG_DROP) != 0)
+ {
+ printf("%sDROP", (or? "|": ""));
+ or = TRUE;
+ }
+ if ((addr.Reflect.Flags & WINDIVERT_FLAG_RECV_ONLY) != 0)
+ {
+ printf("%sRECV_ONLY", (or? "|": ""));
+ or = TRUE;
+ }
+ if ((addr.Reflect.Flags & WINDIVERT_FLAG_SEND_ONLY) != 0)
+ {
+ printf("%sSEND_ONLY", (or? "|": ""));
+ or = TRUE;
+ }
+ if ((addr.Reflect.Flags & WINDIVERT_FLAG_DEBUG) != 0)
+ {
+ printf("%sDEBUG", (or? "|": ""));
+ or = TRUE;
+ }
+ if ((addr.Reflect.Flags & WINDIVERT_FLAG_NO_INSTALL) != 0)
+ {
+ printf("%sNO_INSTALL", (or? "|": ""));
+ or = TRUE;
+ }
+ }
+ SetConsoleTextAttribute(console,
+ FOREGROUND_RED | FOREGROUND_GREEN | FOREGROUND_BLUE);
+ fputs(" priority=", stdout);
+ SetConsoleTextAttribute(console, FOREGROUND_RED | FOREGROUND_GREEN);
+ printf("%d", addr.Reflect.Priority);
+ SetConsoleTextAttribute(console,
+ FOREGROUND_RED | FOREGROUND_GREEN | FOREGROUND_BLUE);
+ fputs(" filter=", stdout);
+ SetConsoleTextAttribute(console, FOREGROUND_RED | FOREGROUND_GREEN);
+ WinDivertHelperParsePacket(packet, packet_len, NULL, NULL, NULL, NULL,
+ NULL, NULL, &object, NULL);
+ if (WinDivertHelperFormatFilter((char *)object, addr.Reflect.Layer,
+ filter_str, sizeof(filter_str)))
+ {
+ printf("\"%s\" \"%s\"", filter_str, (char *)object); // XXX
+ }
+ SetConsoleTextAttribute(console,
+ FOREGROUND_RED | FOREGROUND_GREEN | FOREGROUND_BLUE);
+ putchar('\n');
+
+ if (addr.Event == WINDIVERT_EVENT_REFLECT_CLOSE)
+ {
+ remove_process(addr.Reflect.ProcessId);
+ }
+ if (mode != WATCH && addr.Final)
+ {
+ break;
+ }
+ }
+
+ return 0;
+}
+
diff --git a/include/windivert.h b/include/windivert.h
index 49029f2..68c2131 100644
--- a/include/windivert.h
+++ b/include/windivert.h
@@ -69,6 +69,17 @@ extern "C" {
/* WINDIVERT API */
/****************************************************************************/
+/*
+ * WinDivert layers.
+ */
+typedef enum
+{
+ WINDIVERT_LAYER_NETWORK = 0, /* Network layer. */
+ WINDIVERT_LAYER_NETWORK_FORWARD = 1,/* Network layer (forwarded packets) */
+ WINDIVERT_LAYER_FLOW = 2, /* Flow layer. */
+ WINDIVERT_LAYER_REFLECT = 3, /* Reflect layer. */
+} WINDIVERT_LAYER, *PWINDIVERT_LAYER;
+
/*
* WinDivert NETWORK and NETWORK_FORWARD layer data.
*/
@@ -91,6 +102,18 @@ typedef struct
UINT8 Protocol; /* Protocol. */
} WINDIVERT_FLOW_DATA, *PWINDIVERT_FLOW_DATA;
+/*
+ * WinDivert REFLECTION layer data.
+ */
+typedef struct
+{
+ INT64 Timestamp; /* Handle open time. */
+ UINT32 ProcessId; /* Handle process ID. */
+ WINDIVERT_LAYER Layer; /* Handle layer. */
+ UINT64 Flags; /* Handle flags. */
+ INT16 Priority; /* Handle priority. */
+} WINDIVERT_REFLECT_DATA, *PWINDIVERT_REFLECT_DATA;
+
/*
* WinDivert address.
*/
@@ -106,24 +129,16 @@ typedef struct
UINT32 PseudoIPChecksum:1; /* Packet has pseudo IPv4 checksum? */
UINT32 PseudoTCPChecksum:1; /* Packet has pseudo TCP checksum? */
UINT32 PseudoUDPChecksum:1; /* Packet has pseudo UDP checksum? */
- UINT32 Reserved:9;
+ UINT32 Final:1; /* Packet is final event? */
+ UINT32 Reserved:8;
union
{
WINDIVERT_NETWORK_DATA Network; /* Network layer data. */
WINDIVERT_FLOW_DATA Flow; /* Flow layer data. */
+ WINDIVERT_REFLECT_DATA Reflect; /* Reflect layer data. */
};
} WINDIVERT_ADDRESS, *PWINDIVERT_ADDRESS;
-/*
- * WinDivert layers.
- */
-typedef enum
-{
- WINDIVERT_LAYER_NETWORK = 1, /* Network layer. */
- WINDIVERT_LAYER_NETWORK_FORWARD = 2,/* Network layer (forwarded packets) */
- WINDIVERT_LAYER_FLOW = 3 /* Flow layer. */
-} WINDIVERT_LAYER, *PWINDIVERT_LAYER;
-
/*
* WinDivert events.
*/
@@ -133,24 +148,23 @@ typedef enum
WINDIVERT_EVENT_FLOW_ESTABLISHED = 1,
/* Flow established. */
WINDIVERT_EVENT_FLOW_DELETED = 2, /* Flow deleted. */
+ WINDIVERT_EVENT_REFLECT_ESTABLISHED = 3,
+ /* Previously open WinDivert handle. */
+ WINDIVERT_EVENT_REFLECT_OPEN = 4, /* Open new WinDivert handle. */
+ WINDIVERT_EVENT_REFLECT_CLOSE = 5, /* Close existing WinDivert handle. */
} WINDIVERT_EVENT, *PWINDIVERT_EVENT;
/*
* WinDivert flags.
*/
-#define WINDIVERT_FLAG_SNIFF 1
-#define WINDIVERT_FLAG_DROP 2
-#define WINDIVERT_FLAG_RECV_ONLY 4
+#define WINDIVERT_FLAG_SNIFF 0x01
+#define WINDIVERT_FLAG_DROP 0x02
+#define WINDIVERT_FLAG_RECV_ONLY 0x04
#define WINDIVERT_FLAG_READ_ONLY WINDIVERT_FLAG_RECV_ONLY
-#define WINDIVERT_FLAG_SEND_ONLY 8
+#define WINDIVERT_FLAG_SEND_ONLY 0x08
#define WINDIVERT_FLAG_WRITE_ONLY WINDIVERT_FLAG_SEND_ONLY
-#define WINDIVERT_FLAG_DEBUG 16
-
-#define WINDIVERT_FLAGS_LAYER_NETWORK 0
-#define WINDIVERT_FLAGS_LAYER_NETWORK_FORWARD \
- 0
-#define WINDIVERT_FLAGS_LAYER_FLOW \
- (WINDIVERT_FLAG_SNIFF | WINDIVERT_FLAG_RECV_ONLY)
+#define WINDIVERT_FLAG_DEBUG 0x10
+#define WINDIVERT_FLAG_NO_INSTALL 0x20
/*
* WinDivert parameters.
@@ -430,11 +444,13 @@ extern WINDIVERTEXPORT UINT WinDivertHelperCalcChecksums(
__in UINT64 flags);
/*
- * Check the given filter string.
+ * Compile the given filter string.
*/
-extern WINDIVERTEXPORT BOOL WinDivertHelperCheckFilter(
+extern WINDIVERTEXPORT BOOL WinDivertHelperCompileFilter(
__in const char *filter,
__in WINDIVERT_LAYER layer,
+ __out_opt char *object,
+ __in UINT objLen,
__out_opt const char **errorStr,
__out_opt UINT *errorPos);
@@ -447,6 +463,15 @@ extern WINDIVERTEXPORT BOOL WinDivertHelperEvalFilter(
__in UINT packetLen,
__in PWINDIVERT_ADDRESS pAddr);
+/*
+ * Format the given filter string.
+ */
+extern BOOL WinDivertHelperFormatFilter(
+ __in const char *filter,
+ __in WINDIVERT_LAYER layer,
+ __out char *buffer,
+ __in UINT bufLen);
+
#endif /* WINDIVERT_KERNEL */
#ifdef __cplusplus
diff --git a/include/windivert_device.h b/include/windivert_device.h
index 73cc45f..2741402 100644
--- a/include/windivert_device.h
+++ b/include/windivert_device.h
@@ -128,8 +128,9 @@
#define WINDIVERT_FILTER_FIELD_LOCALPORT 63
#define WINDIVERT_FILTER_FIELD_REMOTEPORT 64
#define WINDIVERT_FILTER_FIELD_PROTOCOL 65
+#define WINDIVERT_FILTER_FIELD_LAYER 66
#define WINDIVERT_FILTER_FIELD_MAX \
- WINDIVERT_FILTER_FIELD_PROTOCOL
+ WINDIVERT_FILTER_FIELD_LAYER
#define WINDIVERT_FILTER_TEST_EQ 0
#define WINDIVERT_FILTER_TEST_NEQ 1
@@ -139,7 +140,7 @@
#define WINDIVERT_FILTER_TEST_GEQ 5
#define WINDIVERT_FILTER_TEST_MAX WINDIVERT_FILTER_TEST_GEQ
-#define WINDIVERT_FILTER_MAXLEN 128
+#define WINDIVERT_FILTER_MAXLEN (0xFF-2)
#define WINDIVERT_FILTER_RESULT_ACCEPT (WINDIVERT_FILTER_MAXLEN+1)
#define WINDIVERT_FILTER_RESULT_REJECT (WINDIVERT_FILTER_MAXLEN+2)
@@ -148,13 +149,15 @@
* WinDivert layers.
*/
#define WINDIVERT_LAYER_DEFAULT WINDIVERT_LAYER_NETWORK
+#define WINDIVERT_LAYER_MAX WINDIVERT_LAYER_REFLECT
/*
* WinDivert flags.
*/
#define WINDIVERT_FLAGS_ALL \
(WINDIVERT_FLAG_SNIFF | WINDIVERT_FLAG_DROP | WINDIVERT_FLAG_RECV_ONLY |\
- WINDIVERT_FLAG_SEND_ONLY | WINDIVERT_FLAG_DEBUG)
+ WINDIVERT_FLAG_SEND_ONLY | WINDIVERT_FLAG_DEBUG | \
+ WINDIVERT_FLAG_NO_INSTALL)
#define WINDIVERT_FLAGS_EXCLUDE(flags, flag1, flag2) \
(((flags) & ((flag1) | (flag2))) != ((flag1) | (flag2)))
#define WINDIVERT_FLAGS_VALID(flags) \
@@ -164,14 +167,24 @@
WINDIVERT_FLAGS_EXCLUDE(flags, WINDIVERT_FLAG_RECV_ONLY, \
WINDIVERT_FLAG_SEND_ONLY))
+/*
+ * WinDivert filter flags.
+ */
+#define WINDIVERT_FILTER_FLAG_INBOUND 0x0000000000000001ull
+#define WINDIVERT_FILTER_FLAG_OUTBOUND 0x0000000000000002ull
+#define WINDIVERT_FILTER_FLAG_IP 0x0000000000000004ull
+#define WINDIVERT_FILTER_FLAG_IPV6 0x0000000000000008ull
+
+#define WINDIVERT_FILTER_FLAGS_ALL \
+ (WINDIVERT_FILTER_FLAG_INBOUND | WINDIVERT_FILTER_FLAG_OUTBOUND | \
+ WINDIVERT_FILTER_FLAG_IP | WINDIVERT_FILTER_FLAG_IPV6)
+
/*
* WinDivert priorities.
*/
-#define WINDIVERT_PRIORITY(priority16) \
- ((UINT32)((INT32)(priority16) + 0x7FFF + 1))
-#define WINDIVERT_PRIORITY_DEFAULT WINDIVERT_PRIORITY(0)
-#define WINDIVERT_PRIORITY_MAX WINDIVERT_PRIORITY(1000)
-#define WINDIVERT_PRIORITY_MIN WINDIVERT_PRIORITY(-1000)
+#define WINDIVERT_PRIORITY_DEFAULT 0
+#define WINDIVERT_PRIORITY_MAX 30000
+#define WINDIVERT_PRIORITY_MIN -WINDIVERT_PRIORITY_MAX
/*
* WinDivert parameters.
@@ -190,27 +203,25 @@
* WinDivert message definitions.
*/
#pragma pack(push, 1)
-struct windivert_ioctl_s
+typedef struct
{
UINT16 magic; // WINDIVERT_IOCTL_MAGIC
UINT8 version; // WINDIVERT_IOCTL_VERSION
UINT8 arg8; // 8-bit argument
UINT64 arg; // 64-bit argument
-};
-typedef struct windivert_ioctl_s *windivert_ioctl_t;
+} WINDIVERT_IOCTL, *PWINDIVERT_IOCTL;
/*
* WinDivert IOCTL structures.
*/
-struct windivert_ioctl_filter_s
+typedef struct
{
UINT8 field; // WINDIVERT_FILTER_FIELD_*
UINT8 test; // WINDIVERT_FILTER_TEST_*
- UINT16 success; // Success continuation.
- UINT16 failure; // Fail continuation.
+ UINT8 success; // Success continuation.
+ UINT8 failure; // Fail continuation.
UINT32 arg[4]; // Argument.
-};
-typedef struct windivert_ioctl_filter_s *windivert_ioctl_filter_t;
+} WINDIVERT_FILTER, *PWINDIVERT_FILTER;
#pragma pack(pop)
/*
diff --git a/mingw-build.sh b/mingw-build.sh
index 74ca900..29a550e 100644
--- a/mingw-build.sh
+++ b/mingw-build.sh
@@ -59,7 +59,7 @@ do
fi
echo "BUILD MINGW-$CPU"
CC="$ENV-gcc"
- COPTS="-shared -Wall -Wno-pointer-to-int-cast -O2 -Iinclude/
+ COPTS="-shared -Wall -Wno-pointer-to-int-cast -Os -Iinclude/
-Wl,--enable-stdcall-fixup -Wl,--entry=${MANGLE}WinDivertDllEntry"
CLIBS="-lgcc -lkernel32 -ladvapi32"
STRIP="$ENV-strip"
@@ -101,6 +101,10 @@ do
$CC -s -O2 -Iinclude/ examples/flowtrack/flowtrack.c \
-o "install/MINGW/$CPU/flowtrack.exe" -lWinDivert -lws2_32 -lpsapi \
-lshlwapi -L"install/MINGW/$CPU/"
+ echo "\tcopy install/MINGW/$CPU/windivertctl.exe..."
+ $CC -s -O2 -Iinclude/ examples/windivertctl/windivertctl.c \
+ -o "install/MINGW/$CPU/windivertctl.exe" -lWinDivert -lws2_32 \
+ -lpsapi -lshlwapi -L"install/MINGW/$CPU/"
echo "\tcopy install/MINGW/$CPU/WinDivert$BITS.sys..."
cp install/WDDK/$CPU/WinDivert$BITS.sys install/MINGW/$CPU
else
diff --git a/sys/sources b/sys/sources
index 21685f8..1461d7a 100644
--- a/sys/sources
+++ b/sys/sources
@@ -19,6 +19,6 @@ NTTARGETFILES=
KMDF_VERSION_MAJOR=1
C_DEFINES=$(C_DEFINES) -DBINARY_COMPATIBLE=0 -DNT -DUNICODE -D_UNICODE \
-DNDIS60 -DNDIS_SUPPORT_NDIS60
-INCLUDES=$(DDK_INC_PATH);..\include
+INCLUDES=$(DDK_INC_PATH);..\include;..\dll
SOURCES=windivert.rc windivert.c
diff --git a/sys/windivert.c b/sys/windivert.c
index 9d4bcc2..08212d9 100644
--- a/sys/windivert.c
+++ b/sys/windivert.c
@@ -32,6 +32,7 @@
* Franklin Street, Fifth Floor, Boston, MA 02110-1301, USA.
*/
+#include
#include
#include
#include
@@ -55,6 +56,7 @@ EVT_WDF_FILE_CLEANUP windivert_cleanup;
EVT_WDF_FILE_CLOSE windivert_close;
EVT_WDF_OBJECT_CONTEXT_DESTROY windivert_destroy;
EVT_WDF_WORKITEM windivert_worker;
+EVT_WDF_WORKITEM windivert_reflect_worker;
/*
* Debugging macros.
@@ -97,27 +99,15 @@ static void DEBUG_ERROR(PCCH format, NTSTATUS status, ...)
#define WINDIVERT_TAG 'viDW'
/*
- * WinDivert packet filter.
+ * WinDivert reflect context information.
*/
-struct filter_s
+struct reflect_context_s
{
- UINT8 protocol:4; // field's protocol
- UINT8 test:4; // Filter test
- UINT8 field; // Field of interest
- UINT16 success; // Success continuation
- UINT16 failure; // Fail continuation
- UINT32 arg[4]; // Comparison argument
+ LIST_ENTRY entry; // Open handle entry.
+ LONGLONG timestamp; // Open timestamp.
+ WINDIVERT_REFLECT_DATA data; // Reflect data.
+ BOOL inserted; // Entry inserted?
};
-typedef struct filter_s *filter_t;
-#define WINDIVERT_FILTER_PROTOCOL_NONE 0
-#define WINDIVERT_FILTER_PROTOCOL_IP 1
-#define WINDIVERT_FILTER_PROTOCOL_IPV6 2
-#define WINDIVERT_FILTER_PROTOCOL_ICMP 3
-#define WINDIVERT_FILTER_PROTOCOL_ICMPV6 4
-#define WINDIVERT_FILTER_PROTOCOL_TCP 5
-#define WINDIVERT_FILTER_PROTOCOL_UDP 6
-#define WINDIVERT_FILTER_PROTOCOL_NETWORK 7
-#define WINDIVERT_FILTER_PROTOCOL_FLOW 8
/*
* WinDivert context information.
@@ -157,22 +147,27 @@ struct context_s
UINT8 worker_curr; // Current read worker.
UINT8 layer; // Context's layer.
UINT64 flags; // Context's flags.
- UINT32 priority; // Context's priority.
+ UINT32 priority; // Context (internal) priority.
+ INT16 priority16; // Context (user) priority.
GUID callout_guid[WINDIVERT_CONTEXT_MAXLAYERS];
// Callout GUIDs.
GUID filter_guid[WINDIVERT_CONTEXT_MAXLAYERS];
// Filter GUIDs.
BOOL installed[WINDIVERT_CONTEXT_MAXLAYERS];// What is installed?
HANDLE engine_handle; // WFP engine handle.
- filter_t filter; // Packet filter.
+ PWINDIVERT_FILTER filter; // Packet filter.
+ UINT8 filter_len; // Length of filter.
+ struct reflect_context_s reflect; // Reflection info.
};
typedef struct context_s context_s;
typedef struct context_s *context_t;
WDF_DECLARE_CONTEXT_TYPE_WITH_NAME(context_s, windivert_context_get);
#define WINDIVERT_TIMEOUT(context, t0, t1) \
- (((t1) >= (t0)? (t1) - (t0): (t0) - (t1)) > \
- (context)->packet_queue_maxcounts)
+ ((context)->layer == WINDIVERT_LAYER_NETWORK || \
+ (context)->layer == WINDIVERT_LAYER_NETWORK_FORWARD? \
+ ((t1) >= (t0)? (t1) - (t0): (t0) - (t1)) > \
+ (context)->packet_queue_maxcounts: FALSE)
/*
* WinDivert Layer information.
@@ -242,6 +237,7 @@ struct packet_s
UINT32 pseudo_ip_checksum:1; // Packet has pseudo IPv4 check?
UINT32 pseudo_tcp_checksum:1; // Packet has pseudo TCP check?
UINT32 pseudo_udp_checksum:1; // Packet has pseudo UDP check?
+ UINT32 final:1; // Packet is final event?
UINT32 match:1; // Packet matches filter?
UINT32 priority; // Packet priority.
UINT32 packet_len; // Length of the packet.
@@ -279,6 +275,18 @@ struct flow_s
};
typedef struct flow_s *flow_t;
+/*
+ * WinDivert reflect event.
+ */
+struct reflect_event_s
+{
+ LIST_ENTRY entry; // Entry for reflect_event_queue.
+ context_t context; // Context.
+ LONGLONG timestamp; // Event timestamp.
+ WINDIVERT_EVENT event; // Event.
+};
+typedef struct reflect_event_s *reflect_event_t;
+
/*
* IPv4/IPv6 pseudo headers.
*/
@@ -320,19 +328,20 @@ static LONGLONG counts_per_ms = 0;
static POOL_TYPE non_paged_pool = NonPagedPool;
/*
- * Priorities.
+ * Priorities & weights.
*/
-#define WINDIVERT_CONTEXT_PRIORITY(priority0) \
- windivert_context_priority(priority0)
-static UINT32 windivert_context_priority(UINT32 priority0)
+static UINT32 windivert_context_priority(INT64 priority64)
{
- UINT16 priority1 = (UINT16)InterlockedIncrement(&priority_counter);
- priority0 -= WINDIVERT_PRIORITY_MIN;
- return ((priority0 << 16) | ((UINT32)priority1 & 0x0000FFFF));
+ UINT32 priority, increment;
+ priority64 += WINDIVERT_PRIORITY_MAX; // Make positive
+ priority = (UINT32)(priority64 << 16);
+ increment = (UINT32)InterlockedIncrement(&priority_counter);
+ priority |= (increment & 0x0000FFFF);
+ return priority;
}
#define WINDIVERT_FILTER_WEIGHT(priority) \
- ((UINT64)(UINT32_MAX - (priority)))
+ ((UINT64)((UINT64)UINT32_MAX - (priority)))
/*
* Prototypes.
@@ -347,7 +356,7 @@ extern VOID windivert_create(IN WDFDEVICE device, IN WDFREQUEST request,
IN WDFFILEOBJECT object);
static NTSTATUS windivert_install_sublayer(layer_t layer);
static NTSTATUS windivert_install_callouts(context_t context, UINT8 layer,
- BOOL inbound, BOOL outbound, BOOL ipv4, BOOL ipv6);
+ UINT64 flags);
static NTSTATUS windivert_install_callout(context_t context, UINT idx,
layer_t layer, UINT32 *callout_id_ptr);
static void windivert_uninstall_callouts(context_t context,
@@ -412,26 +421,29 @@ static void windivert_network_classify(context_t context,
IN PWINDIVERT_NETWORK_DATA network_data, IN BOOL ipv4, IN BOOL outbound,
IN BOOL loopback, IN UINT advance, IN OUT void *data,
OUT FWPS_CLASSIFY_OUT0 *result);
-static BOOL windivert_queue_work(context_t context, PNET_BUFFER buffer,
- PNET_BUFFER_LIST buffers, PWINDIVERT_NETWORK_DATA network_data,
- PWINDIVERT_FLOW_DATA flow_data, WINDIVERT_LAYER layer,
- WINDIVERT_EVENT event, UINT64 flags, UINT32 priority, BOOL ipv4,
- BOOL outbound, BOOL loopback, BOOL impostor, BOOL match,
- LONGLONG timestamp);
+static BOOL windivert_queue_work(context_t context, PVOID packet,
+ ULONG packet_len, PNET_BUFFER_LIST buffers, WINDIVERT_LAYER layer,
+ PVOID layer_data, WINDIVERT_EVENT event, UINT64 flags, UINT32 priority,
+ BOOL ipv4, BOOL outbound, BOOL loopback, BOOL impostor, BOOL final,
+ BOOL match, LONGLONG timestamp);
static void windivert_queue_packet(context_t context, packet_t packet);
static void windivert_reinject_packet(packet_t packet);
static void windivert_free_packet(packet_t packet);
static BOOL windivert_decrement_ttl(PVOID data, BOOL ipv4, BOOL checksum);
static int windivert_big_num_compare(const UINT32 *a, const UINT32 *b);
-static BOOL windivert_filter(PNET_BUFFER buffer,
- PWINDIVERT_NETWORK_DATA network_data, PWINDIVERT_FLOW_DATA flow_data,
- BOOL ipv4, BOOL outbound, BOOL loopback, BOOL impostor, filter_t filter);
-static filter_t windivert_filter_compile(windivert_ioctl_filter_t ioctl_filter,
- size_t ioctl_filter_len);
-static void windivert_filter_analyze(filter_t filter, UINT64 flags,
- BOOL *inbound, BOOL *outbound, BOOL *ipv4, BOOL *ipv6);
-static BOOL windivert_filter_test(filter_t filter, UINT16 ip, UINT8 protocol,
- UINT8 field, UINT32 arg);
+static BOOL windivert_filter(PNET_BUFFER buffer, WINDIVERT_LAYER layer,
+ PVOID layer_data, BOOL ipv4, BOOL outbound, BOOL loopback, BOOL impostor,
+ PWINDIVERT_FILTER filter);
+static PWINDIVERT_FILTER windivert_filter_compile(
+ PWINDIVERT_FILTER ioctl_filter, size_t ioctl_filter_len);
+static NTSTATUS windivert_reflect_init(WDFOBJECT parent);
+static void windivert_reflect_close(void);
+static void windivert_reflect_event(context_t context, WINDIVERT_EVENT event);
+static void windivert_reflect_event_notify(context_t context,
+ LONGLONG timestamp, WINDIVERT_EVENT event);
+static void windivert_reflect_established_notify(context_t context,
+ LONGLONG timestamp);
+static void windivert_reflect_worker(IN WDFWORKITEM item);
/*
* WinDivert sublayer GUIDs
@@ -868,6 +880,12 @@ driver_entry_sublayer_error:
goto driver_entry_exit;
}
+ status = windivert_reflect_init((WDFOBJECT)device);
+ if (!NT_SUCCESS(status))
+ {
+ goto driver_entry_exit;
+ }
+
driver_entry_exit:
if (!NT_SUCCESS(status))
@@ -998,7 +1016,7 @@ extern VOID windivert_create(IN WDFDEVICE device, IN WDFREQUEST request,
context->packet_queue_maxtime = WINDIVERT_PARAM_QUEUE_TIME_DEFAULT;
context->layer = WINDIVERT_LAYER_DEFAULT;
context->flags = 0;
- context->priority = WINDIVERT_CONTEXT_PRIORITY(WINDIVERT_PRIORITY_DEFAULT);
+ context->priority = windivert_context_priority(WINDIVERT_PRIORITY_DEFAULT);
context->filter = NULL;
for (i = 0; i < WINDIVERT_CONTEXT_MAXWORKERS; i++)
{
@@ -1061,6 +1079,7 @@ extern VOID windivert_create(IN WDFDEVICE device, IN WDFREQUEST request,
DEBUG_ERROR("failed to create WFP engine handle", status);
goto windivert_create_exit;
}
+ RtlZeroMemory(&context->reflect, sizeof(context->reflect));
windivert_create_exit:
@@ -1092,13 +1111,19 @@ windivert_create_exit:
* Register all WFP callouts.
*/
static NTSTATUS windivert_install_callouts(context_t context, UINT8 layer,
- BOOL inbound, BOOL outbound, BOOL ipv4, BOOL ipv6)
+ UINT64 flags)
{
UINT8 i, j;
layer_t layers[WINDIVERT_CONTEXT_MAXLAYERS];
UINT32 *callout_ids[WINDIVERT_CONTEXT_MAXLAYERS] = {NULL};
+ BOOL inbound, outbound, ipv4, ipv6;
NTSTATUS status = STATUS_SUCCESS;
+ inbound = ((flags & WINDIVERT_FILTER_FLAG_INBOUND) != 0);
+ outbound = ((flags & WINDIVERT_FILTER_FLAG_OUTBOUND) != 0);
+ ipv4 = ((flags & WINDIVERT_FILTER_FLAG_IP) != 0);
+ ipv6 = ((flags & WINDIVERT_FILTER_FLAG_IPV6) != 0);
+
i = 0;
switch (layer)
{
@@ -1145,6 +1170,9 @@ static NTSTATUS windivert_install_callouts(context_t context, UINT8 layer,
}
break;
+ case WINDIVERT_LAYER_REFLECT:
+ break;
+
default:
return STATUS_INVALID_PARAMETER;
}
@@ -1408,7 +1436,6 @@ extern VOID windivert_cleanup(IN WDFFILEOBJECT object)
DEBUG("CLEANUP: cleaning up WinDivert context (context=%p)", context);
- timestamp = KeQueryPerformanceCounter(NULL).QuadPart;
KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle);
if (context->state != WINDIVERT_CONTEXT_STATE_OPENING &&
context->state != WINDIVERT_CONTEXT_STATE_OPEN)
@@ -1423,6 +1450,10 @@ windivert_cleanup_error:
sniff_mode = ((context->flags & WINDIVERT_FLAG_SNIFF) != 0);
forward = (context->layer == WINDIVERT_LAYER_NETWORK_FORWARD);
priority = context->priority;
+ KeReleaseInStackQueuedSpinLock(&lock_handle);
+ windivert_reflect_event(context, WINDIVERT_EVENT_REFLECT_CLOSE);
+ timestamp = KeQueryPerformanceCounter(NULL).QuadPart;
+ KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle);
while (!IsListEmpty(&context->flow_set))
{
entry = RemoveHeadList(&context->flow_set);
@@ -1530,7 +1561,7 @@ extern VOID windivert_destroy(IN WDFOBJECT object)
{
KLOCK_QUEUE_HANDLE lock_handle;
context_t context = windivert_context_get((WDFFILEOBJECT)object);
- filter_t filter;
+ PWINDIVERT_FILTER filter;
NTSTATUS status;
DEBUG("DESTROY: destroying WinDivert context (context=%p)", context);
@@ -1609,6 +1640,7 @@ static void windivert_read_service_request(packet_t packet, WDFREQUEST request)
{
case WINDIVERT_LAYER_NETWORK:
case WINDIVERT_LAYER_NETWORK_FORWARD:
+ case WINDIVERT_LAYER_REFLECT:
status = WdfRequestRetrieveOutputWdmMdl(request, &dst_mdl);
if (!NT_SUCCESS(status))
@@ -1624,7 +1656,14 @@ static void windivert_read_service_request(packet_t packet, WDFREQUEST request)
goto windivert_read_service_request_exit;
}
- src = WINDIVERT_PACKET_DATA_PTR(WINDIVERT_NETWORK_DATA, packet);
+ if (packet->layer != WINDIVERT_LAYER_REFLECT)
+ {
+ src = WINDIVERT_PACKET_DATA_PTR(WINDIVERT_NETWORK_DATA, packet);
+ }
+ else
+ {
+ src = WINDIVERT_PACKET_DATA_PTR(WINDIVERT_REFLECT_DATA, packet);
+ }
src_len = packet->packet_len;
dst_len = MmGetMdlByteCount(dst_mdl);
dst_len = (src_len < dst_len? src_len: dst_len);
@@ -1651,13 +1690,14 @@ static void windivert_read_service_request(packet_t packet, WDFREQUEST request)
addr->Timestamp = (INT64)packet->timestamp;
addr->Layer = packet->layer;
addr->Event = packet->event;
- addr->Outbound = (packet->outbound? 1: 0);
- addr->Loopback = (packet->loopback? 1: 0);
- addr->Impostor = (packet->impostor? 1: 0);
- addr->IPv6 = (packet->ipv6? 1: 0);
- addr->PseudoIPChecksum = (packet->pseudo_ip_checksum? 1: 0);
- addr->PseudoTCPChecksum = (packet->pseudo_tcp_checksum? 1: 0);
- addr->PseudoUDPChecksum = (packet->pseudo_udp_checksum? 1: 0);
+ addr->Outbound = packet->outbound;
+ addr->Loopback = packet->loopback;
+ addr->Impostor = packet->impostor;
+ addr->IPv6 = packet->ipv6;
+ addr->PseudoIPChecksum = packet->pseudo_ip_checksum;
+ addr->PseudoTCPChecksum = packet->pseudo_tcp_checksum;
+ addr->PseudoUDPChecksum = packet->pseudo_udp_checksum;
+ addr->Final = packet->final;
addr->Reserved = 0;
switch (packet->layer)
{
@@ -1672,6 +1712,11 @@ static void windivert_read_service_request(packet_t packet, WDFREQUEST request)
sizeof(WINDIVERT_FLOW_DATA));
break;
+ case WINDIVERT_LAYER_REFLECT:
+ RtlCopyMemory(&addr->Reflect, layer_data,
+ sizeof(WINDIVERT_REFLECT_DATA));
+ break;
+
default:
break;
}
@@ -1784,11 +1829,15 @@ static NTSTATUS windivert_write(context_t context, WDFREQUEST request,
goto windivert_write_exit;
}
- if (layer == WINDIVERT_LAYER_FLOW)
+ switch (layer)
{
- status = STATUS_INVALID_PARAMETER;
- DEBUG_ERROR("failed to inject at FLOW layer", status);
- goto windivert_write_exit;
+ case WINDIVERT_LAYER_FLOW:
+ case WINDIVERT_LAYER_REFLECT:
+ status = STATUS_INVALID_PARAMETER;
+ DEBUG_ERROR("failed to inject at FLOW layer", status);
+ goto windivert_write_exit;
+ default:
+ break;
}
status = WdfRequestRetrieveOutputWdmMdl(request, &mdl);
@@ -1994,7 +2043,7 @@ VOID windivert_caller_context(IN WDFDEVICE device, IN WDFREQUEST request)
WDF_REQUEST_PARAMETERS params;
WDFMEMORY memobj;
PWINDIVERT_ADDRESS addr = NULL;
- windivert_ioctl_t ioctl;
+ PWINDIVERT_IOCTL ioctl;
WDF_OBJECT_ATTRIBUTES attributes;
req_context_t req_context = NULL;
NTSTATUS status;
@@ -2015,14 +2064,14 @@ VOID windivert_caller_context(IN WDFDEVICE device, IN WDFREQUEST request)
goto windivert_caller_context_error;
}
- if (inbuflen != sizeof(struct windivert_ioctl_s))
+ if (inbuflen != sizeof(WINDIVERT_IOCTL))
{
status = STATUS_INVALID_PARAMETER;
DEBUG_ERROR("input buffer not an ioctl message header", status);
goto windivert_caller_context_error;
}
- ioctl = (windivert_ioctl_t)inbuf;
+ ioctl = (PWINDIVERT_IOCTL)inbuf;
if (ioctl->version != WINDIVERT_IOCTL_VERSION ||
ioctl->magic != WINDIVERT_IOCTL_MAGIC)
{
@@ -2115,11 +2164,13 @@ extern VOID windivert_ioctl(IN WDFQUEUE queue, IN WDFREQUEST request,
KLOCK_QUEUE_HANDLE lock_handle;
PCHAR inbuf, outbuf;
size_t inbuflen, outbuflen, filter0_len;
- windivert_ioctl_t ioctl;
- windivert_ioctl_filter_t filter0;
- filter_t filter;
+ PWINDIVERT_IOCTL ioctl;
+ PWINDIVERT_FILTER filter0;
+ PWINDIVERT_FILTER filter;
UINT8 layer;
- UINT32 priority;
+ INT16 priority;
+ UINT32 priority32;
+ INT64 priority64;
UINT64 flags;
PWINDIVERT_ADDRESS addr;
req_context_t req_context;
@@ -2180,7 +2231,19 @@ extern VOID windivert_ioctl(IN WDFQUEUE queue, IN WDFREQUEST request,
case IOCTL_WINDIVERT_START_FILTER:
{
BOOL inbound, outbound, ipv4, ipv6;
-
+ PIRP irp;
+ LONGLONG timestamp;
+ UINT32 process_id;
+ UINT8 filter_len;
+
+ ioctl = (PWINDIVERT_IOCTL)inbuf;
+ if ((ioctl->arg & ~WINDIVERT_FILTER_FLAGS_ALL) != 0)
+ {
+ status = STATUS_INVALID_PARAMETER;
+ DEBUG_ERROR("failed to start filter; invalid flags", status);
+ goto windivert_ioctl_exit;
+ }
+
filter = NULL;
KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle);
if (context->state != WINDIVERT_CONTEXT_STATE_OPENING)
@@ -2191,9 +2254,10 @@ windivert_ioctl_bad_start_state:
status = STATUS_INVALID_DEVICE_STATE;
goto windivert_ioctl_exit;
}
+ context->state = WINDIVERT_CONTEXT_STATE_OPEN;
KeReleaseInStackQueuedSpinLock(&lock_handle);
- filter0 = (windivert_ioctl_filter_t)outbuf;
+ filter0 = (PWINDIVERT_FILTER)outbuf;
filter0_len = outbuflen;
filter = windivert_filter_compile(filter0, filter0_len);
if (filter == NULL)
@@ -2202,9 +2266,13 @@ windivert_ioctl_bad_start_state:
DEBUG_ERROR("failed to compile filter", status);
goto windivert_ioctl_exit;
}
+ filter_len = filter0_len / sizeof(WINDIVERT_FILTER);
+ irp = WdfRequestWdmGetIrp(request);
+ process_id = (UINT32)IoGetRequestorProcessId(irp);
+ timestamp = KeQueryPerformanceCounter(NULL).QuadPart;
KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle);
- if (context->state != WINDIVERT_CONTEXT_STATE_OPENING)
+ if (context->state != WINDIVERT_CONTEXT_STATE_OPEN)
{
goto windivert_ioctl_bad_start_state;
}
@@ -2213,34 +2281,43 @@ windivert_ioctl_bad_start_state:
switch (layer)
{
case WINDIVERT_LAYER_FLOW:
+ case WINDIVERT_LAYER_REFLECT:
if ((flags & WINDIVERT_FLAG_SNIFF) == 0 ||
(flags & WINDIVERT_FLAG_RECV_ONLY) == 0)
{
goto windivert_ioctl_bad_start_state;
}
break;
+
default:
break;
}
- context->state = WINDIVERT_CONTEXT_STATE_OPEN;
- context->filter = filter;
+ context->filter = filter;
+ context->filter_len = filter_len;
+ context->reflect.data.Timestamp = timestamp;
+ context->reflect.data.ProcessId = process_id;
+ context->reflect.data.Layer = context->layer;
+ context->reflect.data.Flags = context->flags;
+ context->reflect.data.Priority = context->priority16;
+ context->reflect.inserted = FALSE;
KeReleaseInStackQueuedSpinLock(&lock_handle);
- windivert_filter_analyze(filter, flags, &inbound, &outbound,
- &ipv4, &ipv6);
- status = windivert_install_callouts(context, layer, inbound,
- outbound, ipv4, ipv6);
+ windivert_reflect_event(context, WINDIVERT_EVENT_REFLECT_OPEN);
+
+ flags = ioctl->arg;
+ status = windivert_install_callouts(context, layer, flags);
break;
}
case IOCTL_WINDIVERT_SET_LAYER:
- ioctl = (windivert_ioctl_t)inbuf;
+ ioctl = (PWINDIVERT_IOCTL)inbuf;
switch (ioctl->arg)
{
case WINDIVERT_LAYER_NETWORK:
case WINDIVERT_LAYER_NETWORK_FORWARD:
case WINDIVERT_LAYER_FLOW:
+ case WINDIVERT_LAYER_REFLECT:
break;
default:
status = STATUS_INVALID_PARAMETER;
@@ -2260,16 +2337,17 @@ windivert_ioctl_bad_start_state:
break;
case IOCTL_WINDIVERT_SET_PRIORITY:
- ioctl = (windivert_ioctl_t)inbuf;
- if (ioctl->arg < WINDIVERT_PRIORITY_MIN ||
- ioctl->arg > WINDIVERT_PRIORITY_MAX)
+ ioctl = (PWINDIVERT_IOCTL)inbuf;
+ priority64 = (INT64)ioctl->arg - WINDIVERT_PRIORITY_MAX;
+ if (priority64 < WINDIVERT_PRIORITY_MIN ||
+ priority64 > WINDIVERT_PRIORITY_MAX)
{
status = STATUS_INVALID_PARAMETER;
DEBUG_ERROR("failed to set priority; value out of range",
status);
goto windivert_ioctl_exit;
}
- priority = WINDIVERT_CONTEXT_PRIORITY((UINT32)ioctl->arg);
+ priority32 = windivert_context_priority(priority64);
KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle);
if (context->state != WINDIVERT_CONTEXT_STATE_OPENING)
{
@@ -2277,12 +2355,13 @@ windivert_ioctl_bad_start_state:
status = STATUS_INVALID_DEVICE_STATE;
goto windivert_ioctl_exit;
}
- context->priority = priority;
+ context->priority16 = (INT16)priority64;
+ context->priority = priority32;
KeReleaseInStackQueuedSpinLock(&lock_handle);
break;
case IOCTL_WINDIVERT_SET_FLAGS:
- ioctl = (windivert_ioctl_t)inbuf;
+ ioctl = (PWINDIVERT_IOCTL)inbuf;
if (!WINDIVERT_FLAGS_VALID(ioctl->arg))
{
status = STATUS_INVALID_PARAMETER;
@@ -2303,7 +2382,7 @@ windivert_ioctl_bad_start_state:
break;
case IOCTL_WINDIVERT_SET_PARAM:
- ioctl = (windivert_ioctl_t)inbuf;
+ ioctl = (PWINDIVERT_IOCTL)inbuf;
value = ioctl->arg;
KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle);
if (context->state != WINDIVERT_CONTEXT_STATE_OPEN)
@@ -2366,7 +2445,7 @@ windivert_ioctl_bad_start_state:
break;
case IOCTL_WINDIVERT_GET_PARAM:
- ioctl = (windivert_ioctl_t)inbuf;
+ ioctl = (PWINDIVERT_IOCTL)inbuf;
if (outbuflen != sizeof(UINT64))
{
status = STATUS_INVALID_PARAMETER;
@@ -2629,7 +2708,7 @@ static void windivert_network_classify(context_t context,
BOOL impostor, sniff_mode, ok;
WDFOBJECT object;
PLIST_ENTRY old_entry;
- filter_t filter;
+ PWINDIVERT_FILTER filter;
LONGLONG timestamp;
NTSTATUS status;
@@ -2716,8 +2795,8 @@ static void windivert_network_classify(context_t context,
buffer_fst = buffer;
do
{
- BOOL match = windivert_filter(buffer_fst, network_data,
- /*flow_data=*/NULL, ipv4, outbound, loopback, impostor, filter);
+ BOOL match = windivert_filter(buffer_fst, layer, (PVOID)network_data,
+ ipv4, outbound, loopback, impostor, filter);
if (match)
{
break;
@@ -2745,10 +2824,11 @@ static void windivert_network_classify(context_t context,
sniff_mode = ((flags & WINDIVERT_FLAG_SNIFF) != 0);
while (!sniff_mode && buffer_itr != buffer_fst)
{
- ok = windivert_queue_work(context, buffer_itr, buffers, network_data,
- /*flow_data=*/NULL, layer, /*event=*/WINDIVERT_EVENT_NETWORK_PACKET,
+ ok = windivert_queue_work(context, (PVOID)buffer_itr,
+ NET_BUFFER_DATA_LENGTH(buffer_itr), buffers, layer,
+ (PVOID)network_data, /*event=*/WINDIVERT_EVENT_NETWORK_PACKET,
flags, priority, ipv4, outbound, loopback, impostor,
- /*match=*/FALSE, timestamp);
+ /*final=*/FALSE, /*match=*/FALSE, timestamp);
if (!ok)
{
goto windivert_network_classify_exit;
@@ -2757,10 +2837,11 @@ static void windivert_network_classify(context_t context,
}
// STEP (2): Queue the first matching packet buffer_fst:
- ok = windivert_queue_work(context, buffer_itr, buffers, network_data,
- /*flow_data=*/NULL, layer, /*event=*/WINDIVERT_EVENT_NETWORK_PACKET,
- flags, priority, ipv4, outbound, loopback, impostor, /*match=*/TRUE,
- timestamp);
+ ok = windivert_queue_work(context, (PVOID)buffer_itr,
+ NET_BUFFER_DATA_LENGTH(buffer_itr), buffers, layer,
+ (PVOID)network_data, /*event=*/WINDIVERT_EVENT_NETWORK_PACKET,
+ flags, priority, ipv4, outbound, loopback, impostor, /*final=*/FALSE,
+ /*match=*/TRUE, timestamp);
if (advance != 0)
{
// Advance the NET_BUFFER to its original position. Note that we can
@@ -2778,12 +2859,13 @@ static void windivert_network_classify(context_t context,
buffer_itr = NET_BUFFER_NEXT_NB(buffer_fst);
while (buffer_itr != NULL)
{
- BOOL match = windivert_filter(buffer_itr, network_data,
- /*flow_data=*/NULL, ipv4, outbound, loopback, impostor, filter);
- ok = windivert_queue_work(context, buffer_itr, buffers, network_data,
- /*flow_data=*/NULL, layer, /*event=*/WINDIVERT_EVENT_NETWORK_PACKET,
- flags, priority, ipv4, outbound, loopback, impostor, match,
- timestamp);
+ BOOL match = windivert_filter(buffer_itr, layer, (PVOID)network_data,
+ ipv4, outbound, loopback, impostor, filter);
+ ok = windivert_queue_work(context, (PVOID)buffer_itr,
+ NET_BUFFER_DATA_LENGTH(buffer_itr), buffers, layer,
+ (PVOID)network_data, /*event=*/WINDIVERT_EVENT_NETWORK_PACKET,
+ flags, priority, ipv4, outbound, loopback, impostor,
+ /*FINAL=*/FALSE, match, timestamp);
if (!ok)
{
goto windivert_network_classify_exit;
@@ -2907,7 +2989,7 @@ static void windivert_flow_established_classify(context_t context,
UINT16 layer_id;
BOOL match, ok;
WDFOBJECT object;
- filter_t filter;
+ PWINDIVERT_FILTER filter;
LONGLONG timestamp;
flow_t flow;
NTSTATUS status;
@@ -2941,14 +3023,15 @@ static void windivert_flow_established_classify(context_t context,
WdfObjectReference(object);
KeReleaseInStackQueuedSpinLock(&lock_handle);
- match = windivert_filter(/*buffer=*/NULL, /*network_data=*/NULL,
- flow_data, ipv4, outbound, loopback, /*impostor=*/FALSE, filter);
+ match = windivert_filter(/*buffer=*/NULL, /*layer=*/WINDIVERT_LAYER_FLOW,
+ (PVOID)flow_data, ipv4, outbound, loopback, /*impostor=*/FALSE, filter);
if (match)
{
- ok = windivert_queue_work(context, /*buffer=*/NULL, /*buffers=*/NULL,
- /*network_data=*/NULL, flow_data, /*layer=*/WINDIVERT_LAYER_FLOW,
+ ok = windivert_queue_work(context, /*packet=*/NULL, /*packet_len=*/0,
+ /*buffers=*/NULL, /*layer=*/WINDIVERT_LAYER_FLOW, (PVOID)flow_data,
/*event=*/WINDIVERT_EVENT_FLOW_ESTABLISHED, flags, /*priority=*/0,
- ipv4, outbound, loopback, /*impostor=*/FALSE, match, timestamp);
+ ipv4, outbound, loopback, /*impostor=*/FALSE, /*final=*/FALSE,
+ match, timestamp);
if (!ok)
{
WdfObjectDereference(object);
@@ -3021,7 +3104,7 @@ static void windivert_flow_delete_notify(UINT16 layer_id, UINT32 callout_id,
BOOL match, cleanup;
WDFOBJECT object;
context_t context;
- filter_t filter;
+ PWINDIVERT_FILTER filter;
LONGLONG timestamp;
flow_t flow;
@@ -3051,16 +3134,16 @@ static void windivert_flow_delete_notify(UINT16 layer_id, UINT32 callout_id,
flags = context->flags;
KeReleaseInStackQueuedSpinLock(&lock_handle);
- match = windivert_filter(/*buffer=*/NULL, /*network_data=*/NULL,
- &flow->data, !flow->ipv6, flow->outbound, flow->loopback,
+ match = windivert_filter(/*buffer=*/NULL, /*layer=*/WINDIVERT_LAYER_FLOW,
+ (PVOID)&flow->data, !flow->ipv6, flow->outbound, flow->loopback,
/*impostor=*/FALSE, filter);
if (match)
{
- (VOID)windivert_queue_work(context, /*buffer=*/NULL, /*buffers=*/NULL,
- /*network_data=*/NULL, &flow->data, /*layer=*/WINDIVERT_LAYER_FLOW,
- /*event=*/WINDIVERT_EVENT_FLOW_DELETED, flags, /*priority=*/0,
- !flow->ipv6, flow->outbound, flow->loopback, /*impostor=*/FALSE,
- match, timestamp);
+ (VOID)windivert_queue_work(context, /*packet=*/NULL, /*packet_len=*/0,
+ /*buffers=*/NULL, /*layer=*/WINDIVERT_LAYER_FLOW,
+ (PVOID)&flow->data, /*event=*/WINDIVERT_EVENT_FLOW_DELETED, flags,
+ /*priority=*/0, !flow->ipv6, flow->outbound, flow->loopback,
+ /*impostor=*/FALSE, /*final=*/FALSE, match, timestamp);
}
windivert_flow_delete_notify_exit:
@@ -3109,20 +3192,22 @@ VOID windivert_worker(IN WDFWORKITEM item)
/*
* Queue work.
*/
-static BOOL windivert_queue_work(context_t context, PNET_BUFFER buffer,
- PNET_BUFFER_LIST buffers, PWINDIVERT_NETWORK_DATA network_data,
- PWINDIVERT_FLOW_DATA flow_data, WINDIVERT_LAYER layer,
- WINDIVERT_EVENT event, UINT64 flags, UINT32 priority, BOOL ipv4,
- BOOL outbound, BOOL loopback, BOOL impostor, BOOL match,
- LONGLONG timestamp)
+static BOOL windivert_queue_work(context_t context, PVOID packet,
+ ULONG packet_len, PNET_BUFFER_LIST buffers, WINDIVERT_LAYER layer,
+ PVOID layer_data, WINDIVERT_EVENT event, UINT64 flags, UINT32 priority,
+ BOOL ipv4, BOOL outbound, BOOL loopback, BOOL impostor, BOOL final,
+ BOOL match, LONGLONG timestamp)
{
KLOCK_QUEUE_HANDLE lock_handle;
+ PNET_BUFFER buffer;
packet_t work;
- ULONG packet_len;
PVOID packet_data;
UINT8 *data;
PLIST_ENTRY old_entry;
NDIS_TCP_IP_CHECKSUM_NET_BUFFER_LIST_INFO checksums;
+ PWINDIVERT_NETWORK_DATA network_data;
+ PWINDIVERT_FLOW_DATA flow_data;
+ PWINDIVERT_REFLECT_DATA reflect_data;
BOOL pseudo_ip_checksum, pseudo_tcp_checksum, pseudo_udp_checksum;
if (!match && (flags & WINDIVERT_FLAG_SNIFF) != 0)
@@ -3139,7 +3224,8 @@ static BOOL windivert_queue_work(context_t context, PNET_BUFFER buffer,
{
case WINDIVERT_LAYER_NETWORK:
case WINDIVERT_LAYER_NETWORK_FORWARD:
- packet_len = NET_BUFFER_DATA_LENGTH(buffer);
+ buffer = (PNET_BUFFER)packet;
+ network_data = (PWINDIVERT_NETWORK_DATA)layer_data;
if (packet_len > UINT16_MAX)
{
// Cannot handle oversized packet
@@ -3185,6 +3271,7 @@ static BOOL windivert_queue_work(context_t context, PNET_BUFFER buffer,
break;
case WINDIVERT_LAYER_FLOW:
+ flow_data = (PWINDIVERT_FLOW_DATA)layer_data;
work = (packet_t)windivert_malloc(
WINDIVERT_PACKET_SIZE(WINDIVERT_FLOW_DATA, 0), FALSE);
if (work == NULL)
@@ -3198,6 +3285,24 @@ static BOOL windivert_queue_work(context_t context, PNET_BUFFER buffer,
FALSE;
break;
+ case WINDIVERT_LAYER_REFLECT:
+ reflect_data = (PWINDIVERT_REFLECT_DATA)layer_data;
+ work = (packet_t)windivert_malloc(
+ WINDIVERT_PACKET_SIZE(WINDIVERT_REFLECT_DATA, packet_len),
+ FALSE);
+ if (work == NULL)
+ {
+ return TRUE;
+ }
+ work->packet_len = packet_len;
+ data = WINDIVERT_LAYER_DATA_PTR(work);
+ RtlCopyMemory(data, reflect_data, sizeof(WINDIVERT_REFLECT_DATA));
+ data = WINDIVERT_PACKET_DATA_PTR(WINDIVERT_REFLECT_DATA, work);
+ RtlCopyMemory(data, packet, packet_len);
+ pseudo_ip_checksum = TRUE;
+ pseudo_tcp_checksum = pseudo_udp_checksum = FALSE;
+ break;
+
default:
return TRUE;
}
@@ -3211,6 +3316,7 @@ static BOOL windivert_queue_work(context_t context, PNET_BUFFER buffer,
work->pseudo_ip_checksum = (pseudo_ip_checksum? 1: 0);
work->pseudo_tcp_checksum = (pseudo_tcp_checksum? 1: 0);
work->pseudo_udp_checksum = (pseudo_udp_checksum? 1: 0);
+ work->final = (final? 1: 0);
work->match = match;
work->priority = priority;
work->timestamp = timestamp;
@@ -3235,7 +3341,7 @@ static BOOL windivert_queue_work(context_t context, PNET_BUFFER buffer,
context->worker_curr =
(context->worker_curr + 1) % WINDIVERT_CONTEXT_MAXWORKERS;
KeReleaseInStackQueuedSpinLock(&lock_handle);
-
+
if (old_entry != NULL)
{
work = CONTAINING_RECORD(old_entry, struct packet_s, entry);
@@ -3505,6 +3611,11 @@ static BOOL windivert_parse_headers(PNET_BUFFER buffer, BOOL ipv4,
NTSTATUS status;
// Parse the headers:
+ if (buffer == NULL)
+ {
+ DEBUG("FILTER: REJECT (packet is NULL)");
+ return FALSE;
+ }
tot_len = NET_BUFFER_DATA_LENGTH(buffer);
if (tot_len < sizeof(WINDIVERT_IPHDR))
{
@@ -3660,9 +3771,9 @@ static BOOL windivert_parse_headers(PNET_BUFFER buffer, BOOL ipv4,
/*
* Checks if the given network packet is of interest.
*/
-static BOOL windivert_filter(PNET_BUFFER buffer,
- PWINDIVERT_NETWORK_DATA network_data, PWINDIVERT_FLOW_DATA flow_data,
- BOOL ipv4, BOOL outbound, BOOL loopback, BOOL impostor, filter_t filter)
+static BOOL windivert_filter(PNET_BUFFER buffer, WINDIVERT_LAYER layer,
+ PVOID layer_data, BOOL ipv4, BOOL outbound, BOOL loopback, BOOL impostor,
+ PWINDIVERT_FILTER filter)
{
PWINDIVERT_IPHDR ip_header = NULL;
PWINDIVERT_IPV6HDR ipv6_header = NULL;
@@ -3672,21 +3783,32 @@ static BOOL windivert_filter(PNET_BUFFER buffer,
PWINDIVERT_UDPHDR udp_header = NULL;
UINT payload_len = 0;
UINT16 ip, ttl;
+ PWINDIVERT_NETWORK_DATA network_data = NULL;
+ PWINDIVERT_FLOW_DATA flow_data = NULL;
+ PWINDIVERT_REFLECT_DATA reflect_data = NULL;
NTSTATUS status;
- if (network_data != NULL)
+ switch (layer)
{
- if (!windivert_parse_headers(buffer, ipv4, &ip_header, &ipv6_header,
- &icmp_header, &icmpv6_header, &tcp_header, &udp_header,
- &payload_len))
- {
+ case WINDIVERT_LAYER_NETWORK:
+ case WINDIVERT_LAYER_NETWORK_FORWARD:
+ if (!windivert_parse_headers(buffer, ipv4, &ip_header, &ipv6_header,
+ &icmp_header, &icmpv6_header, &tcp_header, &udp_header,
+ &payload_len))
+ {
+ return FALSE;
+ }
+ network_data = (PWINDIVERT_NETWORK_DATA)layer_data;
+ break;
+ case WINDIVERT_LAYER_FLOW:
+ flow_data = (PWINDIVERT_FLOW_DATA)layer_data;
+ break;
+ case WINDIVERT_LAYER_REFLECT:
+ reflect_data = (PWINDIVERT_REFLECT_DATA)layer_data;
+ break;
+ default:
+ DEBUG("FILTER: REJECT (invalid parameter)");
return FALSE;
- }
- }
- else if (flow_data == NULL)
- {
- DEBUG("FILTER: REJECT (invalid parameter)");
- return FALSE;
}
// Execute the filter:
@@ -3701,38 +3823,122 @@ static BOOL windivert_filter(PNET_BUFFER buffer,
field[1] = 0;
field[2] = 0;
field[3] = 0;
- switch (filter[ip].protocol)
+
+ switch (filter[ip].field)
{
- case WINDIVERT_FILTER_PROTOCOL_NONE:
+ case WINDIVERT_FILTER_FIELD_ZERO:
result = TRUE;
break;
- case WINDIVERT_FILTER_PROTOCOL_NETWORK:
- result = (network_data != NULL);
+ case WINDIVERT_FILTER_FIELD_INBOUND:
+ case WINDIVERT_FILTER_FIELD_OUTBOUND:
+ case WINDIVERT_FILTER_FIELD_LOOPBACK:
+ case WINDIVERT_FILTER_FIELD_IMPOSTOR:
+ case WINDIVERT_FILTER_FIELD_IP:
+ case WINDIVERT_FILTER_FIELD_IPV6:
+ case WINDIVERT_FILTER_FIELD_ICMP:
+ case WINDIVERT_FILTER_FIELD_ICMPV6:
+ case WINDIVERT_FILTER_FIELD_TCP:
+ case WINDIVERT_FILTER_FIELD_UDP:
+ result = (layer != WINDIVERT_LAYER_REFLECT);
break;
- case WINDIVERT_FILTER_PROTOCOL_FLOW:
- result = (flow_data != NULL);
+ case WINDIVERT_FILTER_FIELD_IFIDX:
+ case WINDIVERT_FILTER_FIELD_SUBIFIDX:
+ result = (layer == WINDIVERT_LAYER_NETWORK ||
+ layer == WINDIVERT_LAYER_NETWORK_FORWARD);
+ result = result && (network_data != NULL);
break;
- case WINDIVERT_FILTER_PROTOCOL_IP:
- result = (ip_header != NULL);
+ case WINDIVERT_FILTER_FIELD_LOCALADDR:
+ case WINDIVERT_FILTER_FIELD_REMOTEADDR:
+ case WINDIVERT_FILTER_FIELD_LOCALPORT:
+ case WINDIVERT_FILTER_FIELD_REMOTEPORT:
+ case WINDIVERT_FILTER_FIELD_PROTOCOL:
+ result = (layer == WINDIVERT_LAYER_FLOW);
+ result = result && (flow_data != NULL);
break;
- case WINDIVERT_FILTER_PROTOCOL_IPV6:
- result = (ipv6_header != NULL);
+ case WINDIVERT_FILTER_FIELD_PROCESSID:
+ result = ((layer == WINDIVERT_LAYER_FLOW &&
+ flow_data != NULL) ||
+ (layer == WINDIVERT_LAYER_REFLECT &&
+ reflect_data != NULL));
break;
- case WINDIVERT_FILTER_PROTOCOL_ICMP:
- result = (icmp_header != NULL);
+ case WINDIVERT_FILTER_FIELD_LAYER:
+ result = (layer == WINDIVERT_LAYER_REFLECT);
+ result = result && (reflect_data != NULL);
break;
- case WINDIVERT_FILTER_PROTOCOL_ICMPV6:
- result = (icmpv6_header != NULL);
+ case WINDIVERT_FILTER_FIELD_IP_HDRLENGTH:
+ case WINDIVERT_FILTER_FIELD_IP_TOS:
+ case WINDIVERT_FILTER_FIELD_IP_LENGTH:
+ case WINDIVERT_FILTER_FIELD_IP_ID:
+ case WINDIVERT_FILTER_FIELD_IP_DF:
+ case WINDIVERT_FILTER_FIELD_IP_MF:
+ case WINDIVERT_FILTER_FIELD_IP_FRAGOFF:
+ case WINDIVERT_FILTER_FIELD_IP_TTL:
+ case WINDIVERT_FILTER_FIELD_IP_PROTOCOL:
+ case WINDIVERT_FILTER_FIELD_IP_CHECKSUM:
+ case WINDIVERT_FILTER_FIELD_IP_SRCADDR:
+ case WINDIVERT_FILTER_FIELD_IP_DSTADDR:
+ result = (layer == WINDIVERT_LAYER_NETWORK ||
+ layer == WINDIVERT_LAYER_NETWORK_FORWARD);
+ result = result && (ip_header != NULL);
break;
- case WINDIVERT_FILTER_PROTOCOL_TCP:
- result = (tcp_header != NULL);
+ case WINDIVERT_FILTER_FIELD_IPV6_TRAFFICCLASS:
+ case WINDIVERT_FILTER_FIELD_IPV6_FLOWLABEL:
+ case WINDIVERT_FILTER_FIELD_IPV6_LENGTH:
+ case WINDIVERT_FILTER_FIELD_IPV6_NEXTHDR:
+ case WINDIVERT_FILTER_FIELD_IPV6_HOPLIMIT:
+ case WINDIVERT_FILTER_FIELD_IPV6_SRCADDR:
+ case WINDIVERT_FILTER_FIELD_IPV6_DSTADDR:
+ result = (layer == WINDIVERT_LAYER_NETWORK ||
+ layer == WINDIVERT_LAYER_NETWORK_FORWARD);
+ result = result && (ipv6_header != NULL);
break;
- case WINDIVERT_FILTER_PROTOCOL_UDP:
- result = (udp_header != NULL);
+ case WINDIVERT_FILTER_FIELD_ICMP_TYPE:
+ case WINDIVERT_FILTER_FIELD_ICMP_CODE:
+ case WINDIVERT_FILTER_FIELD_ICMP_CHECKSUM:
+ case WINDIVERT_FILTER_FIELD_ICMP_BODY:
+ result = (layer == WINDIVERT_LAYER_NETWORK ||
+ layer == WINDIVERT_LAYER_NETWORK_FORWARD);
+ result = result && (icmp_header != NULL);
+ break;
+ case WINDIVERT_FILTER_FIELD_ICMPV6_TYPE:
+ case WINDIVERT_FILTER_FIELD_ICMPV6_CODE:
+ case WINDIVERT_FILTER_FIELD_ICMPV6_CHECKSUM:
+ case WINDIVERT_FILTER_FIELD_ICMPV6_BODY:
+ result = (layer == WINDIVERT_LAYER_NETWORK ||
+ layer == WINDIVERT_LAYER_NETWORK_FORWARD);
+ result = result && (icmpv6_header != NULL);
+ break;
+ case WINDIVERT_FILTER_FIELD_TCP_SRCPORT:
+ case WINDIVERT_FILTER_FIELD_TCP_DSTPORT:
+ case WINDIVERT_FILTER_FIELD_TCP_SEQNUM:
+ case WINDIVERT_FILTER_FIELD_TCP_ACKNUM:
+ case WINDIVERT_FILTER_FIELD_TCP_HDRLENGTH:
+ case WINDIVERT_FILTER_FIELD_TCP_URG:
+ case WINDIVERT_FILTER_FIELD_TCP_ACK:
+ case WINDIVERT_FILTER_FIELD_TCP_PSH:
+ case WINDIVERT_FILTER_FIELD_TCP_RST:
+ case WINDIVERT_FILTER_FIELD_TCP_SYN:
+ case WINDIVERT_FILTER_FIELD_TCP_FIN:
+ case WINDIVERT_FILTER_FIELD_TCP_WINDOW:
+ case WINDIVERT_FILTER_FIELD_TCP_CHECKSUM:
+ case WINDIVERT_FILTER_FIELD_TCP_URGPTR:
+ case WINDIVERT_FILTER_FIELD_TCP_PAYLOADLENGTH:
+ result = (layer == WINDIVERT_LAYER_NETWORK ||
+ layer == WINDIVERT_LAYER_NETWORK_FORWARD);
+ result = result && (tcp_header != NULL);
+ break;
+ case WINDIVERT_FILTER_FIELD_UDP_SRCPORT:
+ case WINDIVERT_FILTER_FIELD_UDP_DSTPORT:
+ case WINDIVERT_FILTER_FIELD_UDP_LENGTH:
+ case WINDIVERT_FILTER_FIELD_UDP_CHECKSUM:
+ case WINDIVERT_FILTER_FIELD_UDP_PAYLOADLENGTH:
+ result = (layer == WINDIVERT_LAYER_NETWORK ||
+ layer == WINDIVERT_LAYER_NETWORK_FORWARD);
+ result = result && (udp_header != NULL);
break;
default:
- error = TRUE;
result = FALSE;
+ error = TRUE;
break;
}
if (result)
@@ -3971,7 +4177,12 @@ static BOOL windivert_filter(PNET_BUFFER buffer,
field[0] = (UINT32)flow_data->Protocol;
break;
case WINDIVERT_FILTER_FIELD_PROCESSID:
- field[0] = flow_data->ProcessId;
+ field[0] = (flow_data != NULL?
+ flow_data->ProcessId:
+ reflect_data->ProcessId);
+ break;
+ case WINDIVERT_FILTER_FIELD_LAYER:
+ field[0] = reflect_data->Layer;
break;
default:
error = TRUE;
@@ -4028,166 +4239,28 @@ static BOOL windivert_filter(PNET_BUFFER buffer,
return FALSE;
}
-/*
- * Analyze the given filter.
- */
-static void windivert_filter_analyze(filter_t filter, UINT64 flags,
- BOOL *inbound, BOOL *outbound, BOOL *ipv4, BOOL *ipv6)
-{
- BOOL result;
-
- // Send-only?
- if ((flags & WINDIVERT_FLAG_SEND_ONLY) != 0)
- {
-windivert_filter_analyze_send_only:
- *inbound = FALSE;
- *outbound = FALSE;
- *ipv4 = FALSE;
- *ipv6 = FALSE;
- return;
- }
-
- // False filter?
- result = windivert_filter_test(filter, 0, WINDIVERT_FILTER_PROTOCOL_NONE,
- WINDIVERT_FILTER_FIELD_ZERO, 0);
- if (!result)
- {
- goto windivert_filter_analyze_send_only;
- }
-
- // Inbound?
- result = windivert_filter_test(filter, 0, WINDIVERT_FILTER_PROTOCOL_NONE,
- WINDIVERT_FILTER_FIELD_INBOUND, 1);
- if (result)
- {
- result = windivert_filter_test(filter, 0,
- WINDIVERT_FILTER_PROTOCOL_NONE, WINDIVERT_FILTER_FIELD_OUTBOUND,
- 0);
- }
- *inbound = result;
-
- // Outbound?
- result = windivert_filter_test(filter, 0, WINDIVERT_FILTER_PROTOCOL_NONE,
- WINDIVERT_FILTER_FIELD_OUTBOUND, 1);
- if (result)
- {
- result = windivert_filter_test(filter, 0,
- WINDIVERT_FILTER_PROTOCOL_NONE, WINDIVERT_FILTER_FIELD_INBOUND, 0);
- }
- *outbound = result;
-
- // IPv4?
- result = windivert_filter_test(filter, 0, WINDIVERT_FILTER_PROTOCOL_NONE,
- WINDIVERT_FILTER_FIELD_IP, 1);
- if (result)
- {
- result = windivert_filter_test(filter, 0,
- WINDIVERT_FILTER_PROTOCOL_NONE, WINDIVERT_FILTER_FIELD_IPV6, 0);
- }
- *ipv4 = result;
-
- // Ipv6?
- result = windivert_filter_test(filter, 0, WINDIVERT_FILTER_PROTOCOL_NONE,
- WINDIVERT_FILTER_FIELD_IPV6, 1);
- if (result)
- {
- result = windivert_filter_test(filter, 0,
- WINDIVERT_FILTER_PROTOCOL_NONE, WINDIVERT_FILTER_FIELD_IP, 0);
- }
- *ipv6 = result;
-}
-
-/*
- * Test a filter for any packet where field = arg.
- */
-static BOOL windivert_filter_test(filter_t filter, UINT16 ip, UINT8 protocol,
- UINT8 field, UINT32 arg)
-{
- BOOL known = FALSE;
- BOOL result = FALSE;
-
- if (ip == WINDIVERT_FILTER_RESULT_ACCEPT)
- {
- return TRUE;
- }
- if (ip == WINDIVERT_FILTER_RESULT_REJECT)
- {
- return FALSE;
- }
- if (ip > WINDIVERT_FILTER_MAXLEN)
- {
- return FALSE;
- }
-
- if (filter[ip].protocol == protocol &&
- filter[ip].field == field)
- {
- known = TRUE;
- switch (filter[ip].test)
- {
- case WINDIVERT_FILTER_TEST_EQ:
- result = (arg == filter[ip].arg[0]);
- break;
- case WINDIVERT_FILTER_TEST_NEQ:
- result = (arg != filter[ip].arg[0]);
- break;
- case WINDIVERT_FILTER_TEST_LT:
- result = (arg < filter[ip].arg[0]);
- break;
- case WINDIVERT_FILTER_TEST_LEQ:
- result = (arg <= filter[ip].arg[0]);
- break;
- case WINDIVERT_FILTER_TEST_GT:
- result = (arg > filter[ip].arg[0]);
- break;
- case WINDIVERT_FILTER_TEST_GEQ:
- result = (arg >= filter[ip].arg[0]);
- break;
- default:
- result = FALSE;
- break;
- }
- }
-
- if (!known)
- {
- result = windivert_filter_test(filter, filter[ip].success, protocol,
- field, arg);
- if (result)
- {
- return TRUE;
- }
- return windivert_filter_test(filter, filter[ip].failure, protocol,
- field, arg);
- }
- else
- {
- ip = (result? filter[ip].success: filter[ip].failure);
- return windivert_filter_test(filter, ip, protocol, field, arg);
- }
-}
-
/*
* Compile a WinDivert filter from an IOCTL.
*/
-static filter_t windivert_filter_compile(windivert_ioctl_filter_t ioctl_filter,
- size_t ioctl_filter_len)
+static PWINDIVERT_FILTER windivert_filter_compile(
+ PWINDIVERT_FILTER ioctl_filter, size_t ioctl_filter_len)
{
- filter_t filter = NULL;
+ PWINDIVERT_FILTER filter = NULL;
UINT16 i;
size_t length;
- if (ioctl_filter_len % sizeof(struct windivert_ioctl_filter_s) != 0)
+ if (ioctl_filter_len % sizeof(WINDIVERT_FILTER) != 0)
{
goto windivert_filter_compile_error;
}
- length = ioctl_filter_len / sizeof(struct windivert_ioctl_filter_s);
+ length = ioctl_filter_len / sizeof(WINDIVERT_FILTER);
if (length >= WINDIVERT_FILTER_MAXLEN || length == 0)
{
goto windivert_filter_compile_error;
}
- filter = (filter_t)windivert_malloc(length*sizeof(struct filter_s), FALSE);
+ filter = (PWINDIVERT_FILTER)windivert_malloc(
+ length * sizeof(WINDIVERT_FILTER), FALSE);
if (filter == NULL)
{
goto windivert_filter_compile_error;
@@ -4275,6 +4348,12 @@ static filter_t windivert_filter_compile(windivert_ioctl_filter_t ioctl_filter,
goto windivert_filter_compile_error;
}
break;
+ case WINDIVERT_FILTER_FIELD_LAYER:
+ if (ioctl_filter[i].arg[0] > WINDIVERT_LAYER_MAX)
+ {
+ goto windivert_filter_compile_error;
+ }
+ break;
case WINDIVERT_FILTER_FIELD_IP_HDRLENGTH:
case WINDIVERT_FILTER_FIELD_TCP_HDRLENGTH:
if (ioctl_filter[i].arg[0] > 0x0F)
@@ -4345,96 +4424,6 @@ static filter_t windivert_filter_compile(windivert_ioctl_filter_t ioctl_filter,
filter[i].arg[1] = ioctl_filter[i].arg[1];
filter[i].arg[2] = ioctl_filter[i].arg[2];
filter[i].arg[3] = ioctl_filter[i].arg[3];
-
- // Protocol selection:
- switch (ioctl_filter[i].field)
- {
- case WINDIVERT_FILTER_FIELD_ZERO:
- case WINDIVERT_FILTER_FIELD_INBOUND:
- case WINDIVERT_FILTER_FIELD_OUTBOUND:
- case WINDIVERT_FILTER_FIELD_LOOPBACK:
- case WINDIVERT_FILTER_FIELD_IMPOSTOR:
- case WINDIVERT_FILTER_FIELD_IP:
- case WINDIVERT_FILTER_FIELD_IPV6:
- case WINDIVERT_FILTER_FIELD_ICMP:
- case WINDIVERT_FILTER_FIELD_ICMPV6:
- case WINDIVERT_FILTER_FIELD_TCP:
- case WINDIVERT_FILTER_FIELD_UDP:
- filter[i].protocol = WINDIVERT_FILTER_PROTOCOL_NONE;
- break;
- case WINDIVERT_FILTER_FIELD_IFIDX:
- case WINDIVERT_FILTER_FIELD_SUBIFIDX:
- filter[i].protocol = WINDIVERT_FILTER_PROTOCOL_NETWORK;
- break;
- case WINDIVERT_FILTER_FIELD_LOCALADDR:
- case WINDIVERT_FILTER_FIELD_REMOTEADDR:
- case WINDIVERT_FILTER_FIELD_LOCALPORT:
- case WINDIVERT_FILTER_FIELD_REMOTEPORT:
- case WINDIVERT_FILTER_FIELD_PROTOCOL:
- filter[i].protocol = WINDIVERT_FILTER_PROTOCOL_FLOW;
- break;
- case WINDIVERT_FILTER_FIELD_IP_HDRLENGTH:
- case WINDIVERT_FILTER_FIELD_IP_TOS:
- case WINDIVERT_FILTER_FIELD_IP_LENGTH:
- case WINDIVERT_FILTER_FIELD_IP_ID:
- case WINDIVERT_FILTER_FIELD_IP_DF:
- case WINDIVERT_FILTER_FIELD_IP_MF:
- case WINDIVERT_FILTER_FIELD_IP_FRAGOFF:
- case WINDIVERT_FILTER_FIELD_IP_TTL:
- case WINDIVERT_FILTER_FIELD_IP_PROTOCOL:
- case WINDIVERT_FILTER_FIELD_IP_CHECKSUM:
- case WINDIVERT_FILTER_FIELD_IP_SRCADDR:
- case WINDIVERT_FILTER_FIELD_IP_DSTADDR:
- filter[i].protocol = WINDIVERT_FILTER_PROTOCOL_IP;
- break;
- case WINDIVERT_FILTER_FIELD_IPV6_TRAFFICCLASS:
- case WINDIVERT_FILTER_FIELD_IPV6_FLOWLABEL:
- case WINDIVERT_FILTER_FIELD_IPV6_LENGTH:
- case WINDIVERT_FILTER_FIELD_IPV6_NEXTHDR:
- case WINDIVERT_FILTER_FIELD_IPV6_HOPLIMIT:
- case WINDIVERT_FILTER_FIELD_IPV6_SRCADDR:
- case WINDIVERT_FILTER_FIELD_IPV6_DSTADDR:
- filter[i].protocol = WINDIVERT_FILTER_PROTOCOL_IPV6;
- break;
- case WINDIVERT_FILTER_FIELD_ICMP_TYPE:
- case WINDIVERT_FILTER_FIELD_ICMP_CODE:
- case WINDIVERT_FILTER_FIELD_ICMP_CHECKSUM:
- case WINDIVERT_FILTER_FIELD_ICMP_BODY:
- filter[i].protocol = WINDIVERT_FILTER_PROTOCOL_ICMP;
- break;
- case WINDIVERT_FILTER_FIELD_ICMPV6_TYPE:
- case WINDIVERT_FILTER_FIELD_ICMPV6_CODE:
- case WINDIVERT_FILTER_FIELD_ICMPV6_CHECKSUM:
- case WINDIVERT_FILTER_FIELD_ICMPV6_BODY:
- filter[i].protocol = WINDIVERT_FILTER_PROTOCOL_ICMPV6;
- break;
- case WINDIVERT_FILTER_FIELD_TCP_SRCPORT:
- case WINDIVERT_FILTER_FIELD_TCP_DSTPORT:
- case WINDIVERT_FILTER_FIELD_TCP_SEQNUM:
- case WINDIVERT_FILTER_FIELD_TCP_ACKNUM:
- case WINDIVERT_FILTER_FIELD_TCP_HDRLENGTH:
- case WINDIVERT_FILTER_FIELD_TCP_URG:
- case WINDIVERT_FILTER_FIELD_TCP_ACK:
- case WINDIVERT_FILTER_FIELD_TCP_PSH:
- case WINDIVERT_FILTER_FIELD_TCP_RST:
- case WINDIVERT_FILTER_FIELD_TCP_SYN:
- case WINDIVERT_FILTER_FIELD_TCP_FIN:
- case WINDIVERT_FILTER_FIELD_TCP_WINDOW:
- case WINDIVERT_FILTER_FIELD_TCP_CHECKSUM:
- case WINDIVERT_FILTER_FIELD_TCP_URGPTR:
- case WINDIVERT_FILTER_FIELD_TCP_PAYLOADLENGTH:
- filter[i].protocol = WINDIVERT_FILTER_PROTOCOL_TCP;
- break;
- case WINDIVERT_FILTER_FIELD_UDP_SRCPORT:
- case WINDIVERT_FILTER_FIELD_UDP_DSTPORT:
- case WINDIVERT_FILTER_FIELD_UDP_LENGTH:
- case WINDIVERT_FILTER_FIELD_UDP_CHECKSUM:
- case WINDIVERT_FILTER_FIELD_UDP_PAYLOADLENGTH:
- filter[i].protocol = WINDIVERT_FILTER_PROTOCOL_UDP;
- break;
- default:
- goto windivert_filter_compile_error;
- }
}
return filter;
@@ -4445,3 +4434,330 @@ windivert_filter_compile_error:
return NULL;
}
+/****************************************************************************/
+/* WINDIVERT REFLECT MANAGER IMPLEMENTATION */
+/****************************************************************************/
+
+#include "windivert_shared.c"
+
+/*
+ * WinDivert reflect state.
+ */
+static BOOL reflect_inited = FALSE; // Reflection initialized?
+static KSPIN_LOCK reflect_lock; // Reflect lock.
+static LIST_ENTRY reflect_event_queue; // Reflect event queue.
+static LIST_ENTRY reflect_contexts; // All open (non-REFLECT) contexts.
+static LIST_ENTRY reflect_waiters; // All open REFLECT contexts.
+static WDFWORKITEM reflect_worker; // Reflect work item.
+
+/*
+ * Initialize the reflection layer implementation.
+ */
+static NTSTATUS windivert_reflect_init(WDFOBJECT parent)
+{
+ WDF_WORKITEM_CONFIG item_config;
+ WDF_OBJECT_ATTRIBUTES obj_attrs;
+ NTSTATUS status;
+
+ KeInitializeSpinLock(&reflect_lock);
+ InitializeListHead(&reflect_event_queue);
+ InitializeListHead(&reflect_contexts);
+ InitializeListHead(&reflect_waiters);
+ WDF_WORKITEM_CONFIG_INIT(&item_config, windivert_reflect_worker);
+ item_config.AutomaticSerialization = TRUE;
+ WDF_OBJECT_ATTRIBUTES_INIT(&obj_attrs);
+ obj_attrs.ParentObject = parent;
+ status = WdfWorkItemCreate(&item_config, &obj_attrs, &reflect_worker);
+ if (!NT_SUCCESS(status))
+ {
+ DEBUG_ERROR("failed to create reflection work item", status);
+ return status;
+ }
+ reflect_inited = TRUE;
+ return STATUS_SUCCESS;
+}
+
+/*
+ * Cleanup the reflection layer implementation.
+ */
+static void windivert_reflect_close(void)
+{
+ if (!reflect_inited)
+ {
+ return;
+ }
+ WdfWorkItemFlush(reflect_worker);
+ WdfObjectDelete(reflect_worker);
+}
+
+/*
+ * WinDivert handle reflect event.
+ */
+static void windivert_reflect_event(context_t context, WINDIVERT_EVENT event)
+{
+ KLOCK_QUEUE_HANDLE lock_handle;
+ WDFOBJECT object;
+ reflect_event_t reflect_event;
+
+ KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle);
+ object = (WDFOBJECT)context->object;
+ if (event == WINDIVERT_EVENT_REFLECT_OPEN)
+ {
+ // To be released on WINDIVERT_EVENT_REFLECT_CLOSE. This ensures the
+ // context object remains valid until the close event has been handled.
+ WdfObjectReference(object);
+ }
+ KeReleaseInStackQueuedSpinLock(&lock_handle);
+
+ // Queue the event:
+ reflect_event = (reflect_event_t)windivert_malloc(
+ sizeof(struct reflect_event_s), FALSE);
+ if (reflect_event == NULL)
+ {
+ WdfObjectDereference(object);
+ return;
+ }
+ reflect_event->context = context;
+ reflect_event->event = event;
+ KeAcquireInStackQueuedSpinLock(&reflect_lock, &lock_handle);
+ InsertTailList(&reflect_event_queue, &reflect_event->entry);
+ KeReleaseInStackQueuedSpinLock(&lock_handle);
+ WdfWorkItemEnqueue(reflect_worker);
+}
+
+/*
+ * Create REFLECT layer "pseudo" packet to pass the filter.
+ */
+static PWINDIVERT_IPHDR windivert_reflect_pseudo_packet(context_t context,
+ ULONG *len_ptr)
+{
+ KLOCK_QUEUE_HANDLE lock_handle;
+ UINT16 total_len;
+ UINT8 *packet;
+ char *object;
+ PWINDIVERT_FILTER filter;
+ UINT8 filter_len;
+ PWINDIVERT_IPHDR iphdr;
+ WINDIVERT_STREAM stream;
+
+ // The filter is returned in a pseudo-IP packet. This is just to make
+ // the interface consistent, i.e., WinDivertRecv() always receives IP
+ // packets.
+
+ total_len = sizeof(WINDIVERT_IPHDR) + WINDIVERT_OBJECT_MAXLEN;
+ packet = windivert_malloc(total_len, TRUE);
+ if (packet == NULL)
+ {
+ return NULL;
+ }
+
+ iphdr = (PWINDIVERT_IPHDR)packet;
+ object = (char *)(iphdr + 1);
+
+ stream.data = object;
+ stream.pos = 0;
+ stream.max = WINDIVERT_OBJECT_MAXLEN;
+ stream.overflow = FALSE;
+
+ KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle);
+ filter = context->filter;
+ filter_len = context->filter_len;
+ KeReleaseInStackQueuedSpinLock(&lock_handle);
+
+ WinDivertSerializeFilter(&stream, filter, filter_len);
+
+ if (stream.overflow)
+ {
+ windivert_free(packet);
+ return NULL;
+ }
+
+ total_len = sizeof(WINDIVERT_IPHDR) + (UINT16)stream.pos;
+ RtlZeroMemory(iphdr, sizeof(WINDIVERT_IPHDR));
+ iphdr->Version = 4;
+ iphdr->HdrLength = sizeof(WINDIVERT_IPHDR) / sizeof(UINT32);
+ iphdr->Length = RtlUshortByteSwap(total_len);
+ iphdr->TTL = 1;
+ iphdr->Protocol = 254; // "experimental"
+
+ *len_ptr = total_len;
+
+ return iphdr;
+}
+
+/*
+ * Notify all REFLECT layer contexts a new event.
+ */
+static void windivert_reflect_event_notify(context_t context,
+ LONGLONG timestamp, WINDIVERT_EVENT event)
+{
+ KLOCK_QUEUE_HANDLE lock_handle;
+ PLIST_ENTRY entry;
+ context_t waiter;
+ PWINDIVERT_FILTER filter;
+ PWINDIVERT_IPHDR packet = NULL;
+ ULONG packet_len;
+ BOOL match;
+
+ entry = reflect_waiters.Flink;
+ while (entry != &reflect_waiters)
+ {
+ waiter = CONTAINING_RECORD(entry, struct context_s, reflect.entry);
+ entry = entry->Flink;
+ KeAcquireInStackQueuedSpinLock(&waiter->lock, &lock_handle);
+ filter = waiter->filter;
+ KeReleaseInStackQueuedSpinLock(&lock_handle);
+ match = windivert_filter(/*buffer=*/NULL,
+ /*layer=*/WINDIVERT_LAYER_REFLECT, (PVOID)&context->reflect.data,
+ /*ipv4=*/TRUE, /*outbound=*/FALSE, /*loopback=*/FALSE,
+ /*impostor=*/FALSE, filter);
+ if (!match)
+ {
+ continue;
+ }
+ if (packet == NULL)
+ {
+ packet = windivert_reflect_pseudo_packet(context, &packet_len);
+ if (packet == NULL)
+ {
+ return;
+ }
+ }
+ (VOID)windivert_queue_work(waiter, (PVOID)packet, packet_len,
+ /*buffers=*/NULL, /*layer=*/WINDIVERT_LAYER_REFLECT,
+ (PVOID)&context->reflect.data, event, /*flags=*/0, /*priority=*/0,
+ /*ipv4=*/TRUE, /*outbound=*/FALSE, /*loopback=*/FALSE,
+ /*impostor=*/FALSE, /*final=*/FALSE, /*match=*/TRUE, timestamp);
+ }
+
+ windivert_free(packet);
+}
+
+/*
+ * Notify a new REFLECT layer context of all existing open handles.
+ */
+static void windivert_reflect_established_notify(context_t context,
+ LONGLONG timestamp)
+{
+ KLOCK_QUEUE_HANDLE lock_handle;
+ PLIST_ENTRY entry;
+ BOOL match, ok, final;
+ context_t waiter;
+ PWINDIVERT_FILTER filter;
+ PWINDIVERT_IPHDR packet;
+ ULONG packet_len;
+
+ KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle);
+ filter = context->filter;
+ KeReleaseInStackQueuedSpinLock(&lock_handle);
+
+ entry = reflect_contexts.Flink;
+ while (entry != &reflect_contexts)
+ {
+ waiter = CONTAINING_RECORD(entry, struct context_s, reflect.entry);
+ entry = entry->Flink;
+ match = windivert_filter(/*buffer=*/NULL,
+ /*layer=*/WINDIVERT_LAYER_REFLECT, (PVOID)&waiter->reflect.data,
+ /*ipv4=*/TRUE, /*outbound=*/FALSE, /*loopback=*/FALSE,
+ /*impostor=*/FALSE, filter);
+ if (!match)
+ {
+ continue;
+ }
+ packet = windivert_reflect_pseudo_packet(waiter, &packet_len);
+ if (packet == NULL)
+ {
+ continue;
+ }
+ final = (entry == &reflect_contexts);
+ ok = windivert_queue_work(context, (PVOID)packet, packet_len,
+ /*buffers=*/NULL, /*layer=*/WINDIVERT_LAYER_REFLECT,
+ (PVOID)&waiter->reflect.data,
+ /*event=*/WINDIVERT_EVENT_REFLECT_ESTABLISHED, /*flags=*/0,
+ /*priority=*/0, /*ipv4=*/TRUE, /*outbound=*/FALSE,
+ /*loopback=*/FALSE, /*impostor=*/FALSE, final, /*match=*/TRUE,
+ timestamp);
+ windivert_free(packet);
+ if (!ok)
+ {
+ break;
+ }
+ }
+}
+
+/*
+ * WinDivert REFLECT worker.
+ */
+static void windivert_reflect_worker(IN WDFWORKITEM item)
+{
+ KLOCK_QUEUE_HANDLE lock_handle;
+ PLIST_ENTRY entry;
+ context_t context;
+ LONGLONG timestamp;
+ WINDIVERT_EVENT event;
+ reflect_event_t reflect_event;
+ WDFOBJECT object;
+ WINDIVERT_LAYER layer;
+
+ // All reflection events are serialized and handled by this worker.
+ // This ensures that we are always operating on a consistent "snapshot"
+ // of the WinDivert handle state. This worker also has exclusive control
+ // over reflect_contexts/reflect_waiters, so locking is not required.
+
+ KeAcquireInStackQueuedSpinLock(&reflect_lock, &lock_handle);
+ while (!IsListEmpty(&reflect_event_queue))
+ {
+ entry = RemoveHeadList(&reflect_event_queue);
+ KeReleaseInStackQueuedSpinLock(&lock_handle);
+
+ reflect_event = CONTAINING_RECORD(entry, struct reflect_event_s, entry);
+ context = reflect_event->context;
+ event = reflect_event->event;
+ windivert_free(reflect_event);
+
+ DEBUG("REFLECT: %s event for WinDivert context (context=%p)",
+ (event == WINDIVERT_EVENT_REFLECT_OPEN? "open": "close"), context);
+
+ KeAcquireInStackQueuedSpinLock(&context->lock, &lock_handle);
+ object = (WDFOBJECT)context->object;
+ layer = context->layer;
+ KeReleaseInStackQueuedSpinLock(&lock_handle);
+
+ timestamp = KeQueryPerformanceCounter(NULL).QuadPart;
+ switch (event)
+ {
+ case WINDIVERT_EVENT_REFLECT_OPEN:
+ context->reflect.inserted = TRUE;
+ if (layer != WINDIVERT_LAYER_REFLECT)
+ {
+ InsertTailList(&reflect_contexts, &context->reflect.entry);
+ }
+ else
+ {
+ InsertTailList(&reflect_waiters, &context->reflect.entry);
+ windivert_reflect_established_notify(context, timestamp);
+ }
+ break;
+
+ case WINDIVERT_EVENT_REFLECT_CLOSE:
+ if (context->reflect.inserted)
+ {
+ RemoveEntryList(&context->reflect.entry);
+ }
+ break;
+ }
+
+ if (layer != WINDIVERT_LAYER_REFLECT)
+ {
+ windivert_reflect_event_notify(context, timestamp, event);
+ }
+ if (event == WINDIVERT_EVENT_REFLECT_CLOSE)
+ {
+ WdfObjectDereference(object);
+ }
+
+ KeAcquireInStackQueuedSpinLock(&reflect_lock, &lock_handle);
+ }
+ KeReleaseInStackQueuedSpinLock(&lock_handle);
+}
+
diff --git a/test/test.c b/test/test.c
index edd7455..f2bd160 100644
--- a/test/test.c
+++ b/test/test.c
@@ -241,6 +241,9 @@ static struct test tests[] =
"false): false): false): false)", &pkt_http_request, TRUE},
{"(outbound? (ip? (tcp.DstPort == 80? (tcp.PayloadLength == 0? true: "
"false): false): false): false)", &pkt_http_request, FALSE},
+ {"(ipv6? tcp and tcp.DstPort = 1234 and (tcp.SrcPort = 999? !tcp.UrgPtr: "
+ "tcp.Syn) or udp: ip and tcp.DstPort == 80)",
+ &pkt_http_request, TRUE},
{"udp", &pkt_dns_request, TRUE},
{"udp && udp.SrcPort > 1 && ipv6", &pkt_dns_request, FALSE},
{"udp.DstPort == 53", &pkt_dns_request, TRUE},
@@ -388,21 +391,26 @@ static BOOL run_test(HANDLE inject_handle, const char *filter,
OVERLAPPED overlapped;
const char *err_str;
UINT err_pos;
+ PWINDIVERT_IPHDR iphdr = NULL;
HANDLE handle = INVALID_HANDLE_VALUE, handle0 = INVALID_HANDLE_VALUE,
event = NULL;
// (0) Verify the test data:
- if (!WinDivertHelperCheckFilter(filter, WINDIVERT_LAYER_NETWORK, &err_str,
- &err_pos))
+ if (!WinDivertHelperCompileFilter(filter, WINDIVERT_LAYER_NETWORK,
+ NULL, 0, &err_str, &err_pos))
{
fprintf(stderr, "error: filter string \"%s\" is invalid with error "
"\"%s\" (position=%u)\n", filter, err_str, err_pos);
goto failed;
}
+ WinDivertHelperParsePacket((PVOID)packet, packet_len, &iphdr, NULL,
+ NULL, NULL, NULL, NULL, NULL, NULL);
memset(&addr, 0, sizeof(addr));
- addr.Direction = WINDIVERT_DIRECTION_OUTBOUND;
- if (WinDivertHelperEvalFilter(filter, WINDIVERT_LAYER_NETWORK,
- (PVOID)packet, packet_len, &addr) != match)
+ addr.Outbound = TRUE;
+ addr.Layer = WINDIVERT_LAYER_NETWORK;
+ addr.IPv6 = (iphdr == NULL);
+ if (WinDivertHelperEvalFilter(filter, (PVOID)packet, packet_len, &addr)
+ != match)
{
fprintf(stderr, "error: filter \"%s\" does not match the given "
"packet\n", filter);
@@ -481,7 +489,7 @@ read_failed:
}
buf_len = (UINT)iolen;
}
- if (addr.Direction == WINDIVERT_DIRECTION_OUTBOUND)
+ if (addr.Outbound)
{
WinDivertHelperCalcChecksums(buf, buf_len, NULL, 0);
}