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
2 changes: 2 additions & 0 deletions src/Orleans.CodeGenerator/LibraryTypes.cs
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ private LibraryTypes(Compilation compilation, CodeGeneratorOptions options)
RegisterActivatorAttribute = Type("Orleans.RegisterActivatorAttribute");
RegisterConverterAttribute = Type("Orleans.RegisterConverterAttribute");
RegisterCopierAttribute = Type("Orleans.RegisterCopierAttribute");
RegisterProviderAttribute = Type("Orleans.RegisterProviderAttribute");
UseActivatorAttribute = Type("Orleans.UseActivatorAttribute");
SuppressReferenceTrackingAttribute = Type("Orleans.SuppressReferenceTrackingAttribute");
OmitDefaultMemberValuesAttribute = Type("Orleans.OmitDefaultMemberValuesAttribute");
Expand Down Expand Up @@ -251,6 +252,7 @@ INamedTypeSymbol Type(string metadataName)
public INamedTypeSymbol ResponseTimeoutAttribute { get; private set; }
public INamedTypeSymbol RegisterConverterAttribute { get; private set; }
public INamedTypeSymbol RegisterActivatorAttribute { get; private set; }
public INamedTypeSymbol RegisterProviderAttribute { get; private set; }
public INamedTypeSymbol UseActivatorAttribute { get; private set; }
public INamedTypeSymbol SuppressReferenceTrackingAttribute { get; private set; }
public INamedTypeSymbol OmitDefaultMemberValuesAttribute { get; private set; }
Expand Down
44 changes: 41 additions & 3 deletions src/Orleans.CodeGenerator/MetadataGenerator.cs
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,9 @@ internal class MetadataGenerator(MetadataAggregateModel metadataModel, string as
{
private static readonly TypeSyntax TypeManifestOptionsType = ParseTypeName("global::Orleans.Serialization.Configuration.TypeManifestOptions");
private static readonly TypeSyntax TypeManifestProviderBaseType = ParseTypeName("global::Orleans.Serialization.Configuration.TypeManifestProviderBase");
private static readonly TypeSyntax ProviderMetadataProviderType = ParseTypeName("global::Orleans.Serialization.Configuration.IProviderMetadataProvider");
private static readonly TypeSyntax ProviderDictionaryType = ParseTypeName(
"global::System.Collections.Generic.IDictionary<(string Target, string Kind, string Name), global::System.Type>");

private readonly MetadataAggregateModel _metadataModel = metadataModel;
private readonly string _assemblyName = assemblyName ?? "Assembly";
Expand All @@ -23,6 +26,7 @@ private ClassDeclarationSyntax GenerateIncrementalMetadata()
{
var configParam = "config".ToIdentifierName();
var body = new List<StatementSyntax>();
var providerBody = new List<StatementSyntax>();
var model = _metadataModel;
var orderedProxyInterfaces = GetOrderedProxyInterfaces(model.ProxyInterfaces);
var generatedInvokableActivatorMetadataNames = new HashSet<string>(
Expand Down Expand Up @@ -122,23 +126,57 @@ private ClassDeclarationSyntax GenerateIncrementalMetadata()
])))));
}

foreach (var provider in model.ReferenceAssemblyData.RegisteredProviders)
{
var key = TupleExpression(SeparatedList(
[
Argument(LiteralExpression(SyntaxKind.StringLiteralExpression, Literal(provider.Target))),
Argument(LiteralExpression(SyntaxKind.StringLiteralExpression, Literal(provider.Kind))),
Argument(LiteralExpression(SyntaxKind.StringLiteralExpression, Literal(provider.Name))),
]));
var registeredProvider = ElementAccessExpression(IdentifierName("providers"))
.WithArgumentList(BracketedArgumentList(SingletonSeparatedList(Argument(key))));
providerBody.Add(ExpressionStatement(AssignmentExpression(
SyntaxKind.SimpleAssignmentExpression,
registeredProvider,
TypeOfExpression(provider.Type.ToTypeSyntax()))));
}

