// This Source Code Form is subject to the terms of the
// Mozilla Public License, v. 2.0. If a copy of the MPL was not distributed
// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.

const std = @import("std");
const builtin = @import("builtin");
const native_arch = builtin.cpu.arch;
const win = std.os.windows;

const WINAPI: std.builtin.CallingConvention =
    if (native_arch == .x86) .{ .x86_stdcall = .{} } else .c;

var TARGET_HASH: [16]win.BYTE = .{0} ** 16;

var self_module: ?win.HMODULE = null;

const dll_name = std.unicode.utf8ToUtf16LeStringLiteral("autopunch.dll");
const msg_title = std.unicode.utf8ToUtf16LeStringLiteral("autopunch");

// FFI
extern "user32" fn MessageBoxW(
    hWnd: ?win.HWND,
    lpText: ?win.LPCWSTR,
    lpCaption: ?win.LPCWSTR,
    uType: u32,
) callconv(WINAPI) c_int;

extern "kernel32" fn LoadLibraryExW(
    lpLibFileName: win.LPCWSTR,
    hFile: ?win.HANDLE,
    dwFlags: win.DWORD,
) callconv(WINAPI) ?win.HMODULE;

extern "kernel32" fn GetModuleFileNameW(
    hModule: ?win.HMODULE,
    lpFilename: win.LPCWSTR,
    nSize: win.DWORD,
) callconv(WINAPI) win.DWORD;

// HELPERS
fn showMsg(text: win.LPCWSTR) void {
    const MB_ICONINFORMATION = 0x40;
    _ = MessageBoxW(null, text, msg_title, MB_ICONINFORMATION);
}

fn buildDllPath(buf: []win.WCHAR) ?[*:0]win.WCHAR {
    const module = self_module orelse return null;

    const len = GetModuleFileNameW(module, @ptrCast(buf.ptr), @intCast(buf.len));
    if (len == 0 or len >= buf.len) return null;

    const path_slice = buf[0..len];
    const last_sep = std.mem.lastIndexOfScalar(win.WCHAR, path_slice, '\\') orelse return null;
    const dir_end = last_sep + 1; // includes '\'

    const name_len = std.mem.indexOfScalar(win.WCHAR, dll_name, 0) orelse dll_name.len;

    if (dir_end + name_len >= buf.len) return null;

    @memcpy(buf[dir_end .. dir_end + name_len], dll_name[0..name_len]);
    buf[dir_end + name_len] = 0;

    return @ptrCast(buf.ptr);
}

fn loadThread() void {
    var path_buf: [win.MAX_PATH]win.WCHAR = undefined;

    const path = buildDllPath(&path_buf) orelse {
        showMsg(std.unicode.utf8ToUtf16LeStringLiteral("Failed to build autopunch path."));
        return;
    };

    if (LoadLibraryExW(path, null, 0) == null) {
        showMsg(std.unicode.utf8ToUtf16LeStringLiteral("Injecting autopunch failed."));
    }
}

// EXPORTS
export fn CheckVersion(hash: *const [16]win.BYTE) bool {
    return std.mem.eql(win.BYTE, &TARGET_HASH, hash);
}

export fn Initialize(hSelf: win.HMODULE, _: win.HMODULE) bool {
    self_module = @as(?win.HMODULE, @ptrCast(hSelf));

    const thread = std.Thread.spawn(.{}, loadThread, .{}) catch return false;
    thread.detach();
    return true;
}

// DLL ENTRY
fn DllMain(
    hinst: win.HINSTANCE,
    reason: u32,
    _: ?*anyopaque,
) callconv(WINAPI) bool {
    const DLL_PROCESS_ATTACH: win.DWORD = 0x1;
    if (reason == DLL_PROCESS_ATTACH) {
        self_module = @as(?win.HMODULE, @ptrCast(hinst));
    }
    return true;
}
