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
229 changes: 229 additions & 0 deletions EnumerableAsyncProcessor.UnitTests/CancellationAwareSelectorTests.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,229 @@
using System;
using System.Collections.Generic;
using System.Runtime.CompilerServices;
using System.Threading;
using System.Threading.Tasks;
using EnumerableAsyncProcessor.Builders;
using EnumerableAsyncProcessor.Extensions;

namespace EnumerableAsyncProcessor.UnitTests;

public class CancellationAwareSelectorTests
{
[Test]
public async Task CancelAll_Interrupts_InFlight_ForEachAsync_Selector(CancellationToken cancellationToken)
{
var started = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);
var interrupted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);

await using var processor = new[] { 1 }
.ForEachAsync((_, processorToken) => WaitUntilCanceledAsync(processorToken, started, interrupted))
.ProcessInParallel();

await started.Task.WaitAsync(cancellationToken);
processor.CancelAll();

await interrupted.Task.WaitAsync(cancellationToken);
await Assert.ThrowsAsync<TaskCanceledException>(() => processor.WaitAsync());
}

[Test]
public async Task CancelAll_Interrupts_InFlight_SelectAsync_Selector(CancellationToken cancellationToken)
{
var started = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);
var interrupted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);

await using var processor = new[] { 1 }
.SelectAsync((item, processorToken) => WaitUntilCanceledAsync(item, processorToken, started, interrupted))
.ProcessInParallel();

await started.Task.WaitAsync(cancellationToken);
processor.CancelAll();

await interrupted.Task.WaitAsync(cancellationToken);
await Assert.ThrowsAsync<TaskCanceledException>(() => processor.GetResultsAsync());
}

[Test]
public async Task CancelAll_Interrupts_ExecutionCount_Selector(CancellationToken cancellationToken)
{
var started = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);
var interrupted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);

await using var processor = AsyncProcessorBuilder.WithExecutionCount(1)
.ForEachAsync(processorToken => WaitUntilCanceledAsync(processorToken, started, interrupted))
.ProcessInParallel();

await started.Task.WaitAsync(cancellationToken);
processor.CancelAll();

await interrupted.Task.WaitAsync(cancellationToken);
await Assert.ThrowsAsync<TaskCanceledException>(() => processor.WaitAsync());
}

[Test]
public async Task DisposeAsync_Interrupts_InFlight_Selector(CancellationToken cancellationToken)
{
var started = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);
var interrupted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);

var processor = new[] { 1 }
.ForEachAsync((_, processorToken) => WaitUntilCanceledAsync(processorToken, started, interrupted))
.ProcessInParallel();

await started.Task.WaitAsync(cancellationToken);
var disposeTask = processor.DisposeAsync().AsTask();

await interrupted.Task.WaitAsync(cancellationToken);
await disposeTask.WaitAsync(cancellationToken);
}

[Test]
public async Task ExternalCancellation_Interrupts_AsyncEnumerable_Selector(CancellationToken cancellationToken)
{
using var cancellationTokenSource = new CancellationTokenSource();
var started = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);
var interrupted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);
var allowCleanupToFinish = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);

var processor = GetItemsAsync()
.ForEachAsync(
(_, processorToken) => WaitUntilCanceledAfterCleanupAsync(
processorToken,
started,
interrupted,
allowCleanupToFinish),
cancellationTokenSource.Token)
.ProcessInParallel();

var executeTask = processor.ExecuteAsync();
await started.Task.WaitAsync(cancellationToken);
await cancellationTokenSource.CancelAsync();

await AssertExecutionWaitsForCleanupAsync(
executeTask,
interrupted,
allowCleanupToFinish,
cancellationToken);
}

[Test]
public async Task ExternalCancellation_Waits_For_AsyncEnumerable_Result_Selector_Cleanup(CancellationToken cancellationToken)
{
using var cancellationTokenSource = new CancellationTokenSource();
var started = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);
var interrupted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);
var allowCleanupToFinish = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);

var processor = GetItemsAsync()
.SelectAsync(
(item, processorToken) => WaitUntilCanceledAfterCleanupAsync(
item,
processorToken,
started,
interrupted,
allowCleanupToFinish),
cancellationTokenSource.Token)
.ProcessInParallel();