AddCompoundTypeAliases(configParam, body, generatedInvokables);
return CreateMetadataClass(body, configParam);
return CreateMetadataClass(body, providerBody, configParam);
}

private ClassDeclarationSyntax CreateMetadataClass(List<StatementSyntax> body, IdentifierNameSyntax configParam)
private ClassDeclarationSyntax CreateMetadataClass(
List<StatementSyntax> body,
List<StatementSyntax> providerBody,
IdentifierNameSyntax configParam)
{
var configureMethod = MethodDeclaration(PredefinedType(Token(SyntaxKind.VoidKeyword)), "ConfigureInner")
.AddModifiers(Token(SyntaxKind.ProtectedKeyword), Token(SyntaxKind.OverrideKeyword))
.AddParameterListParameters(
Parameter(configParam.Identifier).WithType(TypeManifestOptionsType))
.AddBodyStatements([.. body]);

return ClassDeclaration("Metadata_" + SyntaxGeneration.Identifier.SanitizeIdentifierName(_assemblyName))
var result = ClassDeclaration("Metadata_" + SyntaxGeneration.Identifier.SanitizeIdentifierName(_assemblyName))
.AddBaseListTypes(SimpleBaseType(TypeManifestProviderBaseType))
.AddModifiers(Token(SyntaxKind.InternalKeyword), Token(SyntaxKind.SealedKeyword))
.AddAttributeLists(GeneratedCodeUtilities.GetGeneratedCodeAttributes())
.AddMembers(configureMethod);

if (providerBody.Count > 0)
{
var configureProvidersMethod = MethodDeclaration(PredefinedType(Token(SyntaxKind.VoidKeyword)), "ConfigureProviders")
.AddModifiers(Token(SyntaxKind.PublicKeyword))
.AddParameterListParameters(
Parameter(Identifier("providers")).WithType(ProviderDictionaryType))
.AddBodyStatements([.. providerBody]);

result = result
.AddBaseListTypes(SimpleBaseType(ProviderMetadataProviderType))
.AddMembers(configureProvidersMethod);
}

return result;
}

