Skip to content
72 changes: 42 additions & 30 deletions src/TUnit.Core/ObjectInitializer.cs
Original file line number Diff line number Diff line change
Expand Up @@ -20,10 +20,12 @@ namespace TUnit.Core;
/// </remarks>
internal static class ObjectInitializer
{
// Use Lazy<Task> pattern to ensure InitializeAsync is called exactly once per object,
// even under contention. GetOrAdd's factory can be called multiple times, but with
// Lazy<Task> + ExecutionAndPublication mode, only one initialization actually runs.
private static readonly ConcurrentDictionary<object, Lazy<Task>> InitializationTasks =
// One task per object, published before any user code runs, so InitializeAsync is called
// exactly once per object even under contention. No lock is held while InitializeAsync runs:
// Lazy<Task> + ExecutionAndPublication ran its synchronous part under a lock, so every other
// caller blocked a thread-pool thread until it finished, starving the pool when that part
// blocked on async work itself (#6904).
private static readonly ConcurrentDictionary<object, Task> InitializationTasks =
new(Helpers.ReferenceEqualityComparer.Instance);

/// <summary>
Expand Down Expand Up @@ -88,12 +90,10 @@ internal static bool IsInitialized(object? obj)
return false;
}

// Use Status == RanToCompletion to ensure we don't return true for faulted/canceled tasks
// Use Status == RanToCompletion to ensure we don't return true for pending or failed initializations
// (IsCompletedSuccessfully is not available in netstandard2.0)
// With Lazy<Task>, we need to check if the Lazy has a value AND that value completed successfully
return InitializationTasks.TryGetValue(obj, out var lazyTask) &&
lazyTask.IsValueCreated &&
lazyTask.Value.Status == TaskStatus.RanToCompletion;
return InitializationTasks.TryGetValue(obj, out var initializationTask) &&
initializationTask.Status == TaskStatus.RanToCompletion;
}

/// <summary>
Expand All @@ -107,37 +107,49 @@ internal static void ClearCache()
InitializationTasks.Clear();
}

// Kept async (rather than returning the WaitAsync task as a ValueTask) so that an
// OperationCanceledException thrown by InitializeAsync still completes callers' tasks as
// Canceled, as it did before, instead of Faulted.
private static async ValueTask InitializeCoreAsync(
object obj,
IAsyncInitializer asyncInitializer,
CancellationToken cancellationToken)
{
// Use Lazy<Task> with ExecutionAndPublication mode to ensure InitializeAsync
// is called exactly once, even under contention. GetOrAdd's factory may be
// called multiple times, but Lazy ensures only one initialization runs.
var lazyTask = InitializationTasks.GetOrAdd(obj,
static (_, asyncInitializer) => new Lazy<Task>(
asyncInitializer.InitializeAsync,
LazyThreadSafetyMode.ExecutionAndPublication)
, asyncInitializer);

try
if (!InitializationTasks.TryGetValue(obj, out var initializationTask))
{
// Wait for initialization with cancellation support
await lazyTask.Value.WaitAsync(cancellationToken);
// Waiting tests must not run inline on the thread that completes initialization.
var completionSource = new TaskCompletionSource<bool>(TaskCreationOptions.RunContinuationsAsynchronously);
initializationTask = InitializationTasks.GetOrAdd(obj, completionSource.Task);

if (ReferenceEquals(initializationTask, completionSource.Task))
{
// Only the caller that published the task runs InitializeAsync - inline up to its
// first await, as before, but with no lock held (#6904).
_ = RunInitializerAsync(asyncInitializer, completionSource);
}
}
catch (OperationCanceledException)

// Do NOT remove faulted tasks from the cache - subsequent callers get the same error
// immediately. Removing and retrying can cause hangs when InitializeAsync partially
// initialized resources (e.g. started ports/processes) that block re-initialization (#4715).
// The cancellation token only stops this caller waiting; the initialization keeps running.
await initializationTask.WaitAsync(cancellationToken);
}

