// dllmain.cpp
#include <windows.h>
#include <d3d11.h>
#include <dxgi.h>
#include <intrin.h>
#include <tlhelp32.h>
#include <psapi.h>

#pragma comment(lib, "d3d11.lib")
#pragma comment(lib, "dxgi.lib")
#pragma comment(lib, "psapi.lib")

static void* g_OriginalPresent = nullptr;
static ID3D11Device* g_Device = nullptr;

static void* DetourFunc(void* target, void* hook, size_t len)
{
    DWORD oldProtect;
    VirtualProtect(target, len, PAGE_EXECUTE_READWRITE, &oldProtect);

    void* gateway = VirtualAlloc(0, len + 14, MEM_COMMIT | MEM_RESERVE, PAGE_EXECUTE_READWRITE);
    memcpy(gateway, target, len);
    *(BYTE*)((ULONG_PTR)gateway + len) = 0xE9;
    *(DWORD*)((ULONG_PTR)gateway + len + 1) = (DWORD)((ULONG_PTR)target + len - ((ULONG_PTR)gateway + len + 5));

    memset(target, 0x90, len);
    *(BYTE*)target = 0xE9;
    *(DWORD*)((ULONG_PTR)target + 1) = (DWORD)((ULONG_PTR)hook - ((ULONG_PTR)target + 5));

    VirtualProtect(target, len, oldProtect, &oldProtect);
    return gateway;
}

typedef HRESULT(WINAPI* Present_t)(IDXGISwapChain*, UINT, UINT);

static HRESULT WINAPI HookedPresent(IDXGISwapChain* swapChain, UINT syncInterval, UINT flags)
{
    syncInterval = 0;
    flags |= DXGI_PRESENT_DO_NOT_WAIT;

    static BOOL first = TRUE;
    if (first)
    {
        first = FALSE;
        swapChain->GetDevice(__uuidof(ID3D11Device), (void**)&g_Device);
        if (g_Device)
        {
            IDXGIDevice1* dxgiDevice = nullptr;
            if (SUCCEEDED(g_Device->QueryInterface(__uuidof(IDXGIDevice1), (void**)&dxgiDevice)))
            {
                dxgiDevice->SetMaximumFrameLatency(1);
                dxgiDevice->Release();
            }
        }
    }

    return ((Present_t)g_OriginalPresent)(swapChain, syncInterval, flags);
}

static void BoostThreads(void)
{
    DWORD pid = GetCurrentProcessId();
    HANDLE snap = CreateToolhelp32Snapshot(TH32CS_SNAPTHREAD, 0);
    if (snap == INVALID_HANDLE_VALUE) return;

    THREADENTRY32 te;
    memset(&te, 0, sizeof(te));
    te.dwSize = sizeof(THREADENTRY32);

    if (Thread32First(snap, &te))
    {
        do
        {
            if (te.th32OwnerProcessID == pid)
            {
                HANDLE hTh = OpenThread(THREAD_SET_INFORMATION, FALSE, te.th32ThreadID);
                if (hTh)
                {
                    SetThreadPriority(hTh, THREAD_PRIORITY_TIME_CRITICAL);
                    CloseHandle(hTh);
                }
            }
        } while (Thread32Next(snap, &te));
    }

    CloseHandle(snap);
}

BOOL APIENTRY DllMain(HMODULE hModule, DWORD reason, LPVOID lpReserved)
{
    if (reason != DLL_PROCESS_ATTACH) return TRUE;

    DisableThreadLibraryCalls(hModule);

    HMODULE hDXGI = GetModuleHandleA("dxgi.dll");
    if (hDXGI)
    {
        void* presentAddr = GetProcAddress(hDXGI, "Present");
        if (presentAddr)
        {
            g_OriginalPresent = DetourFunc(presentAddr, HookedPresent, 5);
        }
    }

    SetPriorityClass(GetCurrentProcess(), REALTIME_PRIORITY_CLASS);
    BoostThreads();

    typedef LONG(NTAPI* NtSetTimer_t)(ULONG, BOOLEAN, PULONG);
    HMODULE hNtdll = GetModuleHandleA("ntdll.dll");
    NtSetTimer_t NtSetTimerResolution = (NtSetTimer_t)GetProcAddress(hNtdll, "NtSetTimerResolution");
    if (NtSetTimerResolution)
    {
        ULONG actual;
        NtSetTimerResolution(5000, TRUE, &actual);
    }

    return TRUE;
}