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
20 changes: 20 additions & 0 deletions src/Stunts.CompiledProxy/Processors/CSharpAot.cs
Original file line number Diff line number Diff line change
Expand Up @@ -33,8 +33,27 @@ public SyntaxNode Process(SyntaxNode syntax, ProcessorContext context)
return syntax;

var types = new HashSet<ITypeSymbol>(SymbolEqualityComparer.Default);
var asyncAdapters = new HashSet<string>();
var registrations = new List<StatementSyntax>();
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<TResult>" &&
definition != "System.Threading.Tasks.ValueTask<TResult>")
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 ||
Expand Down Expand Up @@ -82,6 +101,7 @@ type.TypeKind is TypeKind.Pointer or TypeKind.FunctionPointer or TypeKind.TypePa
foreach (var method in symbol.GetMembers().OfType<IMethodSymbol>().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);
}
Expand Down
42 changes: 42 additions & 0 deletions src/Stunts.UnitTests/Castle/ProceedClassTests.cs
Original file line number Diff line number Diff line change
@@ -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<Calculator>().AddBehavior(
(invocation, next) => invocation.Proceed(next, (outcome, again) => new ValueTask<ProceedOutcome>(outcome)),
invocation => !invocation.MethodBase.IsConstructor).ToObject();

var exception = Assert.Throws<InvalidOperationException>(() => stunt.Read());
Assert.Equal("base", exception.Message);
}

public async Task BaseTaskCanBeAdjusted()
{
Calculator stunt = Stunt.For<Calculator>().AddBehavior(
(invocation, next) => invocation.Proceed(next, (outcome, again) =>
new ValueTask<ProceedOutcome>(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<int> GetAsync() => Task.FromResult(3);
}
}
}
290 changes: 290 additions & 0 deletions src/Stunts.UnitTests/ProceedTests.cs
Original file line number Diff line number Diff line change
@@ -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<ProceedOutcome>(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>(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<ProceedOutcome>(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>(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<InvalidOperationException>(() =>
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<ProceedOutcome>(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>(ProceedOutcome.FromValue((int)outcome.Value! + 1, outcome.Elapsed)));

var task = Assert.IsType<Task<int>>(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<int>(original)), (outcome, again) =>
new ValueTask<ProceedOutcome>(outcome));

var task = Assert.IsType<Task<int>>(result.ReturnValue);
Assert.Null(result.Exception);
var actual = await Assert.ThrowsAsync<InvalidOperationException>(() => 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<int>(source.Token)), (outcome, again) =>
new ValueTask<ProceedOutcome>(outcome));

var task = Assert.IsType<Task<int>>(result.ReturnValue);
Assert.True(task.IsCanceled);
var actual = await Assert.ThrowsAsync<TaskCanceledException>(() => task);
Assert.Equal(source.Token, actual.CancellationToken);
}

[Fact]
public async Task NullTaskIsASynchronousFailure()
{
Task<int>? missing = null;
var result = Proceed(nameof(ISample.GetAsync), (invocation, next) => invocation.CreateValueReturn(missing), (outcome, again) =>
new ValueTask<ProceedOutcome>(outcome));

var exception = Assert.IsType<NullReferenceException>(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>(ProceedOutcome.FromValue(4, outcome.Elapsed)));

var task = Assert.IsType<Task<int>>(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<int>(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<Task<int>>(result.ReturnValue));
Assert.Equal(3, calls);
}

[Fact]
public async Task ElapsedIsReportedWhenTheTaskCompletes()
{
var gate = new TaskCompletionSource<int>();
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<Task<int>>(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<ProceedOutcome>(outcome);
});

await Assert.IsAssignableFrom<Task>(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<int>(7)));
var pipeline = new BehaviorPipeline(
(ExecuteHandler)((call, next) => call.Proceed(next, (outcome, again) =>
new ValueTask<ProceedOutcome>(ProceedOutcome.FromValue((int)outcome.Value! + 1, outcome.Elapsed)))),
(ExecuteHandler)((call, next) => call.Proceed(next, (outcome, again) =>
new ValueTask<ProceedOutcome>(outcome))));

var result = pipeline.Invoke(invocation);

Assert.Equal(8, await Assert.IsType<ValueTask<int>>(result.ReturnValue));
}

[Fact]
public async Task RegisteredAdapterWrapsAValueTask()
{
AsyncRegistry.Register<Guid>();
var id = Guid.NewGuid();
var invocation = MethodInvocation.Create(new object(), Sample(nameof(ISample.IdAsync)));
var result = invocation.Proceed((call, next) => call.CreateValueReturn(new ValueTask<Guid>(id)), (outcome, again) =>
new ValueTask<ProceedOutcome>(outcome));

Assert.Equal(id, await Assert.IsType<ValueTask<Guid>>(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<NotSupportedException>(() =>
invocation.Proceed((call, next) => call.CreateValueReturn(Task.FromResult(1)), (outcome, again) =>
new ValueTask<ProceedOutcome>(outcome)));

Assert.Contains("AsyncRegistry.Register<", exception.Message);
}

[Fact]
public void RejectsAMissingCallback()
{
var invocation = MethodInvocation.Create(new object(), Sample(nameof(ISample.Read)));

Assert.Throws<ArgumentNullException>(() => 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<int> GetAsync();

ValueTask<int> ReadAsync();

ValueTask<Guid> IdAsync();

Task<T> Echo<T>(T value);
}
}
}
21 changes: 21 additions & 0 deletions src/Stunts.UnitTests/StuntGeneratorTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,27 @@ public static class Test { public static IDisposable Create() => Stunt.Of<IDispo
assembly, typeof(IDisposable), Array.Empty<Type>(), Array.Empty<object>()));
}

[Fact]
public void RegistersClosedAsyncAdapters()
{
var (diagnostics, compilation) = GetGeneratedOutput(@"
using System.Threading.Tasks;
using Stunts;
public interface IService
{
Task<int> GetAsync();
ValueTask<string> ReadAsync();
Task RunAsync();
}
public static class Test { public static IService Create() => Stunt.Of<IService>(); }");

Assert.Empty(diagnostics);
var text = string.Concat(compilation.SyntaxTrees.Select(tree => tree.ToString()));
Assert.Contains("AsyncRegistry.Register<int>()", text);
Assert.Contains("AsyncRegistry.Register<string>()", text);
Assert.Equal(2, System.Text.RegularExpressions.Regex.Matches(text, "AsyncRegistry.Register<").Count);
}

[Fact]
public void SkipsUncallableConstructorsInRegistration()
{
Expand Down
Loading
Loading