private static async Task RunInitializerAsync(IAsyncInitializer asyncInitializer, TaskCompletionSource<bool> completionSource)
{
try
{
// Propagate cancellation without modification
throw;
await asyncInitializer.InitializeAsync().ConfigureAwait(false);
Comment thread
Sing303 marked this conversation as resolved.
Comment thread
coderabbitai[bot] marked this conversation as resolved.
}
catch
catch (Exception ex)
{
// Do NOT remove from cache - the faulted Lazy<Task> stays so subsequent
// callers get the same error immediately via .WaitAsync() on the faulted task.
// Removing and retrying can cause hangs when InitializeAsync partially initialized
// resources (e.g. started ports/processes) that block re-initialization (#4715).
throw;
// SetException rather than SetCanceled, so callers get the original exception object -
// including an OperationCanceledException thrown by InitializeAsync.
completionSource.SetException(ex);
return;
}

completionSource.SetResult(true);
}
}
2 changes: 1 addition & 1 deletion src/TUnit.Engine/Services/ObjectLifecycleService.cs
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,7 @@ internal sealed class ObjectLifecycleService : IObjectRegistry, IInitializationC
#if NET
// Gates span creation so only the first caller for a given object creates a trace span.
// Subsequent callers (concurrent tests sharing the same object) skip span creation
// and just await ObjectInitializer's deduplicated Lazy<Task>.
// and just await ObjectInitializer's deduplicated initialization task.
// Uses ConditionalWeakTable so per-test objects can be GC'd after their test completes.
private readonly ConditionalWeakTable<object, StrongBox<int>> _spannedObjects = new();
#endif
Expand Down
285 changes: 285 additions & 0 deletions tests/TUnit.UnitTests/ObjectInitializerTests.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,285 @@
using TUnit.Core;
using TUnit.Core.Interfaces;

namespace TUnit.UnitTests;

