667 lines
16 KiB
C++
667 lines
16 KiB
C++
#include "client_pipe.h"
|
|
|
|
|
|
static LONG g_started = 0;
|
|
static HANDLE g_stop_event = nullptr;
|
|
static HANDLE g_queue_event = nullptr;
|
|
static HANDLE g_connected_event = nullptr;
|
|
static HANDLE g_worker_thread = nullptr;
|
|
static HANDLE g_pipe = INVALID_HANDLE_VALUE;
|
|
static std::atomic<DWORD> g_client_id{ 0 };
|
|
static std::atomic<uint64_t> g_session_id{ 0 };
|
|
static std::atomic<uint64_t> g_next_request_id{ 1 };
|
|
static std::mutex g_queue_mutex;
|
|
static std::deque<std::shared_ptr<PendingCommand>> g_queue;
|
|
|
|
|
|
static void LogA(const char* fmt, ...)
|
|
{
|
|
char buf[512]{};
|
|
va_list ap;
|
|
va_start(ap, fmt);
|
|
vsnprintf_s(buf, sizeof(buf), _TRUNCATE, fmt, ap);
|
|
va_end(ap);
|
|
OutputDebugStringA(buf);
|
|
}
|
|
|
|
static void ExitInjectedProcess(UINT exit_code)
|
|
{
|
|
LogA("[dll-client] ExitProcess code=%u\n", exit_code);
|
|
ExitProcess(exit_code);
|
|
}
|
|
|
|
static bool ReadExact(HANDLE pipe, void* buffer, DWORD size)
|
|
{
|
|
BYTE* ptr = static_cast<BYTE*>(buffer);
|
|
DWORD total = 0;
|
|
while (total < size)
|
|
{
|
|
DWORD got = 0;
|
|
if (!ReadFile(pipe, ptr + total, size - total, &got, nullptr))
|
|
return false;
|
|
if (got == 0)
|
|
return false;
|
|
total += got;
|
|
}
|
|
return true;
|
|
}
|
|
|
|
static bool WriteExact(HANDLE pipe, const void* buffer, DWORD size)
|
|
{
|
|
const BYTE* ptr = static_cast<const BYTE*>(buffer);
|
|
DWORD total = 0;
|
|
while (total < size)
|
|
{
|
|
DWORD wrote = 0;
|
|
if (!WriteFile(pipe, ptr + total, size - total, &wrote, nullptr))
|
|
return false;
|
|
if (wrote == 0)
|
|
return false;
|
|
total += wrote;
|
|
}
|
|
return true;
|
|
}
|
|
|
|
|
|
bool WideToUtf8(const std::wstring& input, char* output, size_t output_size)
|
|
{
|
|
if (!output || output_size == 0)
|
|
return false;
|
|
|
|
output[0] = '\0';
|
|
if (input.empty())
|
|
return true;
|
|
|
|
int written = WideCharToMultiByte(
|
|
CP_UTF8,
|
|
0,
|
|
input.c_str(),
|
|
static_cast<int>(input.size()),
|
|
output,
|
|
static_cast<int>(output_size - 1),
|
|
nullptr,
|
|
nullptr);
|
|
|
|
if (written <= 0)
|
|
return false;
|
|
|
|
output[written] = '\0';
|
|
return true;
|
|
}
|
|
|
|
static void QueuePush(const std::shared_ptr<PendingCommand>& cmd)
|
|
{
|
|
std::lock_guard<std::mutex> lock(g_queue_mutex);
|
|
g_queue.push_back(cmd);
|
|
SetEvent(g_queue_event);
|
|
}
|
|
|
|
static std::shared_ptr<PendingCommand> QueuePop()
|
|
{
|
|
std::lock_guard<std::mutex> lock(g_queue_mutex);
|
|
if (g_queue.empty())
|
|
{
|
|
ResetEvent(g_queue_event);
|
|
return {};
|
|
}
|
|
|
|
auto cmd = g_queue.front();
|
|
g_queue.pop_front();
|
|
if (g_queue.empty())
|
|
ResetEvent(g_queue_event);
|
|
return cmd;
|
|
}
|
|
|
|
static void FailPendingCommands(DWORD status, const char* reply)
|
|
{
|
|
(void)reply;
|
|
|
|
for (;;)
|
|
{
|
|
auto cmd = QueuePop();
|
|
if (!cmd)
|
|
break;
|
|
|
|
cmd->status = status;
|
|
if (cmd->done_event)
|
|
SetEvent(cmd->done_event);
|
|
}
|
|
}
|
|
|
|
static void RequestStopNoWait()
|
|
{
|
|
if (g_stop_event)
|
|
SetEvent(g_stop_event);
|
|
}
|
|
|
|
static void CleanupClientStateNoWait()
|
|
{
|
|
RequestStopNoWait();
|
|
|
|
if (g_worker_thread)
|
|
{
|
|
CloseHandle(g_worker_thread);
|
|
g_worker_thread = nullptr;
|
|
}
|
|
|
|
if (g_stop_event)
|
|
{
|
|
CloseHandle(g_stop_event);
|
|
g_stop_event = nullptr;
|
|
}
|
|
|
|
if (g_queue_event)
|
|
{
|
|
CloseHandle(g_queue_event);
|
|
g_queue_event = nullptr;
|
|
}
|
|
|
|
if (g_connected_event)
|
|
{
|
|
CloseHandle(g_connected_event);
|
|
g_connected_event = nullptr;
|
|
}
|
|
}
|
|
|
|
static bool ConnectToServer(HANDLE& pipe)
|
|
{
|
|
for (;;)
|
|
{
|
|
if (WaitForSingleObject(g_stop_event, 0) == WAIT_OBJECT_0)
|
|
return false;
|
|
|
|
pipe = CreateFileW(
|
|
kPipeName,
|
|
GENERIC_READ | GENERIC_WRITE,
|
|
0,
|
|
nullptr,
|
|
OPEN_EXISTING,
|
|
0,
|
|
nullptr);
|
|
|
|
if (pipe != INVALID_HANDLE_VALUE)
|
|
{
|
|
DWORD mode = PIPE_READMODE_MESSAGE;
|
|
SetNamedPipeHandleState(pipe, &mode, nullptr, nullptr);
|
|
return true;
|
|
}
|
|
|
|
DWORD err = GetLastError();
|
|
if (err == ERROR_PIPE_BUSY)
|
|
{
|
|
if (!WaitNamedPipeW(kPipeName, 1000))
|
|
Sleep(250);
|
|
continue;
|
|
}
|
|
|
|
Sleep(500);
|
|
}
|
|
}
|
|
|
|
static bool SendHello(HANDLE pipe)
|
|
{
|
|
Msg hello{};
|
|
hello.type = static_cast<uint32_t>(MsgType::Hello);
|
|
hello.pid = GetCurrentProcessId();
|
|
hello.version = kProtocolVersion;
|
|
hello.tick = static_cast<uint64_t>(NowMs());
|
|
hello.session_id = 0;
|
|
|
|
if (!WriteExact(pipe, &hello, sizeof(hello)))
|
|
{
|
|
LogA("[dll-client] hello write failed: %lu\n", GetLastError());
|
|
return false;
|
|
}
|
|
|
|
Msg ack{};
|
|
if (!ReadExact(pipe, &ack, sizeof(ack)) || ack.type != static_cast<uint32_t>(MsgType::HelloAck))
|
|
{
|
|
LogA("[dll-client] hello ack failed: %lu\n", GetLastError());
|
|
return false;
|
|
}
|
|
|
|
g_client_id.store(ack.client_id);
|
|
g_session_id.store(ack.session_id);
|
|
SetEvent(g_connected_event);
|
|
|
|
LogA("[dll-client] connected client_id=%lu session=%llu\n",
|
|
g_client_id.load(),
|
|
static_cast<unsigned long long>(g_session_id.load()));
|
|
return true;
|
|
}
|
|
|
|
static bool SendHeartbeat(HANDLE pipe)
|
|
{
|
|
Msg heartbeat{};
|
|
heartbeat.type = static_cast<uint32_t>(MsgType::Heartbeat);
|
|
heartbeat.client_id = g_client_id.load();
|
|
heartbeat.pid = GetCurrentProcessId();
|
|
heartbeat.version = kProtocolVersion;
|
|
heartbeat.session_id = g_session_id.load();
|
|
heartbeat.tick = static_cast<uint64_t>(NowMs());
|
|
|
|
if (!WriteExact(pipe, &heartbeat, sizeof(heartbeat)))
|
|
{
|
|
LogA("[dll-client] heartbeat write failed: %lu\n", GetLastError());
|
|
return false;
|
|
}
|
|
|
|
return true;
|
|
}
|
|
|
|
static bool PumpServerMessages(HANDLE pipe)
|
|
{
|
|
for (;;)
|
|
{
|
|
DWORD available = 0;
|
|
if (!PeekNamedPipe(pipe, nullptr, 0, nullptr, &available, nullptr))
|
|
{
|
|
LogA("[dll-client] server disconnected: %lu\n", GetLastError());
|
|
ExitInjectedProcess(0);
|
|
return false;
|
|
}
|
|
|
|
if (available < sizeof(Msg))
|
|
return true;
|
|
|
|
Msg msg{};
|
|
if (!ReadExact(pipe, &msg, sizeof(msg)))
|
|
return false;
|
|
|
|
if (msg.type == static_cast<uint32_t>(MsgType::Exit))
|
|
{
|
|
LogA("[dll-client] exit requested by server\n");
|
|
SetEvent(g_stop_event);
|
|
ExitInjectedProcess(0);
|
|
return false;
|
|
}
|
|
}
|
|
}
|
|
|
|
static bool SendCommand(HANDLE pipe, PendingCommand& cmd)
|
|
{
|
|
Msg req{};
|
|
req.type = static_cast<uint32_t>(MsgType::CommandRequest);
|
|
req.client_id = g_client_id.load();
|
|
req.pid = GetCurrentProcessId();
|
|
req.version = kProtocolVersion;
|
|
req.command = cmd.command;
|
|
req.status = 0;
|
|
req.input_size = static_cast<uint32_t>(cmd.input.size());
|
|
req.output_size = static_cast<uint32_t>(cmd.output.size());
|
|
req.tick = static_cast<uint64_t>(NowMs());
|
|
req.session_id = g_session_id.load();
|
|
req.request_id = cmd.request_id;
|
|
|
|
if (cmd.input.size() > kPayloadBytes || cmd.output.size() > kPayloadBytes)
|
|
return false;
|
|
|
|
if (!cmd.input.empty())
|
|
memcpy(req.payload, cmd.input.data(), cmd.input.size());
|
|
|
|
if (!WriteExact(pipe, &req, sizeof(req)))
|
|
return false;
|
|
|
|
Msg resp{};
|
|
if (!ReadExact(pipe, &resp, sizeof(resp)))
|
|
return false;
|
|
|
|
if (resp.type != static_cast<uint32_t>(MsgType::CommandReply))
|
|
return false;
|
|
|
|
if (resp.request_id != cmd.request_id || resp.command != cmd.command)
|
|
return false;
|
|
|
|
cmd.status = resp.status;
|
|
if (resp.output_size && resp.output_size <= cmd.output.size())
|
|
memcpy(cmd.output.data(), resp.payload, resp.output_size);
|
|
|
|
return true;
|
|
}
|
|
|
|
static DWORD WINAPI WorkerThreadProc(LPVOID)
|
|
{
|
|
for (;;)
|
|
{
|
|
if (WaitForSingleObject(g_stop_event, 0) == WAIT_OBJECT_0)
|
|
break;
|
|
|
|
HANDLE pipe = INVALID_HANDLE_VALUE;
|
|
if (!ConnectToServer(pipe))
|
|
break;
|
|
|
|
g_pipe = pipe;
|
|
ResetEvent(g_connected_event);
|
|
|
|
if (!SendHello(pipe))
|
|
{
|
|
CloseHandle(pipe);
|
|
g_pipe = INVALID_HANDLE_VALUE;
|
|
Sleep(500);
|
|
continue;
|
|
}
|
|
|
|
DWORD last_heartbeat = static_cast<DWORD>(NowMs());
|
|
|
|
for (;;)
|
|
{
|
|
if (WaitForSingleObject(g_stop_event, 0) == WAIT_OBJECT_0)
|
|
break;
|
|
|
|
if (WaitForSingleObject(g_queue_event, 0) == WAIT_OBJECT_0)
|
|
{
|
|
for (;;)
|
|
{
|
|
auto cmd = QueuePop();
|
|
if (!cmd)
|
|
break;
|
|
|
|
if (!SendCommand(pipe, *cmd))
|
|
{
|
|
cmd->status = GetLastError();
|
|
}
|
|
|
|
if (cmd->done_event)
|
|
SetEvent(cmd->done_event);
|
|
|
|
if (WaitForSingleObject(g_stop_event, 0) == WAIT_OBJECT_0)
|
|
break;
|
|
}
|
|
}
|
|
|
|
const DWORD now = static_cast<DWORD>(NowMs());
|
|
if (now - last_heartbeat >= kHeartbeatIntervalMs)
|
|
{
|
|
if (!SendHeartbeat(pipe))
|
|
{
|
|
LogA("[dll-client] server heartbeat failed, exit process\n");
|
|
ExitInjectedProcess(0);
|
|
break;
|
|
}
|
|
|
|
last_heartbeat = now;
|
|
}
|
|
|
|
if (!PumpServerMessages(pipe))
|
|
break;
|
|
|
|
Sleep(10);
|
|
}
|
|
|
|
ResetEvent(g_connected_event);
|
|
|
|
if (pipe != INVALID_HANDLE_VALUE)
|
|
CloseHandle(pipe);
|
|
|
|
g_pipe = INVALID_HANDLE_VALUE;
|
|
|
|
if (WaitForSingleObject(g_stop_event, 0) == WAIT_OBJECT_0)
|
|
break;
|
|
|
|
Sleep(250);
|
|
}
|
|
|
|
FailPendingCommands(ERROR_CANCELLED, "client stopped");
|
|
return 0;
|
|
}
|
|
|
|
BOOL DllClientStart()
|
|
{
|
|
if (InterlockedCompareExchange(&g_started, 1, 0) != 0)
|
|
return TRUE;
|
|
|
|
g_stop_event = CreateEventW(nullptr, TRUE, FALSE, nullptr);
|
|
g_queue_event = CreateEventW(nullptr, TRUE, FALSE, nullptr);
|
|
g_connected_event = CreateEventW(nullptr, TRUE, FALSE, nullptr);
|
|
if (!g_stop_event || !g_queue_event || !g_connected_event)
|
|
{
|
|
CleanupClientStateNoWait();
|
|
return FALSE;
|
|
}
|
|
|
|
g_worker_thread = CreateThread(nullptr, 0, WorkerThreadProc, nullptr, 0, nullptr);
|
|
if (!g_worker_thread)
|
|
{
|
|
CleanupClientStateNoWait();
|
|
return FALSE;
|
|
}
|
|
|
|
return TRUE;
|
|
}
|
|
|
|
BOOL HBClientEnsureStarted()
|
|
{
|
|
return DllClientStart();
|
|
}
|
|
|
|
BOOL HBClientCall(
|
|
uint32_t command,
|
|
const void* input,
|
|
uint32_t input_size,
|
|
void* output,
|
|
uint32_t output_size,
|
|
uint32_t* bytes_returned)
|
|
{
|
|
if (bytes_returned)
|
|
*bytes_returned = 0;
|
|
|
|
if (input_size > kPayloadBytes || output_size > kPayloadBytes)
|
|
return FALSE;
|
|
|
|
if (input_size && !input)
|
|
return FALSE;
|
|
|
|
if (output_size && !output)
|
|
return FALSE;
|
|
|
|
if (InterlockedCompareExchange(&g_started, 0, 0) == 0)
|
|
{
|
|
if (!DllClientStart())
|
|
return FALSE;
|
|
}
|
|
|
|
auto cmd = std::make_shared<PendingCommand>();
|
|
cmd->command = command;
|
|
cmd->request_id = g_next_request_id.fetch_add(1);
|
|
cmd->done_event = CreateEventW(nullptr, TRUE, FALSE, nullptr);
|
|
if (!cmd->done_event)
|
|
return FALSE;
|
|
|
|
if (input_size)
|
|
{
|
|
const auto* ptr = static_cast<const uint8_t*>(input);
|
|
cmd->input.assign(ptr, ptr + input_size);
|
|
}
|
|
|
|
if (output_size)
|
|
cmd->output.resize(output_size);
|
|
|
|
QueuePush(cmd);
|
|
|
|
HANDLE waits[2] = { cmd->done_event, g_stop_event };
|
|
DWORD wait = WaitForMultipleObjects(2, waits, FALSE, 10000);
|
|
BOOL ok = (wait == WAIT_OBJECT_0 && cmd->status == ERROR_SUCCESS);
|
|
|
|
if (output && !cmd->output.empty())
|
|
{
|
|
memcpy(output, cmd->output.data(), cmd->output.size());
|
|
|
|
if (bytes_returned)
|
|
*bytes_returned =
|
|
static_cast<uint32_t>(cmd->output.size());
|
|
}
|
|
|
|
CloseHandle(cmd->done_event);
|
|
return ok;
|
|
}
|
|
|
|
BOOL HBClientSendPrintfW(const wchar_t* text)
|
|
{
|
|
return DllClientPrintfW(text);
|
|
}
|
|
|
|
static bool PipeReadMsgTimeout(HANDLE pipe, Msg& msg, DWORD timeout_ms)
|
|
{
|
|
const DWORD begin = GetTickCount();
|
|
|
|
for (;;)
|
|
{
|
|
DWORD available = 0;
|
|
if (!PeekNamedPipe(pipe, nullptr, 0, nullptr, &available, nullptr))
|
|
return false;
|
|
|
|
if (available >= sizeof(Msg))
|
|
{
|
|
DWORD got = 0;
|
|
return ReadFile(pipe, &msg, sizeof(msg), &got, nullptr) && got == sizeof(msg);
|
|
}
|
|
|
|
if (GetTickCount() - begin >= timeout_ms)
|
|
return false;
|
|
|
|
Sleep(10);
|
|
}
|
|
}
|
|
|
|
static bool PipeWriteMsg(HANDLE pipe, const Msg& msg)
|
|
{
|
|
DWORD wrote = 0;
|
|
return WriteFile(pipe, &msg, sizeof(msg), &wrote, nullptr) && wrote == sizeof(msg);
|
|
}
|
|
|
|
BOOL QueryServerAlreadyOpen()
|
|
{
|
|
HANDLE pipe = CreateFileW(
|
|
kPipeName,
|
|
GENERIC_READ | GENERIC_WRITE,
|
|
0,
|
|
nullptr,
|
|
OPEN_EXISTING,
|
|
0,
|
|
nullptr);
|
|
|
|
if (pipe == INVALID_HANDLE_VALUE)
|
|
return FALSE;
|
|
|
|
DWORD mode = PIPE_READMODE_MESSAGE;
|
|
SetNamedPipeHandleState(pipe, &mode, nullptr, nullptr);
|
|
|
|
Msg hello{};
|
|
hello.type = static_cast<uint32_t>(MsgType::Hello);
|
|
hello.pid = GetCurrentProcessId();
|
|
hello.version = kProtocolVersion;
|
|
hello.tick = static_cast<uint64_t>(NowMs());
|
|
|
|
Msg ack{};
|
|
if (!PipeWriteMsg(pipe, hello) ||
|
|
!PipeReadMsgTimeout(pipe, ack, 500) ||
|
|
ack.type != static_cast<uint32_t>(MsgType::HelloAck))
|
|
{
|
|
CloseHandle(pipe);
|
|
return FALSE;
|
|
}
|
|
|
|
Msg req{};
|
|
req.type = static_cast<uint32_t>(MsgType::CommandRequest);
|
|
req.client_id = ack.client_id;
|
|
req.pid = GetCurrentProcessId();
|
|
req.version = kProtocolVersion;
|
|
req.command = static_cast<uint32_t>(CommandId::QueryStatus);
|
|
req.session_id = ack.session_id;
|
|
req.request_id = 1;
|
|
req.tick = static_cast<uint64_t>(NowMs());
|
|
|
|
Msg resp{};
|
|
BOOL already_open =
|
|
PipeWriteMsg(pipe, req) &&
|
|
PipeReadMsgTimeout(pipe, resp, 500) &&
|
|
resp.type == static_cast<uint32_t>(MsgType::CommandReply) &&
|
|
resp.command == static_cast<uint32_t>(CommandId::QueryStatus);
|
|
|
|
Msg bye{};
|
|
bye.type = static_cast<uint32_t>(MsgType::Goodbye);
|
|
bye.client_id = ack.client_id;
|
|
bye.pid = GetCurrentProcessId();
|
|
bye.version = kProtocolVersion;
|
|
bye.session_id = ack.session_id;
|
|
PipeWriteMsg(pipe, bye);
|
|
|
|
CloseHandle(pipe);
|
|
return already_open;
|
|
}
|
|
|
|
|
|
BOOL HBClientLoadDriverW(PDWORD status)
|
|
{
|
|
|
|
LoadDriverRequest req{};
|
|
LoadDriverRequest response{};
|
|
|
|
uint32_t bytes_returned = 0;
|
|
|
|
BOOL result = HBClientCall(
|
|
static_cast<uint32_t>(CommandId::LoadDriver),
|
|
&req,
|
|
sizeof(req),
|
|
&response,
|
|
sizeof(response),
|
|
&bytes_returned);
|
|
|
|
if (bytes_returned == sizeof(response))
|
|
*status = response.status;
|
|
|
|
return result;
|
|
}
|
|
|
|
BOOL HBClientUnloadDriverW(const wchar_t* service_name)
|
|
{
|
|
if (!service_name)
|
|
return FALSE;
|
|
|
|
UnloadDriverRequest req{};
|
|
wcsncpy_s(req.service_name, service_name, _TRUNCATE);
|
|
|
|
return HBClientCall(
|
|
static_cast<uint32_t>(CommandId::UnloadDriver),
|
|
&req,
|
|
sizeof(req),
|
|
nullptr,
|
|
0,
|
|
nullptr);
|
|
}
|
|
|
|
BOOL HBClientNotifyBreakpointRemoved(uint32_t thread_id, uint64_t dr0, uint64_t dr7)
|
|
{
|
|
BreakpointRemovedRequest req{};
|
|
req.thread_id = thread_id;
|
|
req.dr0 = dr0;
|
|
req.dr7 = dr7;
|
|
|
|
return HBClientCall(
|
|
static_cast<uint32_t>(CommandId::BreakpointRemoved),
|
|
&req,
|
|
sizeof(req),
|
|
nullptr,
|
|
0,
|
|
nullptr);
|
|
}
|
|
|
|
BOOL DllClientPrintfW(const wchar_t* text)
|
|
{
|
|
if (!text)
|
|
return FALSE;
|
|
|
|
PrintfRequest req{};
|
|
if (!WideToUtf8(text, req.text, sizeof(req.text)))
|
|
return FALSE;
|
|
|
|
return HBClientCall(
|
|
static_cast<uint32_t>(CommandId::Printf),
|
|
&req,
|
|
sizeof(req),
|
|
nullptr,
|
|
0,
|
|
nullptr);
|
|
}
|
|
|