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