public class ObjectInitializerTests
{
// Only bound how long a regression can hang the suite - passing runs never wait for them.
// The prefix gate must outlast HangTimeout, or a blocked caller would be released before it is detected.
private static readonly TimeSpan HangTimeout = TimeSpan.FromSeconds(10);
private static readonly TimeSpan PrefixGateTimeout = TimeSpan.FromSeconds(60);

// https://github.com/thomhurst/TUnit/issues/6904
[Test]
public async Task Waiting_Caller_Does_Not_Block_While_InitializeAsync_Runs_Synchronously()
{
using var fixture = new BlockingPrefixInitializer();
var initialization = Task.Run(() => ObjectInitializer.InitializeAsync(fixture).AsTask());
var callReturned = new TaskCompletionSource<bool>(TaskCreationOptions.RunContinuationsAsynchronously);
Task waiter = Task.CompletedTask;
bool callReturnedWhilePrefixBlocked;
bool completedWhilePrefixBlocked;

try
{
await fixture.PrefixEntered.Task.WaitAsync(HangTimeout);

waiter = Task.Run(() =>
{
var pending = ObjectInitializer.InitializeAsync(fixture);
callReturned.SetResult(true);
return pending.AsTask();
});

callReturnedWhilePrefixBlocked = await Task.WhenAny(callReturned.Task, Task.Delay(HangTimeout)) == callReturned.Task;
completedWhilePrefixBlocked = waiter.IsCompleted;
}
finally
{
fixture.ReleasePrefix();
}

await Task.WhenAll(initialization, waiter).WaitAsync(HangTimeout);

await Assert.That(callReturnedWhilePrefixBlocked).IsTrue();
await Assert.That(completedWhilePrefixBlocked).IsFalse();
await Assert.That(fixture.InitializeCount).IsEqualTo(1);
await Assert.That(ObjectInitializer.IsInitialized(fixture)).IsTrue();
}

[Test]
public async Task InitializeAsync_Runs_Once_Under_Contention()
{
var fixtures = Enumerable.Range(0, 100).Select(_ => new YieldingInitializer()).ToArray();

var callers = fixtures
.SelectMany(fixture => Enumerable.Range(0, 8).Select(_ => Task.Run(() => ObjectInitializer.InitializeAsync(fixture).AsTask())))
.ToArray();
await Task.WhenAll(callers).WaitAsync(HangTimeout);

foreach (var fixture in fixtures)
{
await Assert.That(fixture.InitializeCount).IsEqualTo(1);
await Assert.That(ObjectInitializer.IsInitialized(fixture)).IsTrue();
}
}

[Test]
[Arguments(false, true)]
[Arguments(false, false)]
[Arguments(true, true)]
[Arguments(true, false)]
public async Task Failure_Is_Cached_And_Rethrown_As_The_Same_Exception(bool cancellation, bool throwSynchronously)
{
Exception failure = cancellation ? new OperationCanceledException("initializer gave up") : new InvalidOperationException("initialization failed");
var fixture = new ThrowingInitializer(failure, throwSynchronously);

var first = ObjectInitializer.InitializeAsync(fixture).AsTask();
var second = ObjectInitializer.InitializeAsync(fixture).AsTask();

await Assert.That(await CaptureAsync(first)).IsSameReferenceAs(failure);
await Assert.That(await CaptureAsync(second)).IsSameReferenceAs(failure);
// An OperationCanceledException from the initializer still cancels callers' tasks, as before.
await Assert.That(first.IsCanceled).IsEqualTo(cancellation);
await Assert.That(second.IsCanceled).IsEqualTo(cancellation);
await Assert.That(fixture.InitializeCount).IsEqualTo(1);
await Assert.That(ObjectInitializer.IsInitialized(fixture)).IsFalse();
}

[Test]
[Arguments(false)]
[Arguments(true)]
public async Task Cancelling_The_Initializing_Caller_Does_Not_Poison_The_Result(bool initializationFails)
{
var fixture = new GatedInitializer();
var failure = new InvalidOperationException("initialization failed");
using var cancellationTokenSource = new CancellationTokenSource();
var initializingCaller = ObjectInitializer.InitializeAsync(fixture, cancellationTokenSource.Token).AsTask();
var waiter = ObjectInitializer.InitializeAsync(fixture).AsTask();

cancellationTokenSource.Cancel();

await Assert.That(async () => await initializingCaller.WaitAsync(HangTimeout)).Throws<OperationCanceledException>();
await Assert.That(waiter.IsCompleted).IsFalse();

if (initializationFails)
{
fixture.Fail(failure);
var observed = await Assert.That(async () => await waiter.WaitAsync(HangTimeout)).Throws<InvalidOperationException>();
await Assert.That(observed).IsSameReferenceAs(failure);
}
else
{
fixture.Complete();
await waiter.WaitAsync(HangTimeout);
}

await Assert.That(fixture.InitializeCount).IsEqualTo(1);
await Assert.That(ObjectInitializer.IsInitialized(fixture)).IsEqualTo(!initializationFails);
}

[Test]
public async Task Cancelling_A_Waiting_Caller_Does_Not_Cancel_The_Initialization()
{
var fixture = new GatedInitializer();
using var cancellationTokenSource = new CancellationTokenSource();
var initializingCaller = ObjectInitializer.InitializeAsync(fixture).AsTask();
var cancelledWaiter = ObjectInitializer.InitializeAsync(fixture, cancellationTokenSource.Token).AsTask();

cancellationTokenSource.Cancel();

await Assert.That(async () => await cancelledWaiter.WaitAsync(HangTimeout)).Throws<OperationCanceledException>();
await Assert.That(ObjectInitializer.IsInitialized(fixture)).IsFalse();

fixture.Complete();
await initializingCaller.WaitAsync(HangTimeout);
await ObjectInitializer.InitializeAsync(fixture);

await Assert.That(fixture.InitializeCount).IsEqualTo(1);
await Assert.That(ObjectInitializer.IsInitialized(fixture)).IsTrue();
}

[Test]
[Arguments(false)]
[Arguments(true)]
public async Task Waiting_Continuations_Do_Not_Run_On_The_Initializing_Thread(bool cancellableWait)
{
// Arrange
const int waiterCount = 4;
using var fixture = new BlockingPrefixInitializer(completeSynchronously: true);
using var cancellationTokenSource = new CancellationTokenSource();
var cancellationToken = cancellableWait ? cancellationTokenSource.Token : CancellationToken.None;
var initializingThreadId = 0;
Task<int>[] waiters = [];

// A dedicated thread cannot later pick up correctly queued waiter continuations.
var initialization = Task.Factory.StartNew(() =>
{
initializingThreadId = Environment.CurrentManagedThreadId;
return ObjectInitializer.InitializeAsync(fixture).AsTask();
}, CancellationToken.None, TaskCreationOptions.LongRunning, TaskScheduler.Default).Unwrap();

async Task<int> ObserveContinuationAsync()
{
await ObjectInitializer.InitializeAsync(fixture, cancellationToken).ConfigureAwait(false);
return Environment.CurrentManagedThreadId;
}

// Act
try
{
await fixture.PrefixEntered.Task.WaitAsync(HangTimeout);
waiters = Enumerable.Range(0, waiterCount).Select(_ => ObserveContinuationAsync()).ToArray();
}
finally
{
fixture.ReleasePrefix();
await initialization.WaitAsync(HangTimeout);
}

var continuationThreads = await Task.WhenAll(waiters).WaitAsync(HangTimeout);

// Assert
await Assert.That(continuationThreads.Length).IsEqualTo(waiterCount);
foreach (var threadId in continuationThreads)
{
await Assert.That(threadId).IsNotEqualTo(initializingThreadId);
}
}

private static async Task<Exception?> CaptureAsync(Task task)
{
try
{
await task.WaitAsync(HangTimeout);
return null;
}
catch (Exception ex)
{
return ex;
}
}

/// <summary>
/// Blocks inside the synchronous part of InitializeAsync (before its first await), like
/// sync-over-async code in a third-party constructor would.
/// </summary>
private sealed class BlockingPrefixInitializer(bool completeSynchronously = false) : IAsyncInitializer, IDisposable
{
private readonly ManualResetEventSlim _prefixGate = new();
private int _initializeCount;

public TaskCompletionSource<bool> PrefixEntered { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously);

public int InitializeCount => Volatile.Read(ref _initializeCount);

public void ReleasePrefix() => _prefixGate.Set();

public async Task InitializeAsync()
{
Interlocked.Increment(ref _initializeCount);
PrefixEntered.TrySetResult(true);
_prefixGate.Wait(PrefixGateTimeout);
if (!completeSynchronously)
{
await Task.Yield();
}
}

public void Dispose() => _prefixGate.Dispose();
}

/// <summary>
/// Suspends until <see cref="Complete"/> or <see cref="Fail"/> is called.
/// </summary>
private sealed class GatedInitializer : IAsyncInitializer
{
private readonly TaskCompletionSource<bool> _gate = new();
private int _initializeCount;

public int InitializeCount => Volatile.Read(ref _initializeCount);

public void Complete() => _gate.SetResult(true);

public void Fail(Exception exception) => _gate.SetException(exception);

public async Task InitializeAsync()
{
Interlocked.Increment(ref _initializeCount);
await _gate.Task;
}
}

private sealed class YieldingInitializer : IAsyncInitializer
{
private int _initializeCount;

public int InitializeCount => Volatile.Read(ref _initializeCount);

public async Task InitializeAsync()
{
Interlocked.Increment(ref _initializeCount);
await Task.Yield();
}
}

private sealed class ThrowingInitializer(Exception exception, bool throwSynchronously) : IAsyncInitializer
{
private int _initializeCount;

public int InitializeCount => Volatile.Read(ref _initializeCount);

public Task InitializeAsync()
{
Interlocked.Increment(ref _initializeCount);
return throwSynchronously ? throw exception : ThrowAfterYieldAsync();
}

private async Task ThrowAfterYieldAsync()
{
await Task.Yield();
throw exception;
}
}
}
Loading