Skip to content

Commit 64b37f9

Browse files
authored
Add ValueTask/ValueTask<> support for NetCord.Services (#334)
* Add ValueTask/ValueTask<> support for NetCord.Services * Refactor * Refactor invocation * Refactor result resolver providers * Refactor * Fix task handling for result resolver provider * Refactor InteractionResultResolverProviderHelper.cs * Further refactor * Another refactor * Refactor CommandResultResolverProvider.cs * Improve result resolver type handling * Add interaction result resolvers tests * Add tests for CommandResultResolverProvider * Add tests for ValueTask/ValueTask<> to Task/Task<> fallback * Revert a config change in Program.cs * Use params
1 parent c873ef4 commit 64b37f9

7 files changed

Lines changed: 528 additions & 95 deletions

File tree

NetCord.Services/ApplicationCommands/SlashCommandParameter.cs

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,11 @@
1010

1111
namespace NetCord.Services.ApplicationCommands;
1212

13+
file static class Cache<TAutocompleteContext> where TAutocompleteContext : IAutocompleteInteractionContext
14+
{
15+
internal static readonly MethodInfo _getChoicesAsyncMethod = typeof(IAutocompleteProvider<TAutocompleteContext>).GetMethod(nameof(IAutocompleteProvider<>.GetChoicesAsync), BindingFlags.Instance | BindingFlags.Public)!;
16+
}
17+
1318
public class SlashCommandParameter<TContext> where TContext : IApplicationCommandContext
1419
{
1520
public SlashCommandTypeReader<TContext> TypeReader { get; }
@@ -126,10 +131,11 @@ internal void InitializeAutocomplete<TAutocompleteContext>(IServiceResolverProvi
126131
var option = Expression.Parameter(typeof(ApplicationCommandInteractionDataOption));
127132
var context = Expression.Parameter(typeof(TAutocompleteContext));
128133
var serviceProvider = Expression.Parameter(typeof(IServiceProvider));
129-
var getChoicesAsyncMethod = autocompleteProviderBaseType.GetMethod(nameof(IAutocompleteProvider<>.GetChoicesAsync), BindingFlags.Instance | BindingFlags.Public)!;
134+
130135
var call = Expression.Call(TypeHelper.GetCreateInstanceExpression(autocompleteProviderType, serviceProvider, serviceResolverProvider),
131-
getChoicesAsyncMethod,
136+
Cache<TAutocompleteContext>._getChoicesAsyncMethod,
132137
option, context);
138+
133139
var lambda = Expression.Lambda<Func<ApplicationCommandInteractionDataOption, TAutocompleteContext, IServiceProvider?, ValueTask<IEnumerable<ApplicationCommandOptionChoiceProperties>?>>>(call, option, context, serviceProvider);
134140
_invokeAutocompleteAsync = lambda.Compile();
135141
}

NetCord.Services/Commands/CommandResultResolverProvider.cs

Lines changed: 50 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -17,77 +17,98 @@ public bool TryGetResolver(Type type, [MaybeNullWhen(false)] out Func<object?, T
1717
{
1818
if (type == typeof(Task))
1919
{
20-
resolver = (result, context) => new(Unsafe.As<Task>(result!));
20+
resolver = static (result, context) => new(Unsafe.As<Task>(result!));
2121
return true;
2222
}
2323

24-
if (type == typeof(Task<ReplyMessageProperties>))
24+
if (type.IsGenericType && type.GetGenericTypeDefinition() == typeof(Task<>))
25+
return HandleTaskT(type, out resolver);
26+
27+
if (type == typeof(void))
2528
{
26-
resolver = async (result, context) =>
27-
{
28-
var messageProperties = await Unsafe.As<Task<ReplyMessageProperties>>(result!).ConfigureAwait(false);
29-
await context.Message.ReplyAsync(messageProperties).ConfigureAwait(false);
30-
};
29+
resolver = static (_, _) => default;
3130
return true;
3231
}
3332

34-
if (type == typeof(Task<MessageProperties>))
33+
if (type == typeof(ReplyMessageProperties))
3534
{
36-
resolver = async (result, context) =>
35+
resolver = static (result, context) =>
3736
{
38-
var messageProperties = await Unsafe.As<Task<MessageProperties>>(result!).ConfigureAwait(false);
39-
await context.Message.SendAsync(messageProperties).ConfigureAwait(false);
37+
var messageProperties = Unsafe.As<ReplyMessageProperties>(result!);
38+
return new(HandleReplyAsync(context, messageProperties));
4039
};
4140
return true;
4241
}
4342

44-
if (type == typeof(Task<string>))
43+
if (type == typeof(MessageProperties))
4544
{
46-
resolver = async (result, context) =>
45+
resolver = static (result, context) =>
4746
{
48-
var message = await Unsafe.As<Task<string>>(result!).ConfigureAwait(false);
49-
await context.Message.ReplyAsync(message).ConfigureAwait(false);
47+
var messageProperties = Unsafe.As<MessageProperties>(result!);
48+
return new(HandleMessageAsync(context, messageProperties));
5049
};
5150
return true;
5251
}
5352

54-
if (type == typeof(void))
53+
if (type == typeof(string))
5554
{
56-
resolver = (_, _) => default;
55+
resolver = static (result, context) =>
56+
{
57+
var message = Unsafe.As<string>(result!);
58+
return new(HandleReplyAsync(context, message));
59+
};
5760
return true;
5861
}
5962

60-
if (type == typeof(ReplyMessageProperties))
63+
resolver = null;
64+
return false;
65+
}
66+
67+
private static bool HandleTaskT(Type type, [MaybeNullWhen(false)] out Func<object?, TContext, ValueTask> resolver)
68+
{
69+
var genericArgument = type.GetGenericArguments()[0];
70+
71+
if (genericArgument == typeof(ReplyMessageProperties))
6172
{
62-
resolver = (result, context) =>
73+
resolver = static async (result, context) =>
6374
{
64-
var messageProperties = Unsafe.As<ReplyMessageProperties>(result!);
65-
return new(context.Message.ReplyAsync(messageProperties));
75+
var messageProperties = await Unsafe.As<Task<ReplyMessageProperties>>(result!).ConfigureAwait(false);
76+
await HandleReplyAsync(context, messageProperties).ConfigureAwait(false);
6677
};
6778
return true;
6879
}
6980

70-
if (type == typeof(MessageProperties))
81+
if (genericArgument == typeof(MessageProperties))
7182
{
72-
resolver = (result, context) =>
83+
resolver = static async (result, context) =>
7384
{
74-
var messageProperties = Unsafe.As<MessageProperties>(result!);
75-
return new(context.Message.SendAsync(messageProperties));
85+
var messageProperties = await Unsafe.As<Task<MessageProperties>>(result!).ConfigureAwait(false);
86+
await HandleMessageAsync(context, messageProperties).ConfigureAwait(false);
7687
};
7788
return true;
7889
}
7990

80-
if (type == typeof(string))
91+
if (genericArgument == typeof(string))
8192
{
82-
resolver = (result, context) =>
93+
resolver = static async (result, context) =>
8394
{
84-
var message = Unsafe.As<string>(result!);
85-
return new(context.Message.ReplyAsync(message));
95+
var content = await Unsafe.As<Task<string>>(result!).ConfigureAwait(false);
96+
await HandleReplyAsync(context, content).ConfigureAwait(false);
8697
};
8798
return true;
8899
}
89100

90101
resolver = null;
91102
return false;
92103
}
104+
105+
private static Task<RestMessage> HandleReplyAsync(TContext context, ReplyMessageProperties messageProperties)
106+
{
107+
return context.Message.ReplyAsync(messageProperties);
108+
}
109+
110+
private static Task<RestMessage> HandleMessageAsync(TContext context, MessageProperties messageProperties)
111+
{
112+
return context.Message.SendAsync(messageProperties);
113+
}
93114
}

NetCord.Services/Helpers/InvocationHelper.cs

Lines changed: 68 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -4,8 +4,16 @@
44

55
namespace NetCord.Services.Helpers;
66

7-
internal class InvocationHelper
7+
file static class InvocationHelper<TContext>
88
{
9+
internal static readonly MethodInfo _baseModuleSetContextMethod = typeof(IBaseModule<TContext>).GetMethod(nameof(IBaseModule<>.SetContext), BindingFlags.Instance | BindingFlags.NonPublic)!;
10+
}
11+
12+
internal static class InvocationHelper
13+
{
14+
private static readonly MethodInfo _valueTaskAsTaskMethod = typeof(ValueTask).GetMethod(nameof(ValueTask.AsTask), BindingFlags.Instance | BindingFlags.Public)!;
15+
private static readonly MethodInfo _valueTaskOfTAsTaskMethodUnbound = typeof(ValueTask<>).GetMethod(nameof(ValueTask<>.AsTask), BindingFlags.Instance | BindingFlags.Public)!;
16+
917
public static Func<object?[]?, TContext, IServiceProvider?, ValueTask> CreateModuleDelegate<TContext>(MethodInfo method, [DynamicallyAccessedMembers(DynamicallyAccessedMemberTypes.PublicConstructors)] Type declaringType, IEnumerable<Type> parameterTypes, IResultResolverProvider<TContext> resultResolverProvider, IServiceResolverProvider serviceResolverProvider)
1018
{
1119
var parameters = Expression.Parameter(typeof(object?[]));
@@ -19,16 +27,15 @@ internal class InvocationHelper
1927
var module = Expression.Variable(declaringType);
2028
instance = Expression.Block([module],
2129
Expression.Assign(module, TypeHelper.GetCreateInstanceExpression(declaringType, serviceProvider, serviceResolverProvider)),
22-
Expression.Call(module, typeof(IBaseModule<TContext>).GetMethod(nameof(IBaseModule<>.SetContext), BindingFlags.Instance | BindingFlags.NonPublic)!, context),
30+
Expression.Call(module, InvocationHelper<TContext>._baseModuleSetContextMethod, context),
2331
module);
2432
}
2533

2634
var call = Expression.Call(instance,
2735
method,
2836
parameterTypes.Select((p, i) => Expression.Convert(Expression.ArrayIndex(parameters, Expression.Constant(i, typeof(int))), p)));
2937

30-
var resolver = GetResolver(method, resultResolverProvider);
31-
var invokeResolver = GetInvokeResolverExpression(method, context, call, resolver);
38+
var invokeResolver = GetInvokeResolverExpression(method, context, call, resultResolverProvider);
3239

3340
var lambda = Expression.Lambda<Func<object?[]?, TContext, IServiceProvider?, ValueTask>>(invokeResolver, parameters, context, serviceProvider);
3441
return lambda.Compile();
@@ -49,36 +56,76 @@ internal class InvocationHelper
4956

5057
var method = handler.Method;
5158

52-
var resolver = GetResolver(method, resultResolverProvider);
53-
var invokeResolver = GetInvokeResolverExpression(method, context, invoke, resolver);
59+
var invokeResolver = GetInvokeResolverExpression(method, context, invoke, resultResolverProvider);
5460

5561
var lambda = Expression.Lambda<Func<object?[]?, TContext, IServiceProvider?, ValueTask>>(invokeResolver, parameters, context, serviceProvider);
5662
return lambda.Compile();
5763
}
5864

59-
private static Func<object?, TContext, ValueTask> GetResolver<TContext>(MethodInfo method, IResultResolverProvider<TContext> resultResolverProvider)
65+
[DoesNotReturn]
66+
private static void ThrowResolverTypeNotSupportedException<TContext>(MethodInfo method, Type returnType, IResultResolverProvider<TContext> resultResolverProvider)
6067
{
61-
var type = method.ReturnType;
62-
if (resultResolverProvider.TryGetResolver(type, out var resolver))
63-
return resolver;
68+
throw new InvalidDefinitionException($"The return type '{returnType}' is not supported by '{resultResolverProvider.GetType()}'.", method);
69+
}
6470

65-
throw new InvalidDefinitionException($"The return type '{type}' is not supported by '{resultResolverProvider.GetType()}'.", method);
71+
[DoesNotReturn]
72+
private static void ThrowResolverBothTypesNotSupportedException<TContext>(MethodInfo method, Type returnType1, Type returnType2, IResultResolverProvider<TContext> resultResolverProvider)
73+
{
74+
throw new InvalidDefinitionException($"The return types '{returnType1}' and '{returnType2}' are not supported by '{resultResolverProvider.GetType()}'.", method);
6675
}
6776

68-
private static Expression GetInvokeResolverExpression<TContext>(MethodInfo method, ParameterExpression context, Expression call, Func<object?, TContext, ValueTask> resolver)
77+
private static Expression GetInvokeResolverExpression<TContext>(MethodInfo method, ParameterExpression context, Expression call, IResultResolverProvider<TContext> resultResolverProvider)
6978
{
70-
if (method.ReturnType == typeof(void))
79+
var returnType = method.ReturnType;
80+
81+
if (resultResolverProvider.TryGetResolver(returnType, out var resolver))
7182
{
72-
return Expression.Block(call,
73-
Expression.Invoke(Expression.Constant(resolver),
74-
Expression.Constant(null, typeof(object)),
75-
context));
83+
if (returnType == typeof(void))
84+
return Expression.Block(call,
85+
Expression.Invoke(Expression.Constant(resolver, typeof(Func<object?, TContext, ValueTask>)),
86+
Expression.Constant(null, typeof(object)),
87+
context));
88+
89+
return CreateResultResolverCallExpression(context, call, resolver);
7690
}
77-
else
91+
92+
return GetAlternativeInvokeResolverExpression(method, returnType, context, call, resultResolverProvider);
93+
}
94+
95+
private static InvocationExpression GetAlternativeInvokeResolverExpression<TContext>(MethodInfo method, Type returnType, ParameterExpression context, Expression call, IResultResolverProvider<TContext> resultResolverProvider)
96+
{
97+
Func<object?, TContext, ValueTask>? resolver = null;
98+
99+
if (returnType == typeof(ValueTask))
78100
{
79-
return Expression.Invoke(Expression.Constant(resolver),
80-
Expression.Convert(call, typeof(object)),
81-
context);
101+
var alternativeReturnType = typeof(Task);
102+
103+
if (!resultResolverProvider.TryGetResolver(alternativeReturnType, out resolver))
104+
ThrowResolverBothTypesNotSupportedException(method, returnType, alternativeReturnType, resultResolverProvider);
105+
106+
call = Expression.Call(call, _valueTaskAsTaskMethod);
82107
}
108+
else if (returnType.IsGenericType && returnType.GetGenericTypeDefinition() == typeof(ValueTask<>))
109+
{
110+
var asTaskMethod = (MethodInfo)returnType.GetMemberWithSameMetadataDefinitionAs(_valueTaskOfTAsTaskMethodUnbound);
111+
112+
var alternativeReturnType = asTaskMethod.ReturnType;
113+
114+
if (!resultResolverProvider.TryGetResolver(alternativeReturnType, out resolver))
115+
ThrowResolverBothTypesNotSupportedException(method, returnType, alternativeReturnType, resultResolverProvider);
116+
117+
call = Expression.Call(call, asTaskMethod);
118+
}
119+
else
120+
ThrowResolverTypeNotSupportedException(method, returnType, resultResolverProvider);
121+
122+
return CreateResultResolverCallExpression(context, call, resolver);
123+
}
124+
125+
private static InvocationExpression CreateResultResolverCallExpression<TContext>(ParameterExpression context, Expression call, Func<object?, TContext, ValueTask> resolver)
126+
{
127+
return Expression.Invoke(Expression.Constant(resolver, typeof(Func<object?, TContext, ValueTask>)),
128+
Expression.Convert(call, typeof(object)),
129+
context);
83130
}
84131
}

0 commit comments

Comments
 (0)