private void AddCompoundTypeAliases(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,18 @@ private static ReferenceAssemblyModel NormalizeReferenceAssemblyData(ReferenceAs
.ThenBy(static entry => entry.Kind)
.ToImmutableArray();

var registeredProviders = new Dictionary<(string Target, string Kind, string Name), RegisteredProviderModel>();
foreach (var provider in referenceData.RegisteredProviders)
{
registeredProviders[(provider.Target, provider.Kind, provider.Name)] = provider;
}

var normalizedRegisteredProviders = registeredProviders.Values
.OrderBy(static entry => entry.Target, StringComparer.Ordinal)
.ThenBy(static entry => entry.Kind, StringComparer.Ordinal)
.ThenBy(static entry => entry.Name, StringComparer.Ordinal)
.ToImmutableArray();

var interfaceImplementations = referenceData.InterfaceImplementations
.Distinct()
.OrderBy(static entry => entry.ImplementationType.SyntaxString, StringComparer.Ordinal)
Expand All @@ -108,6 +120,7 @@ private static ReferenceAssemblyModel NormalizeReferenceAssemblyData(ReferenceAs
ReferencedSerializableTypes: referencedSerializableTypes,
ReferencedProxyInterfaces: referencedProxyInterfaces,
RegisteredCodecs: registeredCodecs,
RegisteredProviders: normalizedRegisteredProviders,
InterfaceImplementations: interfaceImplementations);
}

Expand Down
6 changes: 6 additions & 0 deletions src/Orleans.CodeGenerator/Model/ReferenceAssemblyModel.cs
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,11 @@ public InterfaceImplementationModel(TypeRef implementationType, SourceLocationMo
public SourceLocationModel SourceLocation { get; }
}

/// <summary>
/// Describes a provider registration.
/// </summary>
internal readonly record struct RegisteredProviderModel(string Target, string Kind, string Name, TypeRef Type);

/// <summary>
/// Aggregated data extracted from referenced assemblies via <c>[GenerateCodeForDeclaringAssembly]</c>
/// and <c>[ApplicationPart]</c> attributes. This model is produced by a <c>CompilationProvider</c>-based
Expand All @@ -70,4 +75,5 @@ internal sealed record class ReferenceAssemblyModel(
EquatableArray<SerializableTypeModel> ReferencedSerializableTypes,
EquatableArray<ProxyInterfaceModel> ReferencedProxyInterfaces,
EquatableArray<RegisteredCodecModel> RegisteredCodecs,
EquatableArray<RegisteredProviderModel> RegisteredProviders,
EquatableArray<InterfaceImplementationModel> InterfaceImplementations);
110 changes: 95 additions & 15 deletions src/Orleans.CodeGenerator/Model/ReferenceAssemblyModelExtractor.cs
Original file line number Diff line number Diff line change
Expand Up @@ -44,9 +44,13 @@ internal static ReferenceAssemblyModel ExtractReferenceAssemblyData(
var referencedSerializableTypes = new HashSet<SerializableTypeModel>();
var referencedProxyInterfaces = new HashSet<ProxyInterfaceModel>();
var registeredCodecs = new HashSet<RegisteredCodecModel>();
var registeredProviders = new Dictionary<(string Target, string Kind, string Name), RegisteredProviderModel>();
var hasUnrepresentableProviderRegistration = false;
var interfaceImplementations = new HashSet<InterfaceImplementationModel>();
var diagnosticBuilder = ImmutableArray.CreateBuilder<Diagnostic>();

CollectAssemblyAttributes(compilation.Assembly, includeApplicationParts: false);

foreach (var reference in compilation.References)
{
cancellationToken.ThrowIfCancellationRequested();
Expand All @@ -56,20 +60,7 @@ internal static ReferenceAssemblyModel ExtractReferenceAssemblyData(
continue;
}

if (!asm.GetAttributes(libraryTypes.ApplicationPartAttribute, out var attrs))
{
continue;
}

AddApplicationPart(asm.MetadataName);
foreach (var attr in attrs)
{
if (attr.ConstructorArguments.Length > 0
&& attr.ConstructorArguments[0].Value is string partName)
{
AddApplicationPart(partName);
}
}
CollectAssemblyAttributes(asm, includeApplicationParts: true);
}

foreach (var asm in assembliesToExamine)
Expand Down Expand Up @@ -179,6 +170,14 @@ internal static ReferenceAssemblyModel ExtractReferenceAssemblyData(
.ThenBy(static entry => entry.Kind)
.ToImmutableArray();

var sortedRegisteredProviders = hasUnrepresentableProviderRegistration
? []
: registeredProviders.Values
.OrderBy(static entry => entry.Target, StringComparer.Ordinal)
.ThenBy(static entry => entry.Kind, StringComparer.Ordinal)
.ThenBy(static entry => entry.Name, StringComparer.Ordinal)
.ToImmutableArray();

var sortedInterfaceImplementations = interfaceImplementations
.OrderBy(static entry => entry.ImplementationType.SyntaxString, StringComparer.Ordinal)
.ToImmutableArray();
Expand All @@ -194,6 +193,7 @@ internal static ReferenceAssemblyModel ExtractReferenceAssemblyData(
ReferencedSerializableTypes: sortedReferencedSerializableTypes,
ReferencedProxyInterfaces: sortedReferencedProxyInterfaces,
RegisteredCodecs: sortedRegisteredCodecs,
RegisteredProviders: sortedRegisteredProviders,
InterfaceImplementations: sortedInterfaceImplementations);

void AddApplicationPart(string applicationPart)
Expand All @@ -203,6 +203,87 @@ void AddApplicationPart(string applicationPart)
applicationParts.Add(applicationPart);
}
}

void CollectAssemblyAttributes(IAssemblySymbol assembly, bool includeApplicationParts)
{
var hasApplicationPartAttribute = false;

foreach (var attribute in assembly.GetAttributes())
{
cancellationToken.ThrowIfCancellationRequested();

if (includeApplicationParts
&& SymbolEqualityComparer.Default.Equals(attribute.AttributeClass, libraryTypes.ApplicationPartAttribute))
{
if (!hasApplicationPartAttribute)
{
AddApplicationPart(assembly.MetadataName);
hasApplicationPartAttribute = true;
}

if (attribute.ConstructorArguments.Length > 0
&& attribute.ConstructorArguments[0].Value is string partName)
{
AddApplicationPart(partName);
}

continue;
}

if (!SymbolEqualityComparer.Default.Equals(assembly, compilation.Assembly)
|| !SymbolEqualityComparer.Default.Equals(attribute.AttributeClass, libraryTypes.RegisterProviderAttribute))
{
continue;
}

var arguments = attribute.ConstructorArguments;
if (arguments.Length < 4
|| arguments[0].Value is not string name
|| arguments[1].Value is not string kind
|| arguments[2].Value is not string target
|| arguments[3].Value is not INamedTypeSymbol type
|| !compilation.IsSymbolAccessibleWithin(type, compilation.Assembly)
|| !CanReferenceProviderType(type))
{
hasUnrepresentableProviderRegistration = true;
continue;
}

var provider = new RegisteredProviderModel(
target,
kind,
name,
new TypeRef(type.ToDisplayString(SymbolDisplayFormat.FullyQualifiedFormat)));
var key = (target, kind, name);
registeredProviders[key] = provider;
}
}

bool CanReferenceProviderType(ITypeSymbol type)
{
return type switch
{
INamedTypeSymbol namedType => CanReferenceNamedType(namedType)
&& (namedType.ContainingType is null || CanReferenceProviderType(namedType.ContainingType))
&& namedType.TypeArguments.All(CanReferenceProviderType),
IArrayTypeSymbol arrayType => CanReferenceProviderType(arrayType.ElementType),
IPointerTypeSymbol pointerType => CanReferenceProviderType(pointerType.PointedAtType),
_ => true,
};
}

bool CanReferenceNamedType(INamedTypeSymbol type)
{
if (type.IsFileLocal)
{
return false;
}

var originalDefinition = type.OriginalDefinition;
var metadataName = TypeMetadataIdentity.Create(originalDefinition).MetadataName;
return compilation.GetTypeByMetadataName(metadataName) is { } resolvedType
&& SymbolEqualityComparer.Default.Equals(resolvedType, originalDefinition);
}
}

