diff --git a/src/Stunts.CompiledProxy/Processors/CSharpAot.cs b/src/Stunts.CompiledProxy/Processors/CSharpAot.cs index b729a93..d58c8c5 100644 --- a/src/Stunts.CompiledProxy/Processors/CSharpAot.cs +++ b/src/Stunts.CompiledProxy/Processors/CSharpAot.cs @@ -33,8 +33,27 @@ public SyntaxNode Process(SyntaxNode syntax, ProcessorContext context) return syntax; var types = new HashSet(SymbolEqualityComparer.Default); + var asyncAdapters = new HashSet(); var registrations = new List(); var usesQueryable = false; + void RegisterAwaitable(ITypeSymbol type) + { + if (type is not INamedTypeSymbol named || !named.IsGenericType) + return; + + var definition = named.OriginalDefinition.ToDisplayString(); + if (definition != "System.Threading.Tasks.Task" && + definition != "System.Threading.Tasks.ValueTask") + return; + + var argument = named.TypeArguments[0]; + if (argument.TypeKind is TypeKind.TypeParameter or TypeKind.Error || argument.IsRefLikeType) + return; + + var name = argument.ToDisplayString(SymbolDisplayFormat.FullyQualifiedFormat); + if (asyncAdapters.Add(name)) + registrations.Add(ParseStatement($"global::Stunts.AsyncRegistry.Register<{name}>();")); + } void Register(ITypeSymbol type) { if (type.SpecialType == SpecialType.System_Void || type.IsRefLikeType || @@ -82,6 +101,7 @@ type.TypeKind is TypeKind.Pointer or TypeKind.FunctionPointer or TypeKind.TypePa foreach (var method in symbol.GetMembers().OfType().Where(method => !method.IsGenericMethod && !method.IsStatic)) { Register(method.ReturnType); + RegisterAwaitable(method.ReturnType); foreach (var parameter in method.Parameters.Where(parameter => parameter.RefKind == RefKind.Out)) Register(parameter.Type); } diff --git a/src/Stunts.UnitTests/Castle/ProceedClassTests.cs b/src/Stunts.UnitTests/Castle/ProceedClassTests.cs new file mode 100644 index 0000000..1ec558f --- /dev/null +++ b/src/Stunts.UnitTests/Castle/ProceedClassTests.cs @@ -0,0 +1,42 @@ +using System; +using System.Threading.Tasks; +using Xunit; + +namespace Stunts.UnitTests.Castle +{ + public class ProceedClassTests : IRunnable + { + public void Run() + { + BaseExceptionIsTheOutcome(); + BaseTaskCanBeAdjusted().GetAwaiter().GetResult(); + } + + public void BaseExceptionIsTheOutcome() + { + Calculator stunt = Stunt.For().AddBehavior( + (invocation, next) => invocation.Proceed(next, (outcome, again) => new ValueTask(outcome)), + invocation => !invocation.MethodBase.IsConstructor).ToObject(); + + var exception = Assert.Throws(() => stunt.Read()); + Assert.Equal("base", exception.Message); + } + + public async Task BaseTaskCanBeAdjusted() + { + Calculator stunt = Stunt.For().AddBehavior( + (invocation, next) => invocation.Proceed(next, (outcome, again) => + new ValueTask(ProceedOutcome.FromValue((int)outcome.Value! + 1, outcome.Elapsed))), + invocation => !invocation.MethodBase.IsConstructor).ToObject(); + + Assert.Equal(4, await stunt.GetAsync()); + } + + public class Calculator + { + public virtual int Read() => throw new InvalidOperationException("base"); + + public virtual Task GetAsync() => Task.FromResult(3); + } + } +} diff --git a/src/Stunts.UnitTests/ProceedTests.cs b/src/Stunts.UnitTests/ProceedTests.cs new file mode 100644 index 0000000..81647f4 --- /dev/null +++ b/src/Stunts.UnitTests/ProceedTests.cs @@ -0,0 +1,290 @@ +using System; +using System.IO; +using System.Reflection; +using System.Threading; +using System.Threading.Tasks; +using Xunit; + +namespace Stunts.UnitTests +{ + public class ProceedTests + { + [Fact] + public void SynchronousValuePassesThrough() + { + var result = Proceed(nameof(ISample.Read), (invocation, next) => invocation.CreateValueReturn(3), (outcome, again) => + new ValueTask(outcome)); + + Assert.Equal(3, result.ReturnValue); + Assert.Null(result.Exception); + } + + [Fact] + public void SynchronousCallbackCanReplaceTheValue() + { + var result = Proceed(nameof(ISample.Read), (invocation, next) => invocation.CreateValueReturn(2), (outcome, again) => + new ValueTask(ProceedOutcome.FromValue((int)outcome.Value! + 1, outcome.Elapsed))); + + Assert.Equal(3, result.ReturnValue); + } + + [Fact] + public void SynchronousThrowBecomesTheOutcome() + { + var thrown = new InvalidOperationException("x"); + var seen = default(Exception); + + var result = Proceed(nameof(ISample.Read), (invocation, next) => throw thrown, (outcome, again) => + { + seen = outcome.Exception; + return new ValueTask(outcome); + }); + + Assert.Same(thrown, seen); + Assert.Same(thrown, result.Exception); + } + + [Fact] + public void SynchronousCallbackCanRecoverAnException() + { + var result = Proceed(nameof(ISample.Read), (invocation, next) => throw new InvalidOperationException(), (outcome, again) => + new ValueTask(ProceedOutcome.FromValue(4, outcome.Elapsed))); + + Assert.Null(result.Exception); + Assert.Equal(4, result.ReturnValue); + } + + [Fact] + public async Task SynchronousAgainRetries() + { + var calls = 0; + + var result = Proceed(nameof(ISample.Read), (invocation, next) => invocation.CreateValueReturn(calls++), async (outcome, again) => + { + while ((int)outcome.Value! < 2) + outcome = await again(); + return outcome; + }); + + Assert.Equal(2, result.ReturnValue); + Assert.Equal(3, calls); + } + + [Fact] + public void SynchronousCallbackThatYieldsThrows() + { + var exception = Assert.Throws(() => + Proceed(nameof(ISample.Read), (invocation, next) => invocation.CreateValueReturn(1), async (outcome, again) => + { + await Task.Yield(); + return outcome; + })); + + Assert.Contains(nameof(ISample.Read), exception.Message); + } + + [Fact] + public void RefOutputIsPreserved() + { + var method = Sample(nameof(ISample.Inc)); + var invocation = MethodInvocation.Create(new object(), method, (invocation, next) => + { + var value = (int)invocation.Arguments.GetValue("value")!; + return invocation.CreateValueReturn(null, invocation.Arguments.SetValue("value", value + 1)); + }, 1); + + var result = invocation.Proceed((call, next) => call.CreateInvokeReturn(), (outcome, again) => + new ValueTask(outcome)); + + Assert.Equal(2, result.Outputs.GetValue("value")); + } + + [Fact] + public async Task TaskResultIsUnwrappedAndCanBeReplaced() + { + var result = Proceed(nameof(ISample.GetAsync), (invocation, next) => invocation.CreateValueReturn(Task.FromResult(3)), (outcome, again) => + new ValueTask(ProceedOutcome.FromValue((int)outcome.Value! + 1, outcome.Elapsed))); + + var task = Assert.IsType>(result.ReturnValue); + Assert.Null(result.Exception); + Assert.Equal(4, await task); + } + + [Fact] + public async Task FaultedTaskKeepsTheOriginalException() + { + var original = new InvalidOperationException("boom"); + var result = Proceed(nameof(ISample.GetAsync), (invocation, next) => invocation.CreateValueReturn(Task.FromException(original)), (outcome, again) => + new ValueTask(outcome)); + + var task = Assert.IsType>(result.ReturnValue); + Assert.Null(result.Exception); + var actual = await Assert.ThrowsAsync(() => task); + Assert.Same(original, actual); + } + + [Fact] + public async Task CanceledTaskStaysCanceled() + { + using var source = new CancellationTokenSource(); + source.Cancel(); + var result = Proceed(nameof(ISample.GetAsync), (invocation, next) => + invocation.CreateValueReturn(Task.FromCanceled(source.Token)), (outcome, again) => + new ValueTask(outcome)); + + var task = Assert.IsType>(result.ReturnValue); + Assert.True(task.IsCanceled); + var actual = await Assert.ThrowsAsync(() => task); + Assert.Equal(source.Token, actual.CancellationToken); + } + + [Fact] + public async Task NullTaskIsASynchronousFailure() + { + Task? missing = null; + var result = Proceed(nameof(ISample.GetAsync), (invocation, next) => invocation.CreateValueReturn(missing), (outcome, again) => + new ValueTask(outcome)); + + var exception = Assert.IsType(result.Exception); + Assert.Contains(nameof(ISample.GetAsync), exception.Message); + } + + [Fact] + public async Task SynchronousFailureOnATaskMethodCanBeRecovered() + { + var result = Proceed(nameof(ISample.GetAsync), (invocation, next) => throw new InvalidOperationException("x"), (outcome, again) => + new ValueTask(ProceedOutcome.FromValue(4, outcome.Elapsed))); + + var task = Assert.IsType>(result.ReturnValue); + Assert.Equal(4, await task); + } + + [Fact] + public async Task RetryWaitsForTheAttempt() + { + var calls = 0; + var result = Proceed(nameof(ISample.GetAsync), (invocation, next) => + { + calls++; + return calls < 3 + ? invocation.CreateValueReturn(Task.FromException(new IOException())) + : invocation.CreateValueReturn(Task.FromResult(5)); + }, async (outcome, again) => + { + while (outcome.Exception is IOException) + outcome = await again(); + return outcome; + }); + + Assert.Equal(5, await Assert.IsType>(result.ReturnValue)); + Assert.Equal(3, calls); + } + + [Fact] + public async Task ElapsedIsReportedWhenTheTaskCompletes() + { + var gate = new TaskCompletionSource(); + TimeSpan? elapsed = null; + var result = Proceed(nameof(ISample.GetAsync), (invocation, next) => invocation.CreateValueReturn(gate.Task), async (outcome, again) => + { + elapsed = outcome.Elapsed; + return outcome; + }); + + var task = Assert.IsType>(result.ReturnValue); + Assert.False(task.IsCompleted); + Assert.Null(elapsed); + + gate.SetResult(9); + + Assert.Equal(9, await task); + Assert.NotNull(elapsed); + } + + [Fact] + public async Task NonGenericTaskCompletes() + { + var result = Proceed(nameof(ISample.RunAsync), (invocation, next) => invocation.CreateValueReturn(Task.CompletedTask), (outcome, again) => + { + Assert.Null(outcome.Exception); + Assert.Null(outcome.Value); + return new ValueTask(outcome); + }); + + await Assert.IsAssignableFrom(result.ReturnValue); + } + + [Fact] + public async Task ValueTaskIsConsumedOnceAcrossBehaviors() + { + var method = Sample(nameof(ISample.ReadAsync)); + var invocation = MethodInvocation.Create(new object(), method, (call, next) => call.CreateValueReturn(new ValueTask(7))); + var pipeline = new BehaviorPipeline( + (ExecuteHandler)((call, next) => call.Proceed(next, (outcome, again) => + new ValueTask(ProceedOutcome.FromValue((int)outcome.Value! + 1, outcome.Elapsed)))), + (ExecuteHandler)((call, next) => call.Proceed(next, (outcome, again) => + new ValueTask(outcome)))); + + var result = pipeline.Invoke(invocation); + + Assert.Equal(8, await Assert.IsType>(result.ReturnValue)); + } + + [Fact] + public async Task RegisteredAdapterWrapsAValueTask() + { + AsyncRegistry.Register(); + var id = Guid.NewGuid(); + var invocation = MethodInvocation.Create(new object(), Sample(nameof(ISample.IdAsync))); + var result = invocation.Proceed((call, next) => call.CreateValueReturn(new ValueTask(id)), (outcome, again) => + new ValueTask(outcome)); + + Assert.Equal(id, await Assert.IsType>(result.ReturnValue)); + } + + [Fact] + public void OpenGenericTaskHasNoAdapter() + { + var method = Sample(nameof(ISample.Echo)); + var invocation = new MethodInvocation(new object(), method, new ArgumentCollection(method.GetParameters())); + var exception = Assert.Throws(() => + invocation.Proceed((call, next) => call.CreateValueReturn(Task.FromResult(1)), (outcome, again) => + new ValueTask(outcome))); + + Assert.Contains("AsyncRegistry.Register<", exception.Message); + } + + [Fact] + public void RejectsAMissingCallback() + { + var invocation = MethodInvocation.Create(new object(), Sample(nameof(ISample.Read))); + + Assert.Throws(() => invocation.Proceed((call, next) => call.CreateValueReturn(1), null!)); + } + + static IMethodReturn Proceed(string name, ExecuteHandler next, ProceedHandler callback) + { + var invocation = MethodInvocation.Create(new object(), Sample(name)); + return invocation.Proceed(next, callback); + } + + static MethodInfo Sample(string name) => typeof(ISample).GetMethod(name)!; + + interface ISample + { + int Read(); + + void Inc(ref int value); + + Task RunAsync(); + + Task GetAsync(); + + ValueTask ReadAsync(); + + ValueTask IdAsync(); + + Task Echo(T value); + } + } +} diff --git a/src/Stunts.UnitTests/StuntGeneratorTests.cs b/src/Stunts.UnitTests/StuntGeneratorTests.cs index 2cdfb2f..a693301 100644 --- a/src/Stunts.UnitTests/StuntGeneratorTests.cs +++ b/src/Stunts.UnitTests/StuntGeneratorTests.cs @@ -94,6 +94,27 @@ public static class Test { public static IDisposable Create() => Stunt.Of(), Array.Empty())); } + [Fact] + public void RegistersClosedAsyncAdapters() + { + var (diagnostics, compilation) = GetGeneratedOutput(@" +using System.Threading.Tasks; +using Stunts; +public interface IService +{ + Task GetAsync(); + ValueTask ReadAsync(); + Task RunAsync(); +} +public static class Test { public static IService Create() => Stunt.Of(); }"); + + Assert.Empty(diagnostics); + var text = string.Concat(compilation.SyntaxTrees.Select(tree => tree.ToString())); + Assert.Contains("AsyncRegistry.Register()", text); + Assert.Contains("AsyncRegistry.Register()", text); + Assert.Equal(2, System.Text.RegularExpressions.Regex.Matches(text, "AsyncRegistry.Register<").Count); + } + [Fact] public void SkipsUncallableConstructorsInRegistration() { diff --git a/src/Stunts/AsyncRegistry.cs b/src/Stunts/AsyncRegistry.cs new file mode 100644 index 0000000..16d0ce6 --- /dev/null +++ b/src/Stunts/AsyncRegistry.cs @@ -0,0 +1,266 @@ +using System; +using System.Collections.Concurrent; +using System.ComponentModel; +using System.Diagnostics.CodeAnalysis; +using System.Linq; +using System.Reflection; +using System.Runtime.CompilerServices; +using System.Threading.Tasks; + +namespace Stunts +{ + /// + /// Typed adapters that await and rebuild and . + /// The source generator calls for every closed awaitable a stunt returns. + /// + public static class AsyncRegistry + { + static readonly ConcurrentDictionary adapters = new(); + + static AsyncRegistry() + { + adapters[typeof(Task)] = Adapter.PlainTask(); + adapters[typeof(ValueTask)] = Adapter.PlainValueTask(); + } + + /// + /// Registers await and wrap adapters for and . + /// + /// The result type carried by the awaitable. + [EditorBrowsable(EditorBrowsableState.Never)] + public static void Register() + { + adapters.TryAdd(typeof(Task), Adapter.TaskOf()); + adapters.TryAdd(typeof(ValueTask), Adapter.ValueTaskOf()); + } + + internal static Type? AwaitableReturn(MethodBase method) + { + if (method is not MethodInfo info || info.ReturnType == typeof(void) || info.ReturnType.IsByRef) + return null; + + var type = info.ReturnType; + if (type == typeof(Task) || type == typeof(ValueTask)) + return type; + if (!type.IsGenericType) + return null; + + var definition = type.GetGenericTypeDefinition(); + if (definition == typeof(Task<>) || definition == typeof(ValueTask<>)) + return type; + + return null; + } + + internal static ValueTask<(object? Value, Exception? Error)> Unwrap(Type awaitable, object instance) + => Require(awaitable).Unwrap(instance); + + internal static object Start(Type awaitable, Task<(object? Value, Exception? Error)> pending) + => Require(awaitable).Start(pending); + + internal static object FromOutcome(Type awaitable, ProceedOutcome outcome) + => Start(awaitable, Task.FromResult((outcome.Value, outcome.Exception))); + + [UnconditionalSuppressMessage("Trimming", "IL2026", Justification = "Reflection runs only when dynamic code is supported.")] + [UnconditionalSuppressMessage("AOT", "IL3050", Justification = "Reflection runs only when dynamic code is supported.")] + static Adapter Require(Type awaitable) + { + if (adapters.TryGetValue(awaitable, out var adapter)) + return adapter; + +#if NET8_0_OR_GREATER + if (!RuntimeFeature.IsDynamicCodeSupported) + throw NotRegistered(awaitable); +#endif + RegisterDynamic(awaitable); + return adapters[awaitable]; + } + + [RequiresDynamicCode("Closed task adapters are generated. Call AsyncRegistry.Register() for Native AOT.")] + [RequiresUnreferencedCode("Closed task adapters are generated. Call AsyncRegistry.Register() for Native AOT.")] + static void RegisterDynamic(Type awaitable) + { + var argument = awaitable.GenericTypeArguments[0]; + if (argument.ContainsGenericParameters) + throw NotRegistered(awaitable); + + var open = typeof(AsyncRegistry).GetMethods(BindingFlags.Public | BindingFlags.Static) + .Single(method => method.Name == nameof(Register) && method.IsGenericMethodDefinition); + open.MakeGenericMethod(argument).Invoke(null, null); + } + + static NotSupportedException NotRegistered(Type awaitable) + { + var argument = awaitable.IsGenericType ? awaitable.GenericTypeArguments[0] : awaitable; + return new NotSupportedException(ThisAssembly.Strings.AsyncAdapterNotRegistered( + CSharpTypeName.Format(awaitable), + CSharpTypeName.Format(argument))); + } + + sealed class Adapter + { + readonly Func> unwrap; + readonly Func, object> start; + + Adapter( + Func> unwrap, + Func, object> start) + { + this.unwrap = unwrap; + this.start = start; + } + + public ValueTask<(object? Value, Exception? Error)> Unwrap(object instance) => unwrap(instance); + + public object Start(Task<(object? Value, Exception? Error)> pending) => start(pending); + + public static Adapter PlainTask() => new(UnwrapTask, StartPlainTask); + + public static Adapter PlainValueTask() => new(UnwrapValueTask, pending => new ValueTask((Task)StartPlainTask(pending))); + + public static Adapter TaskOf() => new(UnwrapTaskOf, StartTask); + + public static Adapter ValueTaskOf() => new(UnwrapValueTaskOf, pending => new ValueTask((Task)StartTask(pending))); + + static async ValueTask<(object? Value, Exception? Error)> UnwrapTask(object instance) + { + try + { + await ((Task)instance).ConfigureAwait(false); + return (null, null); + } + catch (Exception exception) + { + return (null, exception); + } + } + + static async ValueTask<(object? Value, Exception? Error)> UnwrapValueTask(object instance) + { + try + { + await ((ValueTask)instance).ConfigureAwait(false); + return (null, null); + } + catch (Exception exception) + { + return (null, exception); + } + } + + static async ValueTask<(object? Value, Exception? Error)> UnwrapTaskOf(object instance) + { + try + { + return (await ((Task)instance).ConfigureAwait(false), null); + } + catch (Exception exception) + { + return (null, exception); + } + } + + static async ValueTask<(object? Value, Exception? Error)> UnwrapValueTaskOf(object instance) + { + try + { + return (await ((ValueTask)instance).ConfigureAwait(false), null); + } + catch (Exception exception) + { + return (null, exception); + } + } + + static object StartPlainTask(Task<(object? Value, Exception? Error)> pending) + { + if (pending.Status == TaskStatus.RanToCompletion) + return Box(pending.Result.Error); + + var source = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + pending.ContinueWith(task => + { + if (task.IsFaulted) + source.TrySetException(task.Exception!.InnerExceptions); + else if (task.IsCanceled) + source.TrySetCanceled(); + else + Complete(source, task.Result.Error, false); + }, TaskScheduler.Default); + return source.Task; + } + + static object StartTask(Task<(object? Value, Exception? Error)> pending) + { + if (pending.Status == TaskStatus.RanToCompletion) + return Box(pending.Result.Value, pending.Result.Error); + + var source = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + pending.ContinueWith(task => + { + if (task.IsFaulted) + source.TrySetException(task.Exception!.InnerExceptions); + else if (task.IsCanceled) + source.TrySetCanceled(); + else + Complete(source, task.Result.Value, task.Result.Error); + }, TaskScheduler.Default); + return source.Task; + } + + static object Box(Exception? error) + { + if (error is OperationCanceledException canceled) + return Task.FromCanceled(canceled.CancellationToken); + if (error != null) + return Task.FromException(error); + return Task.CompletedTask; + } + + static object Box(object? value, Exception? error) + { + if (error is OperationCanceledException canceled) + return Task.FromCanceled(canceled.CancellationToken); + if (error != null) + return Task.FromException(error); + return Task.FromResult((T)value!); + } + + static void Complete(TaskCompletionSource source, Exception? error, bool value) + { + if (error is OperationCanceledException canceled) + { +#if NET8_0_OR_GREATER + source.TrySetCanceled(canceled.CancellationToken); +#else + source.TrySetCanceled(); +#endif + return; + } + + if (error != null) + source.TrySetException(error); + else + source.TrySetResult(value); + } + + static void Complete(TaskCompletionSource source, object? value, Exception? error) + { + if (error is OperationCanceledException canceled) + { +#if NET8_0_OR_GREATER + source.TrySetCanceled(canceled.CancellationToken); +#else + source.TrySetCanceled(); +#endif + return; + } + + if (error != null) + source.TrySetException(error); + else + source.TrySetResult((T)value!); + } + } + } +} diff --git a/src/Stunts/ProceedExtension.cs b/src/Stunts/ProceedExtension.cs new file mode 100644 index 0000000..cd8f27c --- /dev/null +++ b/src/Stunts/ProceedExtension.cs @@ -0,0 +1,236 @@ +using System; +using System.Diagnostics; +using System.Threading.Tasks; + +namespace Stunts +{ + /// + /// The settled result of one attempt performed by . + /// + public readonly struct ProceedOutcome + { + /// + /// Creates a successful attempt. For and , + /// is the result, not the awaitable. + /// + public static ProceedOutcome FromValue(object? value, TimeSpan elapsed = default) + => new(value, null, elapsed, null); + + /// + /// Creates a failed attempt. Returning it from the callback fails the invocation. + /// + public static ProceedOutcome FromException(Exception exception, TimeSpan elapsed = default) + { + if (exception == null) + throw new ArgumentNullException(nameof(exception)); + + return new ProceedOutcome(null, exception, elapsed, null); + } + + internal static ProceedOutcome Success(object? value, TimeSpan elapsed, IArgumentCollection arguments) + => new(value, null, elapsed, arguments); + + internal static ProceedOutcome Failure(Exception exception, TimeSpan elapsed, IArgumentCollection arguments) + => new(null, exception, elapsed, arguments); + + ProceedOutcome(object? value, Exception? exception, TimeSpan elapsed, IArgumentCollection? arguments) + { + Value = value; + Exception = exception; + Elapsed = elapsed; + Arguments = arguments; + } + + /// The attempt's return value when is . + public object? Value { get; } + + /// The attempt's failure, or when it succeeded. + public Exception? Exception { get; } + + /// How long the attempt ran, including the time spent awaiting its task. + public TimeSpan Elapsed { get; } + + internal IArgumentCollection? Arguments { get; } + } + + /// Invokes the rest of the pipeline once and returns that attempt. + public delegate ValueTask ProceedAgain(); + + /// + /// Receives one settled attempt and returns the attempt the caller should observe. + /// + public delegate ValueTask ProceedHandler(ProceedOutcome outcome, ProceedAgain again); + + /// + /// Awaits task-returning members before the behavior inspects the result. + /// + public static class ProceedExtension + { + /// + /// Calls and passes the settled attempt to . + /// Synchronous members run the callback before this method returns. Task and value-task members + /// return a new awaitable of the same type immediately; the callback runs when that attempt finishes. + /// + /// + /// The receives an again delegate that calls once more + /// and waits for that attempt. A synchronous member's callback must itself finish synchronously. + /// Returning an exception fails the call. A synchronous failure is thrown by the stunt. + /// A faulted task stays faulted, with the original exception, instead of throwing before a task is returned. + /// + public static IMethodReturn Proceed(this IMethodInvocation invocation, ExecuteHandler next, ProceedHandler callback) + { + if (invocation == null) + throw new ArgumentNullException(nameof(invocation)); + if (next == null) + throw new ArgumentNullException(nameof(next)); + if (callback == null) + throw new ArgumentNullException(nameof(callback)); + + return new ProceedCall(invocation, next, callback).Execute(); + } + + sealed class ProceedCall + { + readonly IMethodInvocation invocation; + readonly ExecuteHandler next; + readonly ProceedHandler callback; + IArgumentCollection arguments; + + public ProceedCall(IMethodInvocation invocation, ExecuteHandler next, ProceedHandler callback) + { + this.invocation = invocation; + this.next = next; + this.callback = callback; + arguments = invocation.Arguments; + } + + public IMethodReturn Execute() + { + var started = Stopwatch.GetTimestamp(); + var result = InvokeNext(); + arguments = ArgumentsOf(result); + var awaitable = AsyncRegistry.AwaitableReturn(invocation.MethodBase); + + if (result.Exception != null || (awaitable != null && result.ReturnValue is null)) + { + var error = result.Exception ?? new NullReferenceException(ThisAssembly.Strings.NullAwaitable(invocation.MethodBase.Name)); + return Decide(ProceedOutcome.Failure(error, Elapsed(started), arguments), awaitable, synchronousFailure: true); + } + + if (awaitable == null) + return Decide(ProceedOutcome.Success(result.ReturnValue, Elapsed(started), arguments), null, synchronousFailure: true); + + return invocation.CreateValueReturn( + AsyncRegistry.Start(awaitable, FinishAsync(result, started, awaitable)), + invocation.Arguments); + } + + IMethodReturn Decide(ProceedOutcome outcome, Type? awaitable, bool synchronousFailure) + { + var pending = callback(outcome, Again); + if (!pending.IsCompleted) + { + if (awaitable == null) + { + Abandon(pending); + throw new InvalidOperationException(ThisAssembly.Strings.ProceedCallbackMustNotAwait(invocation.MethodBase.Name)); + } + + return invocation.CreateValueReturn(AsyncRegistry.Start(awaitable, Decision(pending)), invocation.Arguments); + } + + var decided = pending.GetAwaiter().GetResult(); + var used = decided.Arguments ?? arguments; + if (decided.Exception != null && (synchronousFailure || awaitable == null)) + return invocation.CreateExceptionReturn(decided.Exception); + if (awaitable == null) + return invocation.CreateValueReturn(decided.Value, used); + + return invocation.CreateValueReturn(AsyncRegistry.FromOutcome(awaitable, decided), invocation.Arguments); + } + + async Task<(object? Value, Exception? Error)> FinishAsync(IMethodReturn result, long started, Type awaitable) + { + var (value, error) = await AsyncRegistry.Unwrap(awaitable, result.ReturnValue!).ConfigureAwait(false); + var outcome = error != null + ? ProceedOutcome.Failure(error, Elapsed(started), arguments) + : ProceedOutcome.Success(value, Elapsed(started), arguments); + try + { + var decided = await callback(outcome, Again).ConfigureAwait(false); + if (decided.Arguments != null) + arguments = decided.Arguments; + return (decided.Value, decided.Exception); + } + catch (Exception exception) + { + return (null, exception); + } + } + + ValueTask Again() + { + var started = Stopwatch.GetTimestamp(); + var result = InvokeNext(); + var used = ArgumentsOf(result); + arguments = used; + var awaitable = AsyncRegistry.AwaitableReturn(invocation.MethodBase); + if (result.Exception != null) + return new ValueTask(ProceedOutcome.Failure(result.Exception, Elapsed(started), used)); + if (awaitable == null) + return new ValueTask(ProceedOutcome.Success(result.ReturnValue, Elapsed(started), used)); + if (result.ReturnValue is null) + return new ValueTask(ProceedOutcome.Failure( + new NullReferenceException(ThisAssembly.Strings.NullAwaitable(invocation.MethodBase.Name)), + Elapsed(started), + used)); + + return AgainAsync(result, started, awaitable, used); + } + + async ValueTask AgainAsync(IMethodReturn result, long started, Type awaitable, IArgumentCollection used) + { + var (value, error) = await AsyncRegistry.Unwrap(awaitable, result.ReturnValue!).ConfigureAwait(false); + return error != null + ? ProceedOutcome.Failure(error, Elapsed(started), used) + : ProceedOutcome.Success(value, Elapsed(started), used); + } + + IMethodReturn InvokeNext() + { + try + { + return next(invocation, next); + } + catch (Exception exception) + { + return invocation.CreateExceptionReturn(exception); + } + } + + IArgumentCollection ArgumentsOf(IMethodReturn result) + { + var merged = invocation.Arguments; + foreach (Argument output in result.Outputs) + merged = merged.SetValue(output.Name, output.RawValue); + return merged; + } + } + + static async Task<(object? Value, Exception? Error)> Decision(ValueTask pending) + { + var decided = await pending.ConfigureAwait(false); + return (decided.Value, decided.Exception); + } + + static void Abandon(ValueTask pending) + => _ = pending.AsTask().ContinueWith( + task => _ = task.Exception, + default, + TaskContinuationOptions.ExecuteSynchronously, + TaskScheduler.Default); + + static TimeSpan Elapsed(long started) + => TimeSpan.FromTicks((long)((Stopwatch.GetTimestamp() - started) * (double)TimeSpan.TicksPerSecond / Stopwatch.Frequency)); + } +} diff --git a/src/Stunts/Resources.resx b/src/Stunts/Resources.resx index f540ec0..c645c3a 100644 --- a/src/Stunts/Resources.resx +++ b/src/Stunts/Resources.resx @@ -162,6 +162,15 @@ Getting the next behavior from a method implementation is not supported. + + The Proceed callback must finish synchronously because {method} does not return a task or value task. + + + No async adapter is registered for {type}. Call AsyncRegistry.Register<{argument}>() from a module initializer. + + + {method} returned a null task. + Argument type '{argType}' is not compatible with its parameter '{paramType} {paramName}'.