Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -124,19 +124,6 @@ public ValueTaskSourceStatus GetStatus(short token)
}

/// <summary>Gets whether the operation has completed.</summary>
/// <remarks>
/// 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
/// </remarks>
internal bool IsCompleted => ReferenceEquals(_continuation, s_completedSentinel);

/// <summary>Gets the result of the operation.</summary>
Expand Down Expand Up @@ -301,15 +288,29 @@ public void OnCompleted(Action<object?> continuation, object? state, short token
}
}

/// <summary>Unregisters from cancellation.</summary>
/// <summary>Unregisters from cancellation and returns whether cancellation already started.</summary>
/// <returns>
/// 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.
/// </returns>
/// <remarks>
/// 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.
/// 2. To avoid having to worry about concurrent completion; once invoked, the caller can be guaranteed
/// 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).
/// </remarks>
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;
}

/// <summary>Completes the operation with a success state and the specified result.</summary>
/// <param name="item">The result value.</param>
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -356,8 +356,7 @@ public override bool TryWrite(T item)
while (!parent._blockedReaders.IsEmpty)
{
AsyncOperation<T> 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;
Expand Down Expand Up @@ -517,8 +516,7 @@ public override ValueTask WriteAsync(T item, CancellationToken cancellationToken
while (!parent._blockedReaders.IsEmpty)
{
AsyncOperation<T> 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;
Expand Down
79 changes: 70 additions & 9 deletions src/libraries/System.Threading.Channels/tests/Stress.cs
Original file line number Diff line number Diff line change
Expand Up @@ -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<object[]> TestData()
public static IEnumerable<object[]> ReadWriteVariations_TestData()
{
foreach (var readDelegate in new Func<ChannelReader<int>, Task<bool>>[] { ReadSynchronous, ReadAsynchronous, ReadSyncAndAsync} )
foreach (var writeDelegate in new Func<ChannelWriter<int>, int, Task>[] { WriteSynchronous, WriteAsynchronous, WriteSyncAndAsync} )
Expand Down Expand Up @@ -122,12 +122,12 @@ private static async Task WriteSyncAndAsync(ChannelWriter<int> 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<ChannelOptions, Channel<int>> channelCreator,
ChannelOptions options,
Func<ChannelReader<int>, Task<bool>> readDelegate,
Func<ChannelWriter<int>, int, Task> writeDelegate)
[MemberData(nameof(ReadWriteVariations_TestData))]
public void ReadWriteVariations(
Func<ChannelOptions, Channel<int>> channelCreator,
ChannelOptions options,
Func<ChannelReader<int>, Task<bool>> readDelegate,
Func<ChannelWriter<int>, int, Task> writeDelegate)
{
Channel<int> channel = channelCreator(options);
ChannelReader<int> reader = channel.Reader;
Expand Down Expand Up @@ -206,7 +206,68 @@ public void RunInStressMode(
{
Assert.InRange(readCount, 0, MaxNumberToWriteToChannel);
}
}

public static IEnumerable<object[]> CanceledReads_TestData()
{
yield return new object[] { new Func<Channel<int>>(() => Channel.CreateUnbounded<int>()) };
yield return new object[] { new Func<Channel<int>>(() => Channel.CreateUnbounded<int>(new UnboundedChannelOptions() { SingleReader = true, SingleWriter = true })) };
yield return new object[] { new Func<Channel<int>>(() => Channel.CreateBounded<int>(int.MaxValue)) };
}

[ConditionalTheory(typeof(TestEnvironment), nameof(TestEnvironment.IsStressModeEnabled))]
[MemberData(nameof(CanceledReads_TestData))]
public async Task CanceledReads(Func<Channel<int>> 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<int> channel = channelFactory();

// Create a bunch of reads, half of which are cancelable
Task<int>[] 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<OperationCanceledException>(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);
}
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -9,12 +9,12 @@
<Compile Include="ChannelTests.cs" />
<Compile Include="TestBase.cs" />
<Compile Include="UnboundedChannelTests.cs" />
<Compile Include="Stress.cs" />
<Compile Include="DebugAttributeTests.cs" />
<Compile Include="$(CommonTestPath)System\Diagnostics\DebuggerAttributes.cs"
Link="Common\System\Diagnostics\DebuggerAttributes.cs" />
</ItemGroup>
<ItemGroup Condition="'$(TargetFramework)' == '$(NetCoreAppCurrent)'">
<Compile Include="Stress.cs" />
<Compile Include="ChannelClosedExceptionTests.netcoreapp.cs" />
<Compile Include="ChannelTestBase.netcoreapp.cs" />
</ItemGroup>
Expand Down