var executeTask = processor.ExecuteAsync().ToListAsync();
await started.Task.WaitAsync(cancellationToken);
await cancellationTokenSource.CancelAsync();

await AssertExecutionWaitsForCleanupAsync(
executeTask,
interrupted,
allowCleanupToFinish,
cancellationToken);
}

private static async Task WaitUntilCanceledAsync(
CancellationToken cancellationToken,
TaskCompletionSource started,
TaskCompletionSource interrupted)
{
started.TrySetResult();
try
{
await Task.Delay(Timeout.InfiniteTimeSpan, cancellationToken);
}
catch (OperationCanceledException)
{
interrupted.TrySetResult();
throw;
}
}

private static async Task<T> WaitUntilCanceledAsync<T>(
T result,
CancellationToken cancellationToken,
TaskCompletionSource started,
TaskCompletionSource interrupted)
{
await WaitUntilCanceledAsync(cancellationToken, started, interrupted);
return result;
}

private static async Task WaitUntilCanceledAfterCleanupAsync(
CancellationToken cancellationToken,
TaskCompletionSource started,
TaskCompletionSource interrupted,
TaskCompletionSource allowCleanupToFinish)
{
started.TrySetResult();
try
{
await Task.Delay(Timeout.InfiniteTimeSpan, cancellationToken);
}
catch (OperationCanceledException)
{
interrupted.TrySetResult();
await allowCleanupToFinish.Task;
throw;
}
}

private static async Task<T> WaitUntilCanceledAfterCleanupAsync<T>(
T result,
CancellationToken cancellationToken,
TaskCompletionSource started,
TaskCompletionSource interrupted,
TaskCompletionSource allowCleanupToFinish)
{
await WaitUntilCanceledAfterCleanupAsync(
cancellationToken,
started,
interrupted,
allowCleanupToFinish);

return result;
}

private static async Task AssertExecutionWaitsForCleanupAsync(
Task executeTask,
TaskCompletionSource interrupted,
TaskCompletionSource allowCleanupToFinish,
CancellationToken cancellationToken)
{
await interrupted.Task.WaitAsync(cancellationToken);
await Task.Delay(100, cancellationToken);

try
{
await Assert.That(executeTask.IsCompleted).IsFalse();
}
finally
{
allowCleanupToFinish.TrySetResult();
}

await Assert.ThrowsAsync<OperationCanceledException>(() => executeTask);
}

private static async IAsyncEnumerable<int> GetItemsAsync(
[EnumeratorCancellation] CancellationToken cancellationToken = default)
{
yield return 1;
await Task.Delay(Timeout.InfiniteTimeSpan, cancellationToken);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,13 @@ public ActionAsyncProcessorBuilder(int count, Func<Task> taskSelector, Cancellat
_cancellationTokenSource = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
}

public ActionAsyncProcessorBuilder(int count, Func<CancellationToken, Task> taskSelector, CancellationToken cancellationToken)
{
_count = count;
_cancellationTokenSource = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
_taskSelector = () => taskSelector(_cancellationTokenSource.Token);
}

public IAsyncProcessor ProcessInBatches(int batchSize)
{
return new BatchAsyncProcessor(batchSize, _count, _taskSelector, _cancellationTokenSource).StartProcessing();
Expand Down Expand Up @@ -77,4 +84,4 @@ public IAsyncProcessor ProcessOneAtATime()
return new OneAtATimeAsyncProcessor(_count, _taskSelector, _cancellationTokenSource).StartProcessing();
}

}
}
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,13 @@ internal ActionAsyncProcessorBuilder(int count, Func<Task<TOutput>> taskSelector
_cancellationTokenSource = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
}

internal ActionAsyncProcessorBuilder(int count, Func<CancellationToken, Task<TOutput>> taskSelector, CancellationToken cancellationToken)
{
_count = count;
_cancellationTokenSource = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
_taskSelector = () => taskSelector(_cancellationTokenSource.Token);
}

public IAsyncProcessor<TOutput> ProcessInBatches(int batchSize)
{
return new ResultBatchAsyncProcessor<TOutput>(batchSize, _count, _taskSelector, _cancellationTokenSource).StartProcessing();
Expand Down Expand Up @@ -77,4 +84,4 @@ public IAsyncProcessor<TOutput> ProcessOneAtATime()
return new ResultOneAtATimeAsyncProcessor<TOutput>(_count, _taskSelector, _cancellationTokenSource).StartProcessing();
}

}
}
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,16 @@ public AsyncEnumerableActionAsyncProcessorBuilder(
_cancellationTokenSource = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
}

