Add thread control surfaces
Adds RemoteThread, ThreadFactory (enumeration, main-thread selection, get-by-id), and FrozenThread scoped freeze. Supports suspend/resume, 32/64-bit context round-trip, TEB query, and reverse-order resume on dispose. Closes section 2 of add-thread-region-finder.
This commit is contained in:
@@ -0,0 +1,230 @@
|
||||
using System;
|
||||
using System.Collections.Generic;
|
||||
using System.Linq;
|
||||
using System.Runtime.InteropServices;
|
||||
using WhiteMagic.Native;
|
||||
|
||||
namespace WhiteMagic.Thread;
|
||||
|
||||
/// <summary>
|
||||
/// Enumerates and selects threads belonging to the target process.
|
||||
/// </summary>
|
||||
public sealed class ThreadFactory
|
||||
{
|
||||
private readonly MemoryBase _memory;
|
||||
|
||||
/// <summary>Creates a factory bound to the target process represented by <paramref name="memory"/>.</summary>
|
||||
public ThreadFactory(MemoryBase memory)
|
||||
{
|
||||
_memory = memory ?? throw new ArgumentNullException(nameof(memory));
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Enumerates every thread that belongs to the target process.
|
||||
/// </summary>
|
||||
public IEnumerable<RemoteThread> 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<ThreadEntry32>()
|
||||
};
|
||||
|
||||
var ids = new List<int>();
|
||||
|
||||
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();
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Returns the thread with the specified operating-system identifier if it belongs
|
||||
/// to the target process.
|
||||
/// </summary>
|
||||
/// <exception cref="InvalidOperationException">The thread does not belong to the target process.</exception>
|
||||
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<ThreadBasicInformation>(),
|
||||
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();
|
||||
}
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Returns the earliest-created thread of the target process.
|
||||
/// </summary>
|
||||
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;
|
||||
}
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Suspends the supplied threads and returns a disposable scope that resumes exactly
|
||||
/// those threads when disposed, including when an exception escapes the guarded body.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// 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.
|
||||
/// </remarks>
|
||||
public FrozenThread Freeze(IEnumerable<RemoteThread> threads)
|
||||
{
|
||||
ArgumentNullException.ThrowIfNull(threads);
|
||||
|
||||
var suspended = new List<RemoteThread>();
|
||||
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;
|
||||
}
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Suspends all threads selected by <paramref name="predicate"/>.
|
||||
/// </summary>
|
||||
public FrozenThread Freeze(Func<RemoteThread, bool> 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;
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user