diff --git a/src/TUnit.Core/ObjectInitializer.cs b/src/TUnit.Core/ObjectInitializer.cs index 0acc3a6afe..7b2ffe777f 100644 --- a/src/TUnit.Core/ObjectInitializer.cs +++ b/src/TUnit.Core/ObjectInitializer.cs @@ -20,10 +20,12 @@ namespace TUnit.Core; /// internal static class ObjectInitializer { - // Use Lazy pattern to ensure InitializeAsync is called exactly once per object, - // even under contention. GetOrAdd's factory can be called multiple times, but with - // Lazy + ExecutionAndPublication mode, only one initialization actually runs. - private static readonly ConcurrentDictionary> 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 + 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 InitializationTasks = new(Helpers.ReferenceEqualityComparer.Instance); /// @@ -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, 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; } /// @@ -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 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( - 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(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 completionSource) + { + try { - // Propagate cancellation without modification - throw; + await asyncInitializer.InitializeAsync().ConfigureAwait(false); } - catch + catch (Exception ex) { - // Do NOT remove from cache - the faulted Lazy 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); } } diff --git a/src/TUnit.Engine/Services/ObjectLifecycleService.cs b/src/TUnit.Engine/Services/ObjectLifecycleService.cs index 98e5ecd8b7..2775b95da8 100644 --- a/src/TUnit.Engine/Services/ObjectLifecycleService.cs +++ b/src/TUnit.Engine/Services/ObjectLifecycleService.cs @@ -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. + // 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> _spannedObjects = new(); #endif diff --git a/tests/TUnit.UnitTests/ObjectInitializerTests.cs b/tests/TUnit.UnitTests/ObjectInitializerTests.cs new file mode 100644 index 0000000000..f80cf3c0e7 --- /dev/null +++ b/tests/TUnit.UnitTests/ObjectInitializerTests.cs @@ -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(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(); + await Assert.That(waiter.IsCompleted).IsFalse(); + + if (initializationFails) + { + fixture.Fail(failure); + var observed = await Assert.That(async () => await waiter.WaitAsync(HangTimeout)).Throws(); + 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(); + 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[] 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 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 CaptureAsync(Task task) + { + try + { + await task.WaitAsync(HangTimeout); + return null; + } + catch (Exception ex) + { + return ex; + } + } + + /// + /// Blocks inside the synchronous part of InitializeAsync (before its first await), like + /// sync-over-async code in a third-party constructor would. + /// + private sealed class BlockingPrefixInitializer(bool completeSynchronously = false) : IAsyncInitializer, IDisposable + { + private readonly ManualResetEventSlim _prefixGate = new(); + private int _initializeCount; + + public TaskCompletionSource 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(); + } + + /// + /// Suspends until or is called. + /// + private sealed class GatedInitializer : IAsyncInitializer + { + private readonly TaskCompletionSource _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; + } + } +}