Merge caller-sta-pump: single-thread pump architecture

This commit is contained in:
nbaertsch
2026-06-10 17:06:04 -04:00
+158 -330
View File
@@ -101,6 +101,36 @@ fn resolve(dll: [*:0]const u8, name: [*:0]const u8) usize {
return if (proc) |p| @intFromPtr(p) else 0;
}
fn releaseComObject(obj: ?*anyopaque) void {
if (obj) |o| {
const obj_ptr: *const usize = @ptrCast(@alignCast(o));
const vtbl_addr = obj_ptr.*;
const release_fn: *const fn (*anyopaque) callconv(.c) u32 =
@ptrFromInt(@as(*const usize, @ptrFromInt(vtbl_addr + 16)).*);
_ = release_fn(o);
}
}
inline fn getTEB() usize {
return asm volatile ("mov %gs:0x30, %[ret]"
: [ret] "=r" (-> usize),
);
}
fn waitNoPump(handle: ?*anyopaque) void {
const NtWaitForSingleObject: *const fn (?*anyopaque, u8, ?*const i64) callconv(.c) i32 =
@ptrFromInt(resolve("ntdll.dll", "NtWaitForSingleObject"));
_ = NtWaitForSingleObject(handle, 0, null);
}
fn waitComAware(handle: ?*anyopaque) HRESULT {
const CoWaitForMultipleHandles: *const fn (u32, u32, u32, [*]const ?*anyopaque, *u32) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoWaitForMultipleHandles"));
const handles = [_]?*anyopaque{handle};
var index: u32 = 0;
return CoWaitForMultipleHandles(0, INFINITE, 1, &handles, &index);
}
// ============================================================================
// MIDL / RPC structures
// ============================================================================
@@ -780,6 +810,10 @@ fn rpb_release(_: **const FakeRPBVtbl) callconv(.c) u32 {
}
fn rpb_connect(_: **const FakeRPBVtbl, pChannel: *anyopaque) callconv(.c) HRESULT {
log("[RPB] Connect called! pChannel=0x{X}\n", .{@intFromPtr(pChannel)});
if (g_ctx.proxy_channel) |old_channel| {
if (old_channel == pChannel) return S_OK;
releaseComObject(old_channel);
}
g_ctx.proxy_channel = pChannel;
// AddRef the channel
const obj_ptr: *const usize = @ptrCast(@alignCast(pChannel));
@@ -942,142 +976,77 @@ fn initSyncVtbl() void {
}
// ============================================================================
// Layer 5: STA command loop
// Layer 5: Caller-STA helpers
// ============================================================================
const StaCommand = enum(u32) { none = 0, sleep = 1, proxy = 2, shutdown = 3 };
const UnmarshalCtx = struct {
stream: ?*anyopaque,
hr: HRESULT = E_FAIL,
done_event: ?*anyopaque,
};
fn staThread(_: ?*anyopaque) callconv(.c) u32 {
fn unmarshalProxyThread(param: ?*anyopaque) callconv(.c) u32 {
const ctx: *UnmarshalCtx = @ptrCast(@alignCast(param));
const CoInitializeEx: *const fn (?*anyopaque, u32) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoInitializeEx"));
_ = CoInitializeEx(null, COINIT_APARTMENTTHREADED);
const CoRegisterPSClsid: *const fn (*const [16]u8, *const [16]u8) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoRegisterPSClsid"));
const CoRegisterClassObject: *const fn (*const [16]u8, *anyopaque, u32, u32, *u32) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoRegisterClassObject"));
const CoGetInterfaceAndReleaseStream: *const fn (*anyopaque, *const [16]u8, **anyopaque) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoGetInterfaceAndReleaseStream"));
const hr_coinit = CoInitializeEx(null, COINIT_MULTITHREADED);
if (hr_coinit < 0) {
ctx.hr = hr_coinit;
if (ctx.done_event) |evt| _ = SetEvent(evt);
return 1;
}
// The helper MTA thread must register the PS mapping/class object in its
// own apartment so CoGetInterfaceAndReleaseStream can resolve the proxy.
_ = CoRegisterPSClsid(&IID_IFiberDispatch, &CLSID_FiberDispatchPS);
var cookie: u32 = 0;
_ = CoRegisterClassObject(&CLSID_FiberDispatchPS, @ptrCast(&g_psf), CLSCTX_INPROC_SERVER, REGCLS_MULTIPLEUSE, &cookie);
// Marshal interface for cross-apartment unmarshal
var proxy: ?*anyopaque = null;
ctx.hr = CoGetInterfaceAndReleaseStream(ctx.stream.?, &IID_IFiberDispatch, @ptrCast(&proxy));
if (ctx.done_event) |evt| _ = SetEvent(evt);
return if (ctx.hr < 0) 1 else 0;
}
fn refreshProxyChannel() bool {
const CoMarshalInterThreadInterfaceInStream: *const fn (*const [16]u8, *anyopaque, **anyopaque) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoMarshalInterThreadInterfaceInStream"));
var stream: ?*anyopaque = null;
const marshal_hr = CoMarshalInterThreadInterfaceInStream(&IID_IFiberDispatch, @ptrCast(&g_fco), @ptrCast(&stream));
if (marshal_hr < 0) {
_ = SetEvent(g_ctx.sta_ready);
return 1;
}
g_ctx.marshal_stream = stream;
if (marshal_hr < 0) return false;
// Signal ready FIRST — let main thread unmarshal
_ = SetEvent(g_ctx.sta_ready);
const done_event = CreateEventA(null, 1, 0, null) orelse return false;
defer _ = CloseHandle(done_event);
// Wait for unmarshal to complete before converting to fiber
_ = WaitForSingleObject(g_ctx.unmarshal_done, INFINITE);
var ctx = UnmarshalCtx{
.stream = stream,
.done_event = done_event,
};
// NOW convert to fiber (after unmarshal is complete)
const main_fiber = ConvertThreadToFiber(null);
if (main_fiber == null) {
return 1;
}
g_ctx.sta_main_fiber = main_fiber;
// Signal init() that the fiber is ready — resolves race between
// init() returning and pump() reading sta_main_fiber
_ = SetEvent(g_ctx.pump_ready_event);
// Command loop — use NtWaitForSingleObject to prevent COM from pumping
// PostCall messages during the idle wait. This is safe because COM
// initialization is complete by this point (all init messages were
// pumped during the COM-aware WaitForSingleObject in setup).
// Messages MUST only be dispatched inside ModalLoop, not during idle waits.
const NtWaitForSingleObject: *const fn (?*anyopaque, u8, ?*const i64) callconv(.c) i32 =
@ptrFromInt(resolve("ntdll.dll", "NtWaitForSingleObject"));
while (true) {
_ = NtWaitForSingleObject(g_ctx.command_event, 0, null);
_ = ResetEvent(g_ctx.command_event);
const cmd = @as(StaCommand, @enumFromInt(@as(u32, @atomicLoad(u32, &g_ctx.current_command, .seq_cst))));
switch (cmd) {
.sleep => {
// Inline the pump setup here (like omega1 does — all in one stack frame)
const addrs = g_ctx.resolved orelse break;
const NoOpReturn0_addr = addrs.NoOpReturn0;
const CCliModalLoopCtor_addr = addrs.CCliModalLoopCtor;
var inner_vtable: [24]usize = undefined;
for (&inner_vtable) |*v| v.* = NoOpReturn0_addr;
var inner_obj: [0x170]u8 align(8) = [_]u8{0} ** 0x170;
@as(*usize, @ptrCast(@alignCast(&inner_obj[0]))).* = @intFromPtr(&inner_vtable);
var cancel_vtable: [16]usize = undefined;
for (&cancel_vtable) |*v| v.* = NoOpReturn0_addr;
var cancel_obj: [8]u8 align(8) = undefined;
@as(*usize, @ptrCast(@alignCast(&cancel_obj[0]))).* = @intFromPtr(&cancel_vtable);
@as(*usize, @ptrCast(@alignCast(&inner_obj[0x160]))).* = @intFromPtr(&cancel_obj);
const pump_event = CreateEventA(null, 0, 0, null);
if (pump_event == null) break;
var fake_client_call: [0x120]u8 align(8) = [_]u8{0} ** 0x120;
@as(*usize, @ptrCast(@alignCast(&fake_client_call[0xC0]))).* = @intFromPtr(&inner_obj);
@as(*usize, @ptrCast(@alignCast(&fake_client_call[0x108]))).* = @intFromPtr(pump_event);
// Construct CML
var cml_buf: [0x200]u8 align(8) = [_]u8{0} ** 0x200;
const CCliModalLoopCtor2: *const fn (*[0x200]u8, u32, u32, u32, i32) callconv(.c) void =
@ptrFromInt(CCliModalLoopCtor_addr);
CCliModalLoopCtor2(&cml_buf, 0, 0x04FF, 0, 0);
// Refresh stub
if (g_ctx.tables) |*tables| tables.refreshStub();
// Create pump fiber with ModalLoop as DIRECT fiber proc — zero user code on stack
const ModalLoop_addr = addrs.ModalLoop;
const pump_fiber = CreateFiber(0x100000, @ptrFromInt(ModalLoop_addr), @ptrCast(&fake_client_call));
if (pump_fiber == null) {
_ = CloseHandle(pump_event);
break;
}
log("[STA] Switching to pump fiber...\n", .{});
SwitchToFiber(pump_fiber);
// Last dispatched call was SwitchToFiber(sta_main_fiber) → returns here
log("[STA] Returned from pump!\n", .{});
DeleteFiber(pump_fiber);
_ = CloseHandle(pump_event);
log("[STA] Pump cycle complete.\n", .{});
_ = SetEvent(g_ctx.completion_event);
},
.proxy => {
executePumpCycle(null, null, null);
_ = SetEvent(g_ctx.completion_event);
},
.shutdown => break,
.none => {},
}
}
_ = ConvertFiberToThread();
return 0;
var helper_tid: u32 = 0;
const helper_handle = CreateThread(null, 0, &unmarshalProxyThread, @ptrCast(&ctx), 0, &helper_tid) orelse return false;
defer _ = CloseHandle(helper_handle);
_ = waitComAware(done_event);
waitNoPump(helper_handle);
waitNoPump(helper_handle);
return ctx.hr >= 0 and g_ctx.proxy_channel != null;
}
/// Pump cycle for proxy mode: constructs CML + ModalLoop fiber, dispatches queued messages.
fn executePumpCycle(_: ?*anyopaque, _: ?*anyopaque, _: ?*anyopaque) void {
fn executePumpCycle() bool {
log("[STA] executePumpCycle starting ({d} calls)...\n", .{g_ctx.queue_count});
const addrs = g_ctx.resolved orelse return;
const addrs = g_ctx.resolved orelse return false;
const NoOpReturn0_addr = addrs.NoOpReturn0;
const CCliModalLoopCtor_addr = addrs.CCliModalLoopCtor;
// Build fake inner object (all vtable slots → NoOpReturn0)
var inner_vtable: [24]usize = undefined;
for (&inner_vtable) |*v| v.* = NoOpReturn0_addr;
@@ -1089,40 +1058,27 @@ fn executePumpCycle(_: ?*anyopaque, _: ?*anyopaque, _: ?*anyopaque) void {
@as(*usize, @ptrCast(@alignCast(&cancel_obj[0]))).* = @intFromPtr(&cancel_vtable);
@as(*usize, @ptrCast(@alignCast(&inner_obj[0x160]))).* = @intFromPtr(&cancel_obj);
// Pump event (auto-reset)
const pump_event = CreateEventA(null, 0, 0, null);
if (pump_event == null) return;
const pump_event = CreateEventA(null, 0, 0, null) orelse return false;
defer _ = CloseHandle(pump_event);
// Fake CSyncClientCall
var fake_client_call: [0x120]u8 align(8) = [_]u8{0} ** 0x120;
@as(*usize, @ptrCast(@alignCast(&fake_client_call[0xC0]))).* = @intFromPtr(&inner_obj);
@as(*usize, @ptrCast(@alignCast(&fake_client_call[0x108]))).* = @intFromPtr(pump_event);
// Construct CML
var cml_buf: [0x200]u8 align(8) = [_]u8{0} ** 0x200;
const CCliModalLoopCtor: *const fn (*[0x200]u8, u32, u32, u32, i32) callconv(.c) void =
@ptrFromInt(CCliModalLoopCtor_addr);
CCliModalLoopCtor(&cml_buf, 0, 0x04FF, 0, 0);
CCliModalLoopCtor(&cml_buf, 0, 0x04FF, 0, 1);
// Refresh stub before pump
if (g_ctx.tables) |*tables| tables.refreshStub();
// Create pump fiber with ModalLoop as DIRECT fiber proc (like omega1)
const ModalLoop_addr = addrs.ModalLoop;
const pump_fiber = CreateFiber(0x100000, @ptrFromInt(ModalLoop_addr), @ptrCast(&fake_client_call));
if (pump_fiber == null) {
_ = CloseHandle(pump_event);
return;
}
const pump_fiber = CreateFiber(0x100000, @ptrFromInt(addrs.ModalLoop), @ptrCast(&fake_client_call)) orelse return false;
defer DeleteFiber(pump_fiber);
log("[STA] Switching to pump fiber (direct ModalLoop)...\n", .{});
SwitchToFiber(pump_fiber);
// SwitchToFiber(main) in the dispatch chain returns us here
log("[STA] Returned from pump!\n", .{});
DeleteFiber(pump_fiber);
_ = CloseHandle(pump_event);
log("[STA] Pump cycle complete.\n", .{});
return true;
}
// Debug wrapper for SwitchToFiber — intercepts the dispatch to verify
@@ -1148,7 +1104,6 @@ const SleepWorkerCtx = struct {
post_call_addr: usize,
calls: [*]QueuedCall,
num_calls: u32,
go_event: ?*anyopaque,
done_event: ?*anyopaque,
};
@@ -1158,15 +1113,10 @@ fn sleepWorkerThread(param: ?*anyopaque) callconv(.c) u32 {
const CoInitializeEx: *const fn (?*anyopaque, u32) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoInitializeEx"));
_ = CoInitializeEx(null, COINIT_MULTITHREADED);
_ = WaitForSingleObject(ctx.go_event, INFINITE);
log("[MTA-sleep] Worker started, posting {d} calls\n", .{ctx.num_calls});
// Read SOleTlsData for PostCall PID store
const teb = asm ("mov %gs:0x30, %[ret]"
: [ret] "=r" (-> usize),
);
const ole_tls_ptr = @as(*const usize, @ptrFromInt(teb + 0x1758)).*;
const ole_tls_ptr = @as(*const usize, @ptrFromInt(getTEB() + 0x1758)).*;
const ch_vtbl_ptr: *const usize = @ptrCast(@alignCast(ctx.channel));
const vtbl_base: usize = ch_vtbl_ptr.*;
@@ -1256,7 +1206,6 @@ const ProxyWorkerCtx = struct {
switch_call: *QueuedCall,
// Output
return_value: u32,
go_event: ?*anyopaque,
done_event: ?*anyopaque,
};
@@ -1266,8 +1215,6 @@ fn proxyWorkerThread(param: ?*anyopaque) callconv(.c) u32 {
const CoInitializeEx: *const fn (?*anyopaque, u32) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoInitializeEx"));
_ = CoInitializeEx(null, COINIT_MULTITHREADED);
_ = WaitForSingleObject(ctx.go_event, INFINITE);
log("[MTA-proxy] Worker started\n", .{});
const ch_vtbl_ptr: *const usize = @ptrCast(@alignCast(ctx.channel));
@@ -1309,10 +1256,7 @@ fn proxyWorkerThread(param: ?*anyopaque) callconv(.c) u32 {
// Phase 2: PostCall SwitchToFiber(main) to return pump control
log("[MTA-proxy] Phase 2: PostCall SwitchToFiber...\n", .{});
const teb = asm ("mov %gs:0x30, %[ret]"
: [ret] "=r" (-> usize),
);
const ole_tls_ptr = @as(*const usize, @ptrFromInt(teb + 0x1758)).*;
const ole_tls_ptr = @as(*const usize, @ptrFromInt(getTEB() + 0x1758)).*;
var switch_msg = std.mem.zeroes(RPCOLEMESSAGE);
switch_msg.iMethod = ctx.switch_call.method;
@@ -1385,14 +1329,8 @@ const ContextInternal = struct {
resolved: ?ResolvedAddrs = null,
proxy_channel: ?*anyopaque = null,
marshal_stream: ?*anyopaque = null,
sta_ready: ?*anyopaque = null,
sta_main_fiber: ?*anyopaque = null,
command_event: ?*anyopaque = null,
completion_event: ?*anyopaque = null,
unmarshal_done: ?*anyopaque = null,
pump_ready_event: ?*anyopaque = null, // signaled after ConvertThreadToFiber — init() waits on this
current_command: u32 = 0,
sta_handle: ?*anyopaque = null,
caller_fiber: ?*anyopaque = null,
mta_usage_cookie: ?*anyopaque = null,
// Queue for sleep mode
queue_buf: [256]QueuedCall = undefined,
@@ -1430,15 +1368,23 @@ pub const Context = struct {
_ = self;
if (g_ctx.state != .queuing) return;
g_ctx.state = .executing;
defer {
g_ctx.queue_count = 0;
g_ctx.state = .ready;
}
// Auto-append SwitchToFiber(main_fiber)
if (!refreshProxyChannel()) return;
const caller_fiber = g_ctx.caller_fiber orelse return;
// Auto-append SwitchToFiber(caller_fiber)
const switch_slot: u16 = RESERVED_METHODS + @as(u16, @truncate(g_ctx.queue_count));
const fn_switchfiber = if (DEBUG)
@intFromPtr(&debugSwitchToFiber)
else
resolve("kernel32.dll", "SwitchToFiber");
var switch_call = QueuedCall{ .method = switch_slot, .n_params = 1, .params = [_]usize{0} ** MAX_PARAMS };
switch_call.params[0] = @intFromPtr(g_ctx.sta_main_fiber);
switch_call.params[0] = @intFromPtr(caller_fiber);
if (g_ctx.tables) |*tables| tables.writeSlot(switch_slot, fn_switchfiber, 1);
g_ctx.queue_buf[g_ctx.queue_count] = switch_call;
g_ctx.queue_count += 1;
@@ -1450,32 +1396,18 @@ pub const Context = struct {
.post_call_addr = addrs2.PostCall,
.calls = &g_ctx.queue_buf,
.num_calls = g_ctx.queue_count,
.go_event = CreateEventA(null, 0, 0, null),
.done_event = CreateEventA(null, 1, 0, null),
};
if (worker_ctx.done_event == null) return;
defer _ = CloseHandle(worker_ctx.done_event);
var worker_tid: u32 = 0;
const worker_handle = CreateThread(null, 0, &sleepWorkerThread, @ptrCast(&worker_ctx), 0, &worker_tid);
_ = SetEvent(worker_ctx.go_event);
const worker_handle = CreateThread(null, 0, &sleepWorkerThread, @ptrCast(&worker_ctx), 0, &worker_tid) orelse return;
defer _ = CloseHandle(worker_handle);
// Wait for worker to finish posting ALL messages, then let it exit
_ = WaitForSingleObject(worker_ctx.done_event, INFINITE);
_ = WaitForSingleObject(worker_handle, 5000);
// Signal STA to pump
@atomicStore(u32, &g_ctx.current_command, @intFromEnum(StaCommand.sleep), .seq_cst);
_ = SetEvent(g_ctx.command_event);
// Wait for pump completion
_ = WaitForSingleObject(g_ctx.completion_event, INFINITE);
_ = ResetEvent(g_ctx.completion_event);
// Cleanup
_ = CloseHandle(worker_ctx.go_event);
_ = CloseHandle(worker_ctx.done_event);
if (worker_handle) |h| _ = CloseHandle(h);
g_ctx.queue_count = 0;
g_ctx.state = .ready;
waitNoPump(worker_ctx.done_event);
if (!executePumpCycle()) return;
waitNoPump(worker_handle);
}
/// Single synchronous call with return value capture (proxy mode).
@@ -1486,10 +1418,12 @@ pub const Context = struct {
if (g_ctx.state != .ready) return 0;
if (params.len > MAX_PARAMS) return 0;
g_ctx.state = .invoking;
defer g_ctx.state = .ready;
log("[invoke] Setting up slots...\n", .{});
const real_slot: u16 = RESERVED_METHODS;
const switch_slot: u16 = RESERVED_METHODS + 1;
const caller_fiber = g_ctx.caller_fiber orelse return 0;
// Write real call into slot 7
var real_call = QueuedCall{ .method = real_slot, .n_params = @truncate(params.len), .params = [_]usize{0} ** MAX_PARAMS };
@@ -1499,13 +1433,10 @@ pub const Context = struct {
// Write SwitchToFiber(main) into slot 8
const fn_switchfiber = resolve("kernel32.dll", "SwitchToFiber");
var switch_call = QueuedCall{ .method = switch_slot, .n_params = 1, .params = [_]usize{0} ** MAX_PARAMS };
switch_call.params[0] = @intFromPtr(g_ctx.sta_main_fiber);
switch_call.params[0] = @intFromPtr(caller_fiber);
if (g_ctx.tables) |*tables| tables.writeSlot(switch_slot, fn_switchfiber, 1);
const addrs3 = g_ctx.resolved orelse {
g_ctx.state = .ready;
return 0;
};
const addrs3 = g_ctx.resolved orelse return 0;
var worker_ctx = ProxyWorkerCtx{
.channel = g_ctx.proxy_channel.?,
@@ -1513,37 +1444,20 @@ pub const Context = struct {
.real_call = &real_call,
.switch_call = &switch_call,
.return_value = 0,
.go_event = CreateEventA(null, 0, 0, null),
.done_event = CreateEventA(null, 0, 0, null),
};
if (worker_ctx.done_event == null) return 0;
defer _ = CloseHandle(worker_ctx.done_event);
var worker_tid: u32 = 0;
const worker_handle = CreateThread(null, 0, &proxyWorkerThread, @ptrCast(&worker_ctx), 0, &worker_tid);
const worker_handle = CreateThread(null, 0, &proxyWorkerThread, @ptrCast(&worker_ctx), 0, &worker_tid) orelse return 0;
defer _ = CloseHandle(worker_handle);
// Signal STA to start pumping FIRST — wait for ModalLoop to be
// actively running before letting the worker post messages.
// Signal STA to start pumping FIRST (so it's ready when SendReceive posts)
@atomicStore(u32, &g_ctx.current_command, @intFromEnum(StaCommand.proxy), .seq_cst);
_ = SetEvent(g_ctx.command_event);
// Signal worker to start
_ = SetEvent(worker_ctx.go_event);
// Wait for worker to complete both phases
_ = WaitForSingleObject(worker_ctx.done_event, INFINITE);
// Wait for STA pump to complete (SwitchToFiber(main) returns control)
_ = WaitForSingleObject(g_ctx.completion_event, INFINITE);
_ = ResetEvent(g_ctx.completion_event);
if (!executePumpCycle()) return 0;
waitNoPump(worker_ctx.done_event);
waitNoPump(worker_handle);
const retval = worker_ctx.return_value;
// Cleanup
_ = CloseHandle(worker_ctx.go_event);
_ = CloseHandle(worker_ctx.done_event);
if (worker_handle) |h| _ = CloseHandle(h);
g_ctx.state = .ready;
return retval;
}
@@ -1551,21 +1465,18 @@ pub const Context = struct {
_ = self;
if (g_ctx.state == .uninitialized) return;
// Signal STA shutdown
@atomicStore(u32, &g_ctx.current_command, @intFromEnum(StaCommand.shutdown), .seq_cst);
_ = SetEvent(g_ctx.command_event);
_ = WaitForSingleObject(g_ctx.sta_handle, 5000);
if (g_ctx.caller_fiber != null) _ = ConvertFiberToThread();
releaseComObject(g_ctx.proxy_channel);
g_ctx.proxy_channel = null;
// Free tables
if (g_ctx.tables) |*tables| tables.deinit();
// Close events
if (g_ctx.sta_ready) |h| _ = CloseHandle(h);
if (g_ctx.command_event) |h| _ = CloseHandle(h);
if (g_ctx.completion_event) |h| _ = CloseHandle(h);
if (g_ctx.unmarshal_done) |h| _ = CloseHandle(h);
if (g_ctx.pump_ready_event) |h| _ = CloseHandle(h);
if (g_ctx.sta_handle) |h| _ = CloseHandle(h);
g_ctx.tables = null;
if (g_ctx.mta_usage_cookie) |cookie| {
const CoDecrementMTAUsage: *const fn (?*anyopaque) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoDecrementMTAUsage"));
_ = CoDecrementMTAUsage(cookie);
g_ctx.mta_usage_cookie = null;
}
g_ctx.state = .uninitialized;
}
@@ -1728,43 +1639,36 @@ pub fn init(config: Config) !Context {
g_ctx.tables = MidlTables.init(heap, config.max_slots) orelse return error.MidlInitFailed;
g_ctx.max_slots = config.max_slots;
// Create synchronization events
g_ctx.sta_ready = CreateEventA(null, 1, 0, null); // manual-reset
g_ctx.command_event = CreateEventA(null, 1, 0, null); // manual-reset
g_ctx.completion_event = CreateEventA(null, 1, 0, null); // manual-reset
g_ctx.unmarshal_done = CreateEventA(null, 1, 0, null); // manual-reset
g_ctx.pump_ready_event = CreateEventA(null, 0, 0, null); // auto-reset
// Register COM on main thread (MTA side)
// Caller thread is now the STA.
const CoInitializeEx: *const fn (?*anyopaque, u32) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoInitializeEx"));
const CoRegisterPSClsid: *const fn (*const [16]u8, *const [16]u8) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoRegisterPSClsid"));
const CoRegisterClassObject: *const fn (*const [16]u8, *anyopaque, u32, u32, *u32) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoRegisterClassObject"));
const CoIncrementMTAUsage: *const fn (*?*anyopaque) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoIncrementMTAUsage"));
const hr_coinit = CoInitializeEx(null, COINIT_APARTMENTTHREADED);
log("[+] CoInitializeEx (caller STA): 0x{X}\n", .{@as(u32, @bitCast(hr_coinit))});
{
const PeekMessageA: *const fn (*[48]u8, ?*anyopaque, u32, u32, u32) callconv(.c) i32 =
@ptrFromInt(resolve("user32.dll", "PeekMessageA"));
const DispatchMessageA: *const fn (*const [48]u8) callconv(.c) isize =
@ptrFromInt(resolve("user32.dll", "DispatchMessageA"));
var msg: [48]u8 = [_]u8{0} ** 48;
while (PeekMessageA(&msg, null, 0, 0, 0x0001) != 0) {
_ = DispatchMessageA(&msg);
}
}
const hr_coinit = CoInitializeEx(null, COINIT_MULTITHREADED);
log("[+] CoInitializeEx (main): 0x{X}\n", .{@as(u32, @bitCast(hr_coinit))});
const hr_psclsid = CoRegisterPSClsid(&IID_IFiberDispatch, &CLSID_FiberDispatchPS);
log("[+] CoRegisterPSClsid (main): 0x{X}\n", .{@as(u32, @bitCast(hr_psclsid))});
log("[+] CoRegisterPSClsid (caller): 0x{X}\n", .{@as(u32, @bitCast(hr_psclsid))});
var cookie: u32 = 0;
const hr_regclass = CoRegisterClassObject(&CLSID_FiberDispatchPS, @ptrCast(&g_psf), CLSCTX_INPROC_SERVER, REGCLS_MULTIPLEUSE, &cookie);
log("[+] CoRegisterClassObject (main): 0x{X}\n", .{@as(u32, @bitCast(hr_regclass))});
log("[+] CoRegisterClassObject (caller): 0x{X}\n", .{@as(u32, @bitCast(hr_regclass))});
// Spawn STA thread
var sta_tid: u32 = 0;
g_ctx.sta_handle = CreateThread(null, 0, &staThread, null, 0, &sta_tid);
if (g_ctx.sta_handle == null) return error.StaThreadFailed;
// Wait for STA to be ready
_ = WaitForSingleObject(g_ctx.sta_ready, INFINITE);
log("[+] STA ready, stream=0x{X}\n", .{@intFromPtr(g_ctx.marshal_stream)});
// Unmarshal proxy on main thread
const CoGetInterfaceAndReleaseStream: *const fn (*anyopaque, *const [16]u8, **anyopaque) callconv(.c) HRESULT =
@ptrFromInt(resolve("ole32.dll", "CoGetInterfaceAndReleaseStream"));
var proxy: ?*anyopaque = null;
// VEH crash handler — only in debug builds
if (DEBUG) {
const AddVEH: *const fn (u32, *const fn (*anyopaque) callconv(.c) c_long) callconv(.c) ?*anyopaque =
@@ -1772,18 +1676,15 @@ pub fn init(config: Config) !Context {
_ = AddVEH(1, &crashHandler);
}
log("[+] Calling CoGetInterfaceAndReleaseStream...\n", .{});
var mta_cookie: ?*anyopaque = null;
const mta_hr = CoIncrementMTAUsage(&mta_cookie);
if (mta_hr < 0) return error.MtaUsageFailed;
g_ctx.mta_usage_cookie = mta_cookie;
const unmarshal_hr = CoGetInterfaceAndReleaseStream(g_ctx.marshal_stream.?, &IID_IFiberDispatch, @ptrCast(&proxy));
log("[+] Unmarshal returned: 0x{X}\n", .{@as(u32, @bitCast(unmarshal_hr))});
if (unmarshal_hr < 0) return error.UnmarshalFailed;
if (!refreshProxyChannel()) return error.UnmarshalFailed;
// Signal STA that unmarshal is done (so it can convert to fiber)
_ = SetEvent(g_ctx.unmarshal_done);
// Wait for STA to finish fiber conversion — without this, pump() can
// read sta_main_fiber as NULL if it runs before ConvertThreadToFiber completes
_ = WaitForSingleObject(g_ctx.pump_ready_event, INFINITE);
g_ctx.caller_fiber = ConvertThreadToFiber(null);
if (g_ctx.caller_fiber == null) return error.CallerFiberFailed;
g_ctx.state = .ready;
return Context{};
@@ -1798,7 +1699,6 @@ pub fn main() void {
// Parse --sleep-ms <N> from command line (default 2000ms)
var sleep_ms: u64 = 2000;
var msgwait_only = false;
const args = std.process.argsAlloc(std.heap.page_allocator) catch &[_][:0]const u8{};
var argi: usize = 1;
while (argi < args.len) : (argi += 1) {
@@ -1807,12 +1707,9 @@ pub fn main() void {
if (argi < args.len) {
sleep_ms = std.fmt.parseInt(u64, args[argi], 10) catch 2000;
}
} else if (std.mem.eql(u8, args[argi], "--msgwait-only")) {
msgwait_only = true;
}
}
print("[+] Sleep time: {d}ms per call\n", .{sleep_ms});
if (msgwait_only) print("[+] Mode: MsgWait-only (skipping NtDelay tests)\n", .{});
var ctx = init(.{}) catch |e| {
print("[-] Init failed: {}\n", .{e});
@@ -1823,88 +1720,19 @@ pub fn main() void {
print("[+] COMegon initialized\n", .{});
print("[+] State: ready\n", .{});
// Test 1: queue() + pump() — sleep mode (NtDelay)
const fn_ntdelay = resolve("ntdll.dll", "NtDelayExecution");
if (fn_ntdelay == 0) {
print("[-] Failed to resolve NtDelayExecution\n", .{});
return;
}
const heap = GetProcessHeap();
if (!msgwait_only) {
const delay_ptr = HeapAlloc(heap, HEAP_ZERO_MEMORY, 8) orelse return;
const delay_100ns: i64 = -@as(i64, @intCast(sleep_ms)) * 10000;
@as(*i64, @ptrCast(@alignCast(delay_ptr))).* = delay_100ns;
ctx.queue(fn_ntdelay, &[_]usize{ 0, @intFromPtr(delay_ptr) });
ctx.queue(fn_ntdelay, &[_]usize{ 0, @intFromPtr(delay_ptr) });
const start = GetTickCount64();
ctx.pump();
const elapsed = GetTickCount64() - start;
print("[+] Sleep mode complete! (total elapsed={d}ms)\n", .{elapsed});
_ = HeapFree(heap, 0, delay_ptr);
// Test 2: invoke() — proxy mode (NtDelay single call with result)
print("\n[*] Testing proxy mode: invoke(NtDelay {d}ms)...\n", .{sleep_ms});
const delay_ptr2 = HeapAlloc(heap, HEAP_ZERO_MEMORY, 8) orelse return;
@as(*i64, @ptrCast(@alignCast(delay_ptr2))).* = delay_100ns;
const start2 = GetTickCount64();
const result = ctx.invoke(fn_ntdelay, &[_]usize{ 0, @intFromPtr(delay_ptr2) });
const elapsed2 = GetTickCount64() - start2;
print("[+] Proxy mode result: 0x{X} (elapsed={d}ms)\n", .{ result, elapsed2 });
_ = HeapFree(heap, 0, delay_ptr2);
}
// Test 3: MsgWait sleep — dispatch MsgWaitForMultipleObjectsEx through COM
// This blocks for sleep_ms inside NdrStubCall2 using win32u syscall
// (NtUserMsgWaitForMultipleObjectsEx) instead of ntdll (NtDelayExecution).
// Stack during sleep: MsgWait ← Invoke_Epv ← NdrStubCall2 ← ModalLoop
// All system code, no user code. Sleep API is win32u, not ntdll.
print("\n[*] Testing MsgWait sleep: queue(MsgWaitForMultipleObjectsEx, {d}ms)...\n", .{sleep_ms});
const fn_msgwait = resolve("user32.dll", "MsgWaitForMultipleObjectsEx");
if (fn_msgwait == 0) {
print("[-] Failed to resolve MsgWaitForMultipleObjectsEx\n", .{});
} else {
// MsgWaitForMultipleObjectsEx(nCount=0, pHandles=NULL, dwMs=sleep_ms, dwWakeMask=0, dwFlags=0)
// With mask=0 and no handles, the ONLY wake condition is the timeout.
const sleep_ms_usize: usize = @intCast(sleep_ms);
ctx.queue(fn_msgwait, &[_]usize{ 0, 0, sleep_ms_usize, 0, 0 });
const start3 = GetTickCount64();
ctx.pump();
const elapsed3 = GetTickCount64() - start3;
print("[+] MsgWait sleep complete! (elapsed={d}ms, expected ~{d}ms)\n", .{ elapsed3, sleep_ms });
return;
}
// Test 4: Full encrypt-sleep-decrypt chain via MsgWait
// Demonstrates the real sleep mask pattern:
// 1. queue(encrypt_fn) — encrypts target memory
// 2. queue(MsgWaitForMultipleObjectsEx, sleep_ms) — timed wait, all system code
// 3. queue(decrypt_fn) — decrypts target memory
// 4. pump() — autonomous dispatch, no worker timing
print("\n[*] Testing encrypt-MsgWait-decrypt chain ({d}ms)...\n", .{sleep_ms});
if (fn_msgwait != 0) {
// Use NtDelayExecution as stand-in for encrypt/decrypt (proves the 3-phase chain works)
const tiny_delay_ptr = HeapAlloc(heap, HEAP_ZERO_MEMORY, 8) orelse return;
const tiny_100ns: i64 = -10000; // 1ms — just proves the dispatch fires
@as(*i64, @ptrCast(@alignCast(tiny_delay_ptr))).* = tiny_100ns;
const sleep_ms_usize: usize = @intCast(sleep_ms);
ctx.queue(fn_ntdelay, &[_]usize{ 0, @intFromPtr(tiny_delay_ptr) }); // "encrypt" (1ms)
ctx.queue(fn_msgwait, &[_]usize{ 0, 0, sleep_ms_usize, 0, 0 }); // sleep via MsgWait
ctx.queue(fn_ntdelay, &[_]usize{ 0, @intFromPtr(tiny_delay_ptr) }); // "decrypt" (1ms)
const start4 = GetTickCount64();
ctx.pump();
const elapsed4 = GetTickCount64() - start4;
print("[+] Chain complete! (elapsed={d}ms, expected ~{d}ms)\n", .{ elapsed4, sleep_ms + 2 });
_ = HeapFree(heap, 0, tiny_delay_ptr);
}
const sleep_ms_usize: usize = @intCast(sleep_ms);
print("[+] Queueing MsgWait + SwitchToFiber\n", .{});
ctx.queue(fn_msgwait, &[_]usize{ 0, 0, sleep_ms_usize, 0, 0 });
const start = GetTickCount64();
ctx.pump();
const elapsed = GetTickCount64() - start;
print("[+] Pump returned (elapsed={d}ms, expected ~{d}ms)\n", .{ elapsed, sleep_ms });
print("[+] COMegon test complete\n", .{});
}