public AsyncEnumerableActionAsyncProcessorBuilder(
IAsyncEnumerable<TInput> items,
Func<TInput, CancellationToken, Task> taskSelector,
CancellationToken cancellationToken)
{
_items = items;
_cancellationTokenSource = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
_taskSelector = item => taskSelector(item, _cancellationTokenSource.Token);
Comment thread
thomhurst marked this conversation as resolved.
}

/// <summary>
/// Process items in parallel without concurrency limits.
/// </summary>
Expand Down Expand Up @@ -70,7 +80,6 @@ public IAsyncEnumerableProcessor ProcessInParallel(int? maxConcurrency, bool sch
_items, _taskSelector, maxConcurrency, scheduleOnThreadPool, _cancellationTokenSource);
}


/// <summary>
/// Process items one at a time (sequential processing).
/// </summary>
Expand All @@ -91,4 +100,4 @@ public IAsyncEnumerableProcessor ProcessInBatches(int batchSize)
_items, _taskSelector, batchSize, _cancellationTokenSource);
}

}
}
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,16 @@ public AsyncEnumerableActionAsyncProcessorBuilder(
_cancellationTokenSource = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
}

public AsyncEnumerableActionAsyncProcessorBuilder(
IAsyncEnumerable<TInput> items,
Func<TInput, CancellationToken, Task<TOutput>> taskSelector,
CancellationToken cancellationToken)
{
_items = items;
_cancellationTokenSource = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
_taskSelector = item => taskSelector(item, _cancellationTokenSource.Token);
}

/// <summary>
/// Process items in parallel without concurrency limits and return results.
/// </summary>
Expand Down Expand Up @@ -70,7 +80,6 @@ public IAsyncEnumerableProcessor<TOutput> ProcessInParallel(int? maxConcurrency,
_items, _taskSelector, maxConcurrency, scheduleOnThreadPool, _cancellationTokenSource);
}


/// <summary>
/// Process items one at a time and return results in order.
/// </summary>
Expand All @@ -91,4 +100,4 @@ public IAsyncEnumerableProcessor<TOutput> ProcessInBatches(int batchSize)
_items, _taskSelector, batchSize, _cancellationTokenSource);
}

}
}
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,19 @@ public AsyncEnumerableActionAsyncProcessorBuilder<TInput, TOutput> SelectAsync<T
return new AsyncEnumerableActionAsyncProcessorBuilder<TInput, TOutput>(_items, taskSelector, cancellationToken);
}

public AsyncEnumerableActionAsyncProcessorBuilder<TInput, TOutput> SelectAsync<TOutput>(
Func<TInput, CancellationToken, Task<TOutput>> taskSelector)
{
return SelectAsync(taskSelector, CancellationToken.None);
}

public AsyncEnumerableActionAsyncProcessorBuilder<TInput, TOutput> SelectAsync<TOutput>(
Func<TInput, CancellationToken, Task<TOutput>> taskSelector,
CancellationToken cancellationToken)
{
return new AsyncEnumerableActionAsyncProcessorBuilder<TInput, TOutput>(_items, taskSelector, cancellationToken);
}

public AsyncEnumerableActionAsyncProcessorBuilder<TInput> ForEachAsync(
Func<TInput, Task> taskSelector)
{
Expand All @@ -34,4 +47,17 @@ public AsyncEnumerableActionAsyncProcessorBuilder<TInput> ForEachAsync(
{
return new AsyncEnumerableActionAsyncProcessorBuilder<TInput>(_items, taskSelector, cancellationToken);
}
}

public AsyncEnumerableActionAsyncProcessorBuilder<TInput> ForEachAsync(
Func<TInput, CancellationToken, Task> taskSelector)
{
return ForEachAsync(taskSelector, CancellationToken.None);
}

public AsyncEnumerableActionAsyncProcessorBuilder<TInput> ForEachAsync(
Func<TInput, CancellationToken, Task> taskSelector,
CancellationToken cancellationToken)
{
return new AsyncEnumerableActionAsyncProcessorBuilder<TInput>(_items, taskSelector, cancellationToken);
}
}
Loading
Loading