#pragma once #ifdef WS2INTERCEPTDLL_EXPORTS #define WS2INTERCEPTDLL_API __declspec(dllexport) #else #define WS2INTERCEPTDLL_API __declspec(dllimport) #endif using namespace std; #define DRIVER_SERVICE_NAME _T("dnsfirewall") static PCTSTR c_szDnsFireWallSecName = _T("dns_firewall_settings"); static PCTSTR c_szDnsPluginStateSecName = _T("dns_plugin_state_settings"); #define DNS_PORT 53 inline u_long hostAddress( u_long nOctet1, u_long nOctet2, u_long nOctet3, u_long nOctet4 ) { u_long nRet; nRet = (nOctet4 & 0xFF) + ((nOctet3 & 0xFF) << 8) + ((nOctet2 & 0xFF) << 16) + ((nOctet1 & 0xFF) << 24); return nRet; } inline u_long netAddress( u_long nOctet1, u_long nOctet2, u_long nOctet3, u_long nOctet4 ) { u_long nRet; nRet = (nOctet1 & 0xFF) + ((nOctet2 & 0xFF) << 8) + ((nOctet3 & 0xFF) << 16) + ((nOctet4 & 0xFF) << 24); return nRet; } static const UINT32 c_uiTestSubnet = netAddress(85, 21, 0, 0); static const UINT32 c_uiTestSubnetMask = netAddress(255, 255, 0, 0); #define DNSFIREWALL_FILE_COUNT 4 #define DNSFIREWALL_THREAD_COUNT 8 #define DNSFIREWALL_COMPLETION_ENTRIES_COUNT 32 typedef int WSAAPI tf_recv( _In_ SOCKET s, _Out_writes_bytes_to_(len, return) __out_data_source(NETWORK) char FAR * buf, _In_ int len, _In_ int flags ); typedef int WSAAPI tf_select( _In_ int nfds, _Inout_opt_ fd_set FAR * readfds, _Inout_opt_ fd_set FAR * writefds, _Inout_opt_ fd_set FAR * exceptfds, _In_opt_ const struct timeval FAR * timeout ); typedef int WSAAPI tf_send( _In_ SOCKET s, _In_reads_bytes_(len) const char FAR * buf, _In_ int len, _In_ int flags ); typedef SOCKET WSAAPI tf_socket( _In_ int af, _In_ int type, _In_ int protocol ); typedef SOCKET WSAAPI tf_accept( _In_ SOCKET s, _Out_ struct sockaddr *addr, _Inout_ int *addrlen ); typedef int WSAAPI tf_bind( _In_ SOCKET s, _In_ const struct sockaddr *name, _In_ int namelen ); typedef int WSAAPI tf_listen( _In_ SOCKET s, _In_ int backlog ); typedef int WSAAPI tf_sendto( _In_ SOCKET s, _In_ const char *buf, _In_ int len, _In_ int flags, _In_ const struct sockaddr *to, _In_ int tolen ); typedef int WSAAPI tf_closesocket( _In_ SOCKET s ); typedef int WSAAPI tf_WSARecvFrom( _In_ SOCKET s, _Inout_ LPWSABUF lpBuffers, _In_ DWORD dwBufferCount, _Out_ LPDWORD lpNumberOfBytesRecvd, _Inout_ LPDWORD lpFlags, _Out_ struct sockaddr *lpFrom, _Inout_ LPINT lpFromlen, _In_ LPWSAOVERLAPPED lpOverlapped, _In_ LPWSAOVERLAPPED_COMPLETION_ROUTINE lpCompletionRoutine ); typedef SOCKET WSAAPI tf_WSASocketW( _In_ int af, _In_ int type, _In_ int protocol, _In_opt_ LPWSAPROTOCOL_INFOW lpProtocolInfo, _In_ GROUP g, _In_ DWORD dwFlags ); typedef int WSAAPI tf_WSACleanup(void); typedef int WSAAPI tf_WSAStartup( _In_ WORD wVersionRequested, _Out_ LPWSADATA lpWSAData ); typedef WINBASEAPI _When_(lpModuleName == NULL, _Ret_notnull_) _When_(lpModuleName != NULL, _Ret_maybenull_) HMODULE WINAPI tf_GetModuleHandleW( _In_opt_ LPCWSTR lpModuleName ); typedef WINADVAPI SERVICE_STATUS_HANDLE WINAPI tf_RegisterServiceCtrlHandlerW( _In_ LPCWSTR lpServiceName, _In_ __callback LPHANDLER_FUNCTION lpHandlerProc ); #include "pshpack1.h" struct TWSAGetOverlappedResultStub { BYTE m_baStubCode[15]; // copy of stub code WSAGetOverlappedResult WORD m_wJumpCode; // Jump Instruction code = 0x25FF DWORD m_dwJumpAddress; // relative offset of jump address = 0 PBYTE m_pJumpAddress; // Jump address }; typedef TWSAGetOverlappedResultStub* PTWSAGetOverlappedResultStub; #include "poppack.h" // Back to 4 byte packing #include "afxmt.h" typedef BOOL WINAPI tf_GetQueuedCompletionStatus( _In_ HANDLE CompletionPort, _Out_ LPDWORD lpNumberOfBytesTransferred, _Out_ PULONG_PTR lpCompletionKey, _Out_ LPOVERLAPPED * lpOverlapped, _In_ DWORD dwMilliseconds ); typedef HANDLE WINAPI tf_CreateIoCompletionPort( _In_ HANDLE FileHandle, _In_opt_ HANDLE ExistingCompletionPort, _In_ ULONG_PTR CompletionKey, _In_ DWORD NumberOfConcurrentThreads ); typedef struct _DNS_APPROVAL_CONFIG { BOOLEAN m_bDiscard; BOOLEAN m_bDataModify; UINT32 m_nDataSize; } DNS_APPROVAL_CONFIG, *PDNS_APPROVAL_CONFIG; #pragma pack(push, 1) typedef struct _TCP_DNS_HEADER { WORD m_wMessageLength; DNS_HEADER m_stDnsHeader; } TCP_DNS_HEADER, *PTCP_DNS_HEADER; #pragma pack(pop) __inline bool bfIPv6AddressInIPv6Subnet( const IN6_ADDR* pIPv6Address, const PIPV6_ADDRESS_STRUCT pIPv6Subnet ) { const UINT32 nMaskIndex = (pIPv6Subnet->_PrefixLength << 1); const PUINT64 pDoubleQword = (PUINT64) pIPv6Address; const PUINT64 aIPv6SubnetArray = (PUINT64) pIPv6Subnet; UINT64 aIPv6Address[2]; bool bReturn; aIPv6Address[0] = pDoubleQword[0]; aIPv6Address[1] = pDoubleQword[1]; aIPv6Address[0] &= c_aQWMaskArray[nMaskIndex]; aIPv6Address[1] &= c_aQWMaskArray[nMaskIndex + 1]; bReturn = (aIPv6SubnetArray[0] == aIPv6Address[0]); bReturn = bReturn && (aIPv6SubnetArray[1] == aIPv6Address[1]); return bReturn; } __inline void IPv6CopyAddress( IN6_ADDR* pDestAddress, const IP6_ADDRESS* pSrcAddress ) { const PUINT64 pDoubleQwordSrc = (PUINT64) pSrcAddress; PUINT64 pDoubleQwordDest = (PUINT64) pDestAddress; pDoubleQwordDest[0] = pDoubleQwordSrc[0]; pDoubleQwordDest[1] = pDoubleQwordSrc[1]; } // // Byte swap routines. These are used to convert from little-endian to // big-endian and vice-versa. // #ifdef __cplusplus extern "C" { #endif unsigned short __cdecl _byteswap_ushort(_In_ unsigned short); unsigned long __cdecl _byteswap_ulong(_In_ unsigned long); unsigned __int64 __cdecl _byteswap_uint64(_In_ unsigned __int64); #ifdef __cplusplus } #endif #pragma intrinsic(_byteswap_ushort) #pragma intrinsic(_byteswap_ulong) #pragma intrinsic(_byteswap_uint64) #define RtlUshortByteSwap(_x) _byteswap_ushort((USHORT)(_x)) #define RtlUlongByteSwap(_x) _byteswap_ulong((_x)) #define RtlUlonglongByteSwap(_x) _byteswap_uint64((_x)) enum EDnsSocketState { enNotInitialized = 0, enUnknownState, enErrorOperation, enIp4Created, enIp6Created, enIp4Binded, enIp6Binded, enDnsIp4Binded, enDnsIp6Binded, enDnsIp4Connected, enDnsIp6Connected, enDnsIp4Listen, enDnsIp6Listen, enDnsIp4Accepted, enDnsIp6Accepted, enDnsTcp4RecvQuery, enDnsTcp6RecvQuery, enDnsIp4RecvData, enDnsIp6RecvData, enDnsRecvNonBlocking, enDnsRecvFromNonBlocking, enDnsTcp4SendResponce, enDnsTcp6SendResponce, enDnsIp4SendData, enDnsIp6SendData, enDnsUdp4RecvQueryPended, enDnsUdp6RecvQueryPended, enDnsUdp4RecvDataPended, enDnsUdp6RecvDataPended, enDnsUdp4RecvQuery, enDnsUdp6RecvQuery, enDnsUdp4SendResponce, enDnsUdp6SendResponce, enDnsUdp4SendData, enDnsUdp6SendData, enDnsUdp4RecvQueryStarted, enDnsUdp6RecvQueryStarted, enDnsUdp4RecvDataStarted, enDnsUdp6RecvDataStarted, enDnsUdp4Received, enDnsUdp6Received, enDnsUdp4RecvQueryDone, enDnsUdp6RecvQueryDone, enDnsUdp4RecvDataDone, enDnsUdp6RecvDataDone, enDnsIp4Selected, enDnsIp6Selected, enDnsTcp4AcceptSelected, enDnsTcp6AcceptSelected, enDnsOverlappedRecvStarted, enDnsPreCreateIoCompletionPort, enDnsCreateIoCompletionPort, enDnsUdp4RecvQueryComplete, enDnsUdp6RecvQueryComplete, enClosed }; class CDnsMessageWB; typedef struct _DNS_APPROVAL_OUTPUT { DNS_APPROVAL_CONFIG m_stApprovalConfig; BYTE m_pData[65536]; } DNS_APPROVAL_OUTPUT, *PDNS_APPROVAL_OUTPUT; struct TRecvQueuedEntry; struct TDnsSocket { TDnsSocket() : m_hSocket(NULL) , m_enSocketState(enNotInitialized) , m_nAddressFamily(0) , m_nSocketType(0) , m_nSocketProtocol(0) , m_nErrorCode(0) , m_bWSACreated(false) , m_bOverlapedIO(false) , m_bBindedToDnsPort(false) , m_bAcceptedSocket(false) , m_hListenSocket(NULL) , m_bSocketBind(false) , m_pProtocolInfo(NULL) , m_bOverlappedRecvStarted(false) , m_hIoCompletionPort(NULL) , m_ulCompletionKey(0) , m_bAbortedQuery(false) , m_bPermitQuery(false) , m_bModifyQuery(false) , m_pDnsQueryMessage(NULL) , m_pCurrRecvQueuedEntry(NULL) , m_bAcceptLocal(false) , m_bTcpFirstPartReceived(false) , m_pDnsApprovalOutput(NULL) , m_nFirstPartLen(0) , m_bDnsDiscardQuery(false) , m_pZoneTransferMessage(NULL) , m_bRecvLengthOnly(false) { RtlZeroMemory( &m_stSocketBindAddress, sizeof(m_stSocketBindAddress) ); RtlZeroMemory( &m_stSocketClientAddress, sizeof(m_stSocketClientAddress) ); } ~TDnsSocket(); SOCKET m_hSocket; EDnsSocketState m_enSocketState; SOCKADDR_STORAGE m_stSocketBindAddress; SOCKADDR_STORAGE m_stSocketClientAddress; int m_nAddressFamily; int m_nSocketType; int m_nSocketProtocol; int m_nErrorCode; bool m_bWSACreated; bool m_bOverlapedIO; bool m_bBindedToDnsPort; bool m_bAcceptedSocket; SOCKET m_hListenSocket; bool m_bSocketBind; WSAPROTOCOL_INFOW* m_pProtocolInfo; bool m_bOverlappedRecvStarted; HANDLE m_hIoCompletionPort; ULONG_PTR m_ulCompletionKey; bool m_bAbortedQuery; bool m_bPermitQuery; bool m_bModifyQuery; CDnsMessageWB* m_pDnsQueryMessage; TRecvQueuedEntry* m_pCurrRecvQueuedEntry; bool m_bAcceptLocal; bool m_bTcpFirstPartReceived; BYTE m_baDnsQueryBuffer[4096]; DNS_APPROVAL_OUTPUT* m_pDnsApprovalOutput; int m_nFirstPartLen; bool m_bDnsDiscardQuery; CDnsMessageWB* m_pZoneTransferMessage; bool m_bRecvLengthOnly; WORD m_wRecvMsgLength; }; typedef unordered_map TDnsSocketArray; struct TRecvQueuedEntry { TRecvQueuedEntry() : m_pOverlapped(NULL) , m_hSocket(NULL) , m_pRecvFromAddress(NULL) , m_pBuffers(NULL) , m_dwBufferCount(0) , m_ulCompletionKey(0) , m_pNumberOfBytesRecvd(NULL) , m_bRemoteIsLocal(false) { RtlZeroMemory( &m_stSocketClientAddress, sizeof(m_stSocketClientAddress) ); } OVERLAPPED* m_pOverlapped; SOCKET m_hSocket; SOCKADDR* m_pRecvFromAddress; LPINT m_pRecvFromAddressLen; SOCKADDR_STORAGE m_stSocketClientAddress; WSABUF* m_pBuffers; DWORD m_dwBufferCount; ULONG_PTR m_ulCompletionKey; DWORD* m_pNumberOfBytesRecvd; bool m_bRemoteIsLocal; }; typedef unordered_map CTRecvQueuedArray; typedef pair TRecvQueuedInsertPair; struct TUdpRemoteClient { TUdpRemoteClient() : m_pDnsSocket(NULL) , m_nQueryCount(0) { } TDnsSocket* m_pDnsSocket; int m_nQueryCount; }; typedef unordered_map CTUdpRemoteClientArray; typedef pair TUdpRemoteClientInsertPair; // CWS2InterceptAPP class CDnsStateThread; class CDnsResorceRecord; class CWS2InterceptAPP : public CWinApp { DECLARE_DYNCREATE(CWS2InterceptAPP) public: CWS2InterceptAPP(); virtual ~CWS2InterceptAPP(); public: virtual BOOL InitInstance(); virtual int ExitInstance(); protected: DECLARE_MESSAGE_MAP() public: CModule32HLP m_oModule32HLP; tf_accept* m_pfSavedAcceptFunction; tf_bind* m_pfSavedBindFunction; tf_closesocket* m_pfSavedClosesocketFunction; tf_recv* m_pfSavedRecvFunction; tf_send* m_pfSavedSendFunction; tf_sendto* m_pfSavedSendtoFunction; tf_socket* m_pfSavedSocketFunction; tf_WSACleanup* m_pfSavedWSACleanupFunction; tf_WSARecvFrom* m_pfSavedWSARecvFromFunction; tf_WSASocketW* m_pfSavedWSASocketWFunction; tf_WSAStartup* m_pfSavedWSAStartupFunction; SOCKET mf_accept( SOCKET s, struct sockaddr* addr, int* addrlen ); int mf_bind( SOCKET s, const struct sockaddr* name, int namelen ); int mf_closesocket(SOCKET s); int mf_recv( SOCKET s, char* buf, int len, int flags ); int mf_send( SOCKET s, const char* buf, int len, int flags ); int mf_sendto( SOCKET s, const char* buf, int len, int flags, const struct sockaddr* to, int tolen ); SOCKET mf_socket( int af, int type, int protocol ); int mf_WSACleanup(); int mf_WSARecvFrom( SOCKET s, LPWSABUF lpBuffers, DWORD dwBufferCount, LPDWORD lpNumberOfBytesRecvd, LPDWORD lpFlags, struct sockaddr* lpFrom, LPINT lpFromlen, LPWSAOVERLAPPED lpOverlapped, LPWSAOVERLAPPED_COMPLETION_ROUTINE lpCompletionRoutine ); SOCKET mf_WSASocketW( int af, int type, int protocol, LPWSAPROTOCOL_INFOW lpProtocolInfo, GROUP g, DWORD dwFlags ); int mf_WSAStartup( WORD wVersionRequested, LPWSADATA lpWSAData ); NTSTATUS UpdateImports( PBYTE oldFunction, PBYTE newFunction ); static SOCKET WSAAPI sf_accept( SOCKET s, struct sockaddr* addr, int* addrlen ); static int WSAAPI sf_bind( SOCKET s, const struct sockaddr* name, int namelen ); static int WSAAPI sf_closesocket(SOCKET s); static int WSAAPI sf_recv( SOCKET s, char* buf, int len, int flags ); static int WSAAPI sf_send( SOCKET s, const char* buf, int len, int flags ); static int WSAAPI sf_sendto( SOCKET s, const char* buf, int len, int flags, const struct sockaddr* to, int tolen ); static SOCKET WSAAPI sf_socket( int af, int type, int protocol ); static int WSAAPI sf_WSACleanup(); static int WSAAPI sf_WSARecvFrom( SOCKET s, LPWSABUF lpBuffers, DWORD dwBufferCount, LPDWORD lpNumberOfBytesRecvd, LPDWORD lpFlags, struct sockaddr* lpFrom, LPINT lpFromlen, LPWSAOVERLAPPED lpOverlapped, LPWSAOVERLAPPED_COMPLETION_ROUTINE lpCompletionRoutine ); static SOCKET WSAAPI sf_WSASocketW( int af, int type, int protocol, LPWSAPROTOCOL_INFOW lpProtocolInfo, GROUP g, DWORD dwFlags ); static int WSAAPI sf_WSAStartup( WORD wVersionRequested, LPWSADATA lpWSAData ); NTSTATUS UpdateWS2_32Imports( PBYTE oldFunction, PBYTE newFunction ); TDnsSocketArray m_oDnsSocketArray; tf_GetQueuedCompletionStatus* m_pfSavedGetQueuedCompletionStatusFunction; static BOOL WINAPI sf_GetQueuedCompletionStatus( HANDLE CompletionPort, LPDWORD lpNumberOfBytes, PULONG_PTR lpCompletionKey, LPOVERLAPPED* lpOverlapped, DWORD dwMilliseconds ); BOOL mf_GetQueuedCompletionStatus( HANDLE CompletionPort, LPDWORD lpNumberOfBytes, PULONG_PTR lpCompletionKey, LPOVERLAPPED* lpOverlapped, DWORD dwMilliseconds ); tf_CreateIoCompletionPort* m_pfSavedCreateIoCompletionPortFunction; static HANDLE WINAPI sf_CreateIoCompletionPort( HANDLE FileHandle, HANDLE ExistingCompletionPort, ULONG_PTR CompletionKey, DWORD NumberOfConcurrentThreads ); HANDLE mf_CreateIoCompletionPort( HANDLE FileHandle, HANDLE ExistingCompletionPort, ULONG_PTR CompletionKey, DWORD NumberOfConcurrentThreads ); CTRecvQueuedArray m_oRecvQueuedArray; CString m_sDnsZoneName; int m_nPublicIpCount; UINT32* m_pPublicIpArray; int m_nPublicRangeCount; FWP_RANGE0* m_pPublicRangeArray; int m_nPublicSubnetCount; FWP_V4_ADDR_AND_MASK* m_pPublicSubnetArray; int m_nPublicIPV6AddrCount; IPV6_ADDRESS_STRUCT* m_pPublicIPV6AddrArray; IPV6_ADDRESS_STRUCT* m_pPrivateIPV6AddrArray; int m_nPrivateIpCount; UINT32* m_pPrivateIpArray; int m_nPrivateRangeCount; FWP_RANGE0* m_pPrivateRangeArray; int m_nPrivateSubnetCount; FWP_V4_ADDR_AND_MASK* m_pPrivateSubnetArray; int m_nPrivateIPV6AddrCount; bool IsRemoteIPv4Local(ULONG ulIp4NAddr); bool IsRemoteIPv6Local(IN6_ADDR* pIPv6Address); CDnsResorceRecord* m_pDnsZoneRecord; CString m_sNamePrimaryServer; CString m_sNameAdministrator; IN_ADDR m_stIp4TestAddr; tf_RegisterServiceCtrlHandlerW* m_pfSavedRegisterServiceCtrlHandlerWFunction; static SERVICE_STATUS_HANDLE WINAPI sf_RegisterServiceCtrlHandlerW( LPCWSTR lpServiceName, LPHANDLER_FUNCTION lpHandlerProc ); SERVICE_STATUS_HANDLE mf_RegisterServiceCtrlHandlerW( LPCWSTR lpServiceName, LPHANDLER_FUNCTION lpHandlerProc ); CEvent m_oDnsSvcStopedEvent; LPHANDLER_FUNCTION m_pfSavedSvcHandleFunction; static void WINAPI sf_SvcHandler(DWORD fdwControl); void mf_SvcHandler(DWORD fdwControl); NTSTATUS UpdateImports( PCSTR szCalledModuleName, PCSTR szFunctionName, PBYTE pNewEntryPoint, PBYTE& rOldEntryPoint ); DWORD m_dwSerialNo; CEvent m_oStateThreadInitDoneEvent; CMutex m_oPluginStateLock; CDnsStateThread* m_pDnsStateThread; CTUdpRemoteClientArray m_oUdpRemoteClientArray; CCriticalSection m_oDataLock; CCriticalSection m_oMainThreadLock; NTSTATUS UpdateImports(PCSTR szFunctionName, PBYTE pNewEntryPoint, PBYTE& rOldEntryPoint); }; extern CWS2InterceptAPP theApp;