using System;
using System.Collections.Generic;
using System.Linq;
using System.Runtime.InteropServices;
using WhiteMagic.Native;
namespace WhiteMagic.Thread;
///
/// Enumerates and selects threads belonging to the target process.
///
public sealed class ThreadFactory
{
private readonly MemoryBase _memory;
/// Creates a factory bound to the target process represented by .
public ThreadFactory(MemoryBase memory)
{
_memory = memory ?? throw new ArgumentNullException(nameof(memory));
}
///
/// Enumerates every thread that belongs to the target process.
///
public IEnumerable Enumerate()
{
foreach (int threadId in CollectThreadIds())
{
SafeMemoryHandle handle = NativeMethods.OpenThread(
ThreadAccess.SuspendResume |
ThreadAccess.GetContext |
ThreadAccess.SetContext |
ThreadAccess.QueryInformation,
false,
threadId);
if (handle.IsInvalid)
continue;
yield return new RemoteThread(_memory, threadId, handle);
}
}
private int[] CollectThreadIds()
{
using SafeMemoryHandle snapshot = NativeMethods.CreateToolhelp32Snapshot(SnapshotFlags.Thread, 0);
if (snapshot.IsInvalid)
{
int error = Marshal.GetLastPInvokeError();
throw new InvalidOperationException($"CreateToolhelp32Snapshot failed: error {error}.");
}
var entry = new ThreadEntry32
{
dwSize = (uint)Marshal.SizeOf()
};
var ids = new List();
if (!NativeMethods.Thread32First(snapshot, ref entry))
{
int error = Marshal.GetLastPInvokeError();
if (error == 18 || error == 259) // ERROR_NO_MORE_FILES / ERROR_NO_MORE_ITEMS
return ids.ToArray();
throw new InvalidOperationException($"Thread32First failed: error {error}.");
}
do
{
if (entry.th32OwnerProcessID == (uint)_memory.ProcessId)
ids.Add((int)entry.th32ThreadID);
}
while (NativeMethods.Thread32Next(snapshot, ref entry));
return ids.ToArray();
}
///
/// Returns the thread with the specified operating-system identifier if it belongs
/// to the target process.
///
/// The thread does not belong to the target process.
public RemoteThread GetThreadById(int threadId)
{
if (threadId <= 0)
throw new ArgumentException("Thread ID must be positive.", nameof(threadId));
SafeMemoryHandle handle = NativeMethods.OpenThread(ThreadAccess.QueryInformation, false, threadId);
if (handle.IsInvalid)
{
int error = Marshal.GetLastPInvokeError();
throw new InvalidOperationException($"OpenThread failed for thread {threadId}: error {error}.");
}
try
{
var info = new ThreadBasicInformation();
int status = NativeMethods.NtQueryInformationThread(
handle,
0,
ref info,
(uint)Marshal.SizeOf(),
out _);
if (status < 0)
{
throw new InvalidOperationException(
$"NtQueryInformationThread failed for thread {threadId} (NTSTATUS {status:X8}).");
}
if ((uint)(nint)info.ClientId.UniqueProcess != (uint)_memory.ProcessId)
{
throw new InvalidOperationException(
$"Thread {threadId} does not belong to process {_memory.ProcessId}.");
}
// Open a handle with the rights the public RemoteThread surface needs.
return new RemoteThread(_memory, threadId);
}
finally
{
handle.Dispose();
}
}
///
/// Returns the earliest-created thread of the target process.
///
public RemoteThread MainThread
{
get
{
RemoteThread? earliest = null;
long earliestTime = long.MaxValue;
foreach (RemoteThread thread in Enumerate())
{
long creationTime = GetCreationTime(thread.Id);
if (creationTime < earliestTime)
{
earliestTime = creationTime;
earliest?.Dispose();
earliest = thread;
}
else
{
thread.Dispose();
}
}
if (earliest is null)
{
throw new InvalidOperationException(
$"Process {_memory.ProcessId} has no observable threads.");
}
return earliest;
}
}
///
/// Suspends the supplied threads and returns a disposable scope that resumes exactly
/// those threads when disposed, including when an exception escapes the guarded body.
///
///
/// Do not freeze the target's threads while executing target code through a remote
/// thread or main-thread pump; doing so can deadlock because the frozen thread is the
/// one responsible for running the code.
///
public FrozenThread Freeze(IEnumerable threads)
{
ArgumentNullException.ThrowIfNull(threads);
var suspended = new List();
try
{
foreach (RemoteThread thread in threads)
{
thread.Suspend();
suspended.Add(thread);
}
return new FrozenThread(suspended);
}
catch
{
foreach (RemoteThread thread in suspended)
{
try
{
thread.Resume();
}
catch
{
// Best-effort unwind.
}
}
throw;
}
}
///
/// Suspends all threads selected by .
///
public FrozenThread Freeze(Func predicate)
{
ArgumentNullException.ThrowIfNull(predicate);
return Freeze(Enumerate().Where(predicate));
}
private long GetCreationTime(int threadId)
{
using SafeMemoryHandle handle = NativeMethods.OpenThread(ThreadAccess.QueryInformation, false, threadId);
if (handle.IsInvalid)
{
int error = Marshal.GetLastPInvokeError();
throw new InvalidOperationException($"OpenThread failed for thread {threadId}: error {error}.");
}
if (!NativeMethods.GetThreadTimes(handle, out long creationTime, out _, out _, out _))
{
int error = Marshal.GetLastPInvokeError();
throw new InvalidOperationException($"GetThreadTimes failed for thread {threadId}: error {error}.");
}
return creationTime;
}
}