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 @@
-
+