From 3e4d6f1c8dc6c00c3b4191954369fede8c4b872e Mon Sep 17 00:00:00 2001 From: Stephen Toub Date: Wed, 10 Jun 2020 17:06:52 -0400 Subject: [PATCH 1/2] Fix BoundedChannel race condition between writes and canceled reads When BoundedChannelWriter has an item to be written, if there are readers waiting, it needs to try to transfer that item to a waiting reader. When it dequeues the next waiting reader, it needs to check to see if that reader has already been canceled, as that's the only way a reader in the queue might not be completable by the writer. Our current check, however, is flawed, in that we're assuming that the reader's IsCompleted will have synchronously transitioned to true as part of cancellation being requested; however, due to an implementation detail of the AsyncOperation that represents the reader, it might _asynchronously_ transition to true, in which case we have a race condition where we might see the reader as not having been canceled even though it has been, and that causes us to lose data by the writer trying to complete the reader with its data and failing to do so (this condition would assert in debug builds but isn't checked in release builds as it should never happen). The fix is simply changing the condition we check to factor in the right information. --- .../Threading/Channels/AsyncOperation.cs | 31 ++++++++++--------- .../Threading/Channels/BoundedChannel.cs | 6 ++-- 2 files changed, 18 insertions(+), 19 deletions(-) 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; From 4941062e346538e1c42a9cabc6b9e8cf20e79238 Mon Sep 17 00:00:00 2001 From: Stephen Toub Date: Thu, 11 Jun 2020 14:09:40 -0400 Subject: [PATCH 2/2] Add stress test for canceled reads This fails quickly before the fix and succeeds after. --- .../System.Threading.Channels/tests/Stress.cs | 79 ++++++++++++++++--- .../System.Threading.Channels.Tests.csproj | 2 +- 2 files changed, 71 insertions(+), 10 deletions(-) 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 @@ - +