How to use Rx.Nex extension ForEachAsync with async action

Viewed 1761

I have code which streams data down from SQL and writes it to a different store. The code is approximately this:

using (var cmd = new SqlCommand("select * from MyTable", connection))
{
     using (var reader = await cmd.ExecuteReaderAsync())
     {
         var list = new List<MyData>();
         while (await reader.ReadAsync())
         {
             var row = GetRow(reader);
             list.Add(row);
             if (list.Count == BatchSize)
             {
                 await WriteDataAsync(list);
                 list.Clear();
             }
         }
         if (list.Count > 0)
         {
             await WriteDataAsync(list);
         }
     }
 }

I would like to use Reactive extensions for this purpose instead. Ideally the code would look like this:

await StreamDataFromSql()
    .Buffer(BatchSize)
    .ForEachAsync(async batch => await WriteDataAsync(batch));

However, it seems that the extension method ForEachAsync only accepts synchronous actions. Would it be possible to write an extension which would accept an async action?

4 Answers

Here is a version of the ForEachAsync method that supports asynchronous actions. It projects the source observable to a nested IObservable<IObservable<Unit>> containing the asynchronous actions, and then flattens it back to an IObservable<Unit> using the Merge operator. The resulting observable is finally converted to a task.

By default the actions are invoked sequentially, but it is possible to invoke them concurrently by configuring the optional maximumConcurrency argument.

Canceling the optional cancellationToken argument results to the immediate completion (cancellation) of the returned Task, potentially before the cancellation of the currently running actions.

Any exception that may occur is propagated through the Task, and causes the cancellation of all currently running actions.

/// <summary>
/// Invokes an asynchronous action for each element in the observable sequence,
/// and returns a 'Task' that represents the completion of the sequence and
/// all the asynchronous actions.
/// </summary>
public static Task ForEachAsync<TSource>(
    this IObservable<TSource> source,
    Func<TSource, CancellationToken, Task> action,
    CancellationToken cancellationToken = default,
    int maximumConcurrency = 1)
{
    // Arguments validation omitted
    return source
        .Select(item => Observable.FromAsync(ct => action(item, ct)))
        .Merge(maximumConcurrency)
        .DefaultIfEmpty()
        .ToTask(cancellationToken);
}

Usage example:

await StreamDataFromSql()
    .Buffer(BatchSize)
    .ForEachAsync(async (batch, token) => await WriteDataAsync(batch, token));
Related