private static void ComputeAssembliesToExamine(
Expand Down Expand Up @@ -321,4 +402,3 @@ internal static RegisteredCodecModel ExtractRegisteredCodec(INamedTypeSymbol sym
kind);
}
}

Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,6 @@ internal static ReferenceAssemblyModel CreateEmptyReferenceAssemblyModel(string
EquatableArray<SerializableTypeModel>.Empty,
EquatableArray<ProxyInterfaceModel>.Empty,
EquatableArray<RegisteredCodecModel>.Empty,
EquatableArray<RegisteredProviderModel>.Empty,
EquatableArray<InterfaceImplementationModel>.Empty);
}

17 changes: 1 addition & 16 deletions src/Orleans.Core/Core/DefaultClientServices.cs
Original file line number Diff line number Diff line change
@@ -1,4 +1,3 @@
using System.Reflection;
using Microsoft.AspNetCore.Connections;
using Microsoft.Extensions.Configuration;
using Microsoft.Extensions.DependencyInjection;
Expand Down Expand Up @@ -233,21 +232,7 @@ static IProviderBuilder<IClientBuilder> GetRequiredProvider(Dictionary<(string K
}

static Dictionary<(string Kind, string Name), Type> GetRegisteredProviders()
{
var result = new Dictionary<(string, string), Type>();
foreach (var asm in ReferencedAssemblyProvider.GetRelevantAssemblies())
{
foreach (var attr in asm.GetCustomAttributes<RegisterProviderAttribute>())
{
if (string.Equals(attr.Target, "Client"))
{
result[(attr.Kind, attr.Name)] = attr.Type;
}
}
}

return result;
}
=> ProviderRegistrationResolver.GetRegisteredProviders(ReferencedAssemblyProvider.GetRelevantAssemblies(), "Client");

static void ApplySubsection(IClientBuilder builder, IConfigurationSection cfg, Dictionary<(string Kind, string Name), Type> knownProviderTypes, string sectionName)
{
Expand Down
68 changes: 68 additions & 0 deletions src/Orleans.Core/Core/ProviderRegistrationResolver.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,68 @@
using System;
using System.Collections.Generic;
using System.Reflection;
using Orleans.Serialization.Configuration;

namespace Orleans
{
/// <summary>
/// Resolves configuration-driven provider registrations from a set of assemblies.
/// </summary>
/// <remarks>
/// Registrations are read from generated provider metadata (types implementing
/// <see cref="IProviderMetadataProvider"/> which are emitted by the Orleans code generator). Provider
/// assemblies produced by an older code generator do not carry that metadata, so their registrations are
/// read from <see cref="RegisterProviderAttribute"/> as a fallback. Only providers implementing
/// <see cref="IProviderMetadataProvider"/> and marked as generated by <c>OrleansCodeGen</c> are activated;
/// other <see cref="TypeManifestProviderAttribute"/> providers are never instantiated during discovery.
/// </remarks>
internal static class ProviderRegistrationResolver
{
/// <summary>
/// Gets the providers registered for the specified <paramref name="target"/>, keyed by kind and name.
/// </summary>
/// <param name="assemblies">The assemblies to scan for provider registrations.</param>
/// <param name="target">The registration target to filter on, for example <c>"Client"</c> or <c>"Silo"</c>.</param>
public static Dictionary<(string Kind, string Name), Type> GetRegisteredProviders(IEnumerable<Assembly> assemblies, string target)
{
var result = new Dictionary<(string, string), Type>();
var registeredProviders = new Dictionary<(string Target, string Kind, string Name), Type>();

// Collect providers from generated metadata.
foreach (var asm in assemblies)
{
var hasGeneratedProviderMetadata = false;
var attrs = asm.GetCustomAttributes<TypeManifestProviderAttribute>();
foreach (var attr in attrs)
{
if (typeof(IProviderMetadataProvider).IsAssignableFrom(attr.ProviderType)
&& attr.ProviderType.GetCustomAttribute<System.CodeDom.Compiler.GeneratedCodeAttribute>() is { Tool: "OrleansCodeGen" }
&& Activator.CreateInstance(attr.ProviderType) is IProviderMetadataProvider provider)
{
provider.ConfigureProviders(registeredProviders);
hasGeneratedProviderMetadata = true;
}
}

// Provider assemblies built with an older generator do not include provider metadata.
if (!hasGeneratedProviderMetadata && asm.IsDefined(typeof(RegisterProviderAttribute)))
{
foreach (var attr in asm.GetCustomAttributes<RegisterProviderAttribute>())
{
registeredProviders[(attr.Target, attr.Kind, attr.Name)] = attr.Type;
}
}
}

foreach (var kvp in registeredProviders)
{
if (string.Equals(kvp.Key.Target, target, StringComparison.Ordinal))
{
result[(kvp.Key.Kind, kvp.Key.Name)] = kvp.Value;
}
}

return result;
}
}
}
Loading