diff --git a/src/libraries/System.Threading.Channels/src/System/Threading/Channels/AsyncOperation.cs b/src/libraries/System.Threading.Channels/src/System/Threading/Channels/AsyncOperation.cs index 4d2a3720605daf..3d08945b9c0f6e 100644 --- a/src/libraries/System.Threading.Channels/src/System/Threading/Channels/AsyncOperation.cs +++ b/src/libraries/System.Threading.Channels/src/System/Threading/Channels/AsyncOperation.cs @@ -124,19 +124,6 @@ public ValueTaskSourceStatus GetStatus(short token) } /// Gets whether the operation has completed. - /// - /// The operation is considered completed if both a) it's in the completed state, - /// AND b) it has a non-null continuation. We need to consider both because they're - /// not set atomically. If we only considered the state, then if we set the state to - /// completed and then set the continuation, it's possible for an awaiter to check - /// IsCompleted, see true, call GetResult, and return the object to the pool, and only - /// then do we try to store the continuation into an object we no longer own. If we - /// only considered the state, then if we set the continuation and then set the state, - /// a racing awaiter could see the continuation set before the state has transitioned - /// to completed and could end up calling GetResult in an incomplete state. And if we - /// only considered the continuation, then we have issues if OnCompleted is used before - /// the operation completes, as the continuation will be - /// internal bool IsCompleted => ReferenceEquals(_continuation, s_completedSentinel); /// Gets the result of the operation. @@ -301,7 +288,11 @@ public void OnCompleted(Action continuation, object? state, short token } } - /// Unregisters from cancellation. + /// Unregisters from cancellation and returns whether cancellation already started. + /// + /// true if either the instance wasn't cancelable or cancellation successfully unregistered without cancellation having started. + /// false if cancellation successfully unregistered after cancellation was initiated. + /// /// /// This is important for two reasons: /// 1. To avoid leaking a registration into a token, so it must be done prior to completing the operation. @@ -309,7 +300,17 @@ public void OnCompleted(Action continuation, object? state, short token /// that no one else will try to complete the operation (assuming the caller is properly constructed /// and themselves guarantees only a single completer other than through cancellation). /// - public void UnregisterCancellation() => _registration.Dispose(); + public bool UnregisterCancellation() + { + if (CancellationToken.CanBeCanceled) + { + _registration.Dispose(); // Dispose rather than Unregister is important to know work has quiesced + return _completionReserved == 0; + } + + Debug.Assert(_registration == default); + return true; + } /// Completes the operation with a success state and the specified result. /// The result value. diff --git a/src/libraries/System.Threading.Channels/src/System/Threading/Channels/BoundedChannel.cs b/src/libraries/System.Threading.Channels/src/System/Threading/Channels/BoundedChannel.cs index 6e911a797121f5..fd17bd1185fe62 100644 --- a/src/libraries/System.Threading.Channels/src/System/Threading/Channels/BoundedChannel.cs +++ b/src/libraries/System.Threading.Channels/src/System/Threading/Channels/BoundedChannel.cs @@ -356,8 +356,7 @@ public override bool TryWrite(T item) while (!parent._blockedReaders.IsEmpty) { AsyncOperation r = parent._blockedReaders.DequeueHead(); - r.UnregisterCancellation(); // ensure that once we grab it, we own its completion - if (!r.IsCompleted) + if (r.UnregisterCancellation()) // ensure that once we grab it, we own its completion { blockedReader = r; break; @@ -517,8 +516,7 @@ public override ValueTask WriteAsync(T item, CancellationToken cancellationToken while (!parent._blockedReaders.IsEmpty) { AsyncOperation r = parent._blockedReaders.DequeueHead(); - r.UnregisterCancellation(); // ensure that once we grab it, we own its completion - if (!r.IsCompleted) + if (r.UnregisterCancellation()) // ensure that once we grab it, we own its completion { blockedReader = r; break; diff --git a/src/libraries/System.Threading.Channels/tests/Stress.cs b/src/libraries/System.Threading.Channels/tests/Stress.cs index d2d80e5c2d89e0..61354e844d22d1 100644 --- a/src/libraries/System.Threading.Channels/tests/Stress.cs +++ b/src/libraries/System.Threading.Channels/tests/Stress.cs @@ -2,16 +2,16 @@ // The .NET Foundation licenses this file to you under the MIT license. // See the LICENSE file in the project root for more information. -using System.Threading.Tasks; using System.Collections.Generic; -using System; +using System.Linq; +using System.Threading.Tasks; using Xunit; namespace System.Threading.Channels.Tests { public class StressTests { - public static IEnumerable TestData() + public static IEnumerable ReadWriteVariations_TestData() { foreach (var readDelegate in new Func, Task>[] { ReadSynchronous, ReadAsynchronous, ReadSyncAndAsync} ) foreach (var writeDelegate in new Func, int, Task>[] { WriteSynchronous, WriteAsynchronous, WriteSyncAndAsync} ) @@ -122,12 +122,12 @@ private static async Task WriteSyncAndAsync(ChannelWriter writer, int value private static readonly int MaxTaskCounts = Math.Max(2, Environment.ProcessorCount); [ConditionalTheory(typeof(TestEnvironment), nameof(TestEnvironment.IsStressModeEnabled))] - [MemberData(nameof(TestData))] - public void RunInStressMode( - Func> channelCreator, - ChannelOptions options, - Func, Task> readDelegate, - Func, int, Task> writeDelegate) + [MemberData(nameof(ReadWriteVariations_TestData))] + public void ReadWriteVariations( + Func> channelCreator, + ChannelOptions options, + Func, Task> readDelegate, + Func, int, Task> writeDelegate) { Channel channel = channelCreator(options); ChannelReader reader = channel.Reader; @@ -206,7 +206,68 @@ public void RunInStressMode( { Assert.InRange(readCount, 0, MaxNumberToWriteToChannel); } + } + + public static IEnumerable CanceledReads_TestData() + { + yield return new object[] { new Func>(() => Channel.CreateUnbounded()) }; + yield return new object[] { new Func>(() => Channel.CreateUnbounded(new UnboundedChannelOptions() { SingleReader = true, SingleWriter = true })) }; + yield return new object[] { new Func>(() => Channel.CreateBounded(int.MaxValue)) }; + } + + [ConditionalTheory(typeof(TestEnvironment), nameof(TestEnvironment.IsStressModeEnabled))] + [MemberData(nameof(CanceledReads_TestData))] + public async Task CanceledReads(Func> channelFactory) + { + const int Attempts = 100; + const int Writes = 1_000; + const int WaitTimeoutMs = 100_000; + + for (int i = 0; i < Attempts; i++) + { + var cts = new CancellationTokenSource(); + Channel channel = channelFactory(); + + // Create a bunch of reads, half of which are cancelable + Task[] reads = Enumerable.Range(0, Writes).Select(i => channel.Reader.ReadAsync(i % 2 == 0 ? cts.Token : default).AsTask()).ToArray(); + + // Queue cancellation + _ = Task.Run(() => cts.Cancel()); + + // Write to complete the rest of the tasks + for (int item = 0; item < Writes; item++) + { + Assert.True(channel.Writer.TryWrite(item)); + } + channel.Writer.Complete(); + + // Wait for all the reads to complete + try + { + Assert.True(Task.WaitAll(reads, WaitTimeoutMs)); + } + catch (AggregateException ae) + { + Assert.All(ae.InnerExceptions, e => Assert.IsAssignableFrom(e)); + } + // Validate all write data showed up + int expected = 0; + int actual = 0; + for (int write = 0; write < Writes; write++) + { + expected += write; + if (reads[write].Status == TaskStatus.RanToCompletion) + { + actual += reads[write].Result; + } + } + await foreach (int remaining in channel.Reader.ReadAllAsync()) + { + actual += remaining; + } + Assert.Equal(expected, actual); + } } } } diff --git a/src/libraries/System.Threading.Channels/tests/System.Threading.Channels.Tests.csproj b/src/libraries/System.Threading.Channels/tests/System.Threading.Channels.Tests.csproj index a2811146aa12f7..5c1178c3d1c646 100644 --- a/src/libraries/System.Threading.Channels/tests/System.Threading.Channels.Tests.csproj +++ b/src/libraries/System.Threading.Channels/tests/System.Threading.Channels.Tests.csproj @@ -9,12 +9,12 @@ - +