diff --git a/src/Orleans.CodeGenerator/LibraryTypes.cs b/src/Orleans.CodeGenerator/LibraryTypes.cs index 4ca83ae54b9..e1950431de7 100644 --- a/src/Orleans.CodeGenerator/LibraryTypes.cs +++ b/src/Orleans.CodeGenerator/LibraryTypes.cs @@ -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"); @@ -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; } diff --git a/src/Orleans.CodeGenerator/MetadataGenerator.cs b/src/Orleans.CodeGenerator/MetadataGenerator.cs index 04d418fa9cf..25ed3c1fff0 100644 --- a/src/Orleans.CodeGenerator/MetadataGenerator.cs +++ b/src/Orleans.CodeGenerator/MetadataGenerator.cs @@ -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"; @@ -23,6 +26,7 @@ private ClassDeclarationSyntax GenerateIncrementalMetadata() { var configParam = "config".ToIdentifierName(); var body = new List(); + var providerBody = new List(); var model = _metadataModel; var orderedProxyInterfaces = GetOrderedProxyInterfaces(model.ProxyInterfaces); var generatedInvokableActivatorMetadataNames = new HashSet( @@ -122,11 +126,30 @@ 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 body, IdentifierNameSyntax configParam) + private ClassDeclarationSyntax CreateMetadataClass( + List body, + List providerBody, + IdentifierNameSyntax configParam) { var configureMethod = MethodDeclaration(PredefinedType(Token(SyntaxKind.VoidKeyword)), "ConfigureInner") .AddModifiers(Token(SyntaxKind.ProtectedKeyword), Token(SyntaxKind.OverrideKeyword)) @@ -134,11 +157,26 @@ private ClassDeclarationSyntax CreateMetadataClass(List body, I 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( diff --git a/src/Orleans.CodeGenerator/Model/MetadataAggregateModelBuilder.cs b/src/Orleans.CodeGenerator/Model/MetadataAggregateModelBuilder.cs index d77db187c41..bdcebde8894 100644 --- a/src/Orleans.CodeGenerator/Model/MetadataAggregateModelBuilder.cs +++ b/src/Orleans.CodeGenerator/Model/MetadataAggregateModelBuilder.cs @@ -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) @@ -108,6 +120,7 @@ private static ReferenceAssemblyModel NormalizeReferenceAssemblyData(ReferenceAs ReferencedSerializableTypes: referencedSerializableTypes, ReferencedProxyInterfaces: referencedProxyInterfaces, RegisteredCodecs: registeredCodecs, + RegisteredProviders: normalizedRegisteredProviders, InterfaceImplementations: interfaceImplementations); } diff --git a/src/Orleans.CodeGenerator/Model/ReferenceAssemblyModel.cs b/src/Orleans.CodeGenerator/Model/ReferenceAssemblyModel.cs index 6d7afd561e5..505a1609eac 100644 --- a/src/Orleans.CodeGenerator/Model/ReferenceAssemblyModel.cs +++ b/src/Orleans.CodeGenerator/Model/ReferenceAssemblyModel.cs @@ -56,6 +56,11 @@ public InterfaceImplementationModel(TypeRef implementationType, SourceLocationMo public SourceLocationModel SourceLocation { get; } } +/// +/// Describes a provider registration. +/// +internal readonly record struct RegisteredProviderModel(string Target, string Kind, string Name, TypeRef Type); + /// /// Aggregated data extracted from referenced assemblies via [GenerateCodeForDeclaringAssembly] /// and [ApplicationPart] attributes. This model is produced by a CompilationProvider-based @@ -70,4 +75,5 @@ internal sealed record class ReferenceAssemblyModel( EquatableArray ReferencedSerializableTypes, EquatableArray ReferencedProxyInterfaces, EquatableArray RegisteredCodecs, + EquatableArray RegisteredProviders, EquatableArray InterfaceImplementations); diff --git a/src/Orleans.CodeGenerator/Model/ReferenceAssemblyModelExtractor.cs b/src/Orleans.CodeGenerator/Model/ReferenceAssemblyModelExtractor.cs index 9fb511dc35a..550ab840702 100644 --- a/src/Orleans.CodeGenerator/Model/ReferenceAssemblyModelExtractor.cs +++ b/src/Orleans.CodeGenerator/Model/ReferenceAssemblyModelExtractor.cs @@ -44,9 +44,13 @@ internal static ReferenceAssemblyModel ExtractReferenceAssemblyData( var referencedSerializableTypes = new HashSet(); var referencedProxyInterfaces = new HashSet(); var registeredCodecs = new HashSet(); + var registeredProviders = new Dictionary<(string Target, string Kind, string Name), RegisteredProviderModel>(); + var hasUnrepresentableProviderRegistration = false; var interfaceImplementations = new HashSet(); var diagnosticBuilder = ImmutableArray.CreateBuilder(); + CollectAssemblyAttributes(compilation.Assembly, includeApplicationParts: false); + foreach (var reference in compilation.References) { cancellationToken.ThrowIfCancellationRequested(); @@ -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) @@ -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(); @@ -194,6 +193,7 @@ internal static ReferenceAssemblyModel ExtractReferenceAssemblyData( ReferencedSerializableTypes: sortedReferencedSerializableTypes, ReferencedProxyInterfaces: sortedReferencedProxyInterfaces, RegisteredCodecs: sortedRegisteredCodecs, + RegisteredProviders: sortedRegisteredProviders, InterfaceImplementations: sortedInterfaceImplementations); void AddApplicationPart(string applicationPart) @@ -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( @@ -321,4 +402,3 @@ internal static RegisteredCodecModel ExtractRegisteredCodec(INamedTypeSymbol sym kind); } } - diff --git a/src/Orleans.CodeGenerator/ReferenceAssemblyDataProvider.cs b/src/Orleans.CodeGenerator/ReferenceAssemblyDataProvider.cs index 578abcdba75..0994a5595d1 100644 --- a/src/Orleans.CodeGenerator/ReferenceAssemblyDataProvider.cs +++ b/src/Orleans.CodeGenerator/ReferenceAssemblyDataProvider.cs @@ -42,6 +42,6 @@ internal static ReferenceAssemblyModel CreateEmptyReferenceAssemblyModel(string EquatableArray.Empty, EquatableArray.Empty, EquatableArray.Empty, + EquatableArray.Empty, EquatableArray.Empty); } - diff --git a/src/Orleans.Core/Core/DefaultClientServices.cs b/src/Orleans.Core/Core/DefaultClientServices.cs index ac17ed8cf8d..f54484de66c 100644 --- a/src/Orleans.Core/Core/DefaultClientServices.cs +++ b/src/Orleans.Core/Core/DefaultClientServices.cs @@ -1,4 +1,3 @@ -using System.Reflection; using Microsoft.AspNetCore.Connections; using Microsoft.Extensions.Configuration; using Microsoft.Extensions.DependencyInjection; @@ -233,21 +232,7 @@ static IProviderBuilder 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()) - { - 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) { diff --git a/src/Orleans.Core/Core/ProviderRegistrationResolver.cs b/src/Orleans.Core/Core/ProviderRegistrationResolver.cs new file mode 100644 index 00000000000..eb1c6ff3eb4 --- /dev/null +++ b/src/Orleans.Core/Core/ProviderRegistrationResolver.cs @@ -0,0 +1,68 @@ +using System; +using System.Collections.Generic; +using System.Reflection; +using Orleans.Serialization.Configuration; + +namespace Orleans +{ + /// + /// Resolves configuration-driven provider registrations from a set of assemblies. + /// + /// + /// Registrations are read from generated provider metadata (types implementing + /// 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 as a fallback. Only providers implementing + /// and marked as generated by OrleansCodeGen are activated; + /// other providers are never instantiated during discovery. + /// + internal static class ProviderRegistrationResolver + { + /// + /// Gets the providers registered for the specified , keyed by kind and name. + /// + /// The assemblies to scan for provider registrations. + /// The registration target to filter on, for example "Client" or "Silo". + public static Dictionary<(string Kind, string Name), Type> GetRegisteredProviders(IEnumerable 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(); + foreach (var attr in attrs) + { + if (typeof(IProviderMetadataProvider).IsAssignableFrom(attr.ProviderType) + && attr.ProviderType.GetCustomAttribute() 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()) + { + 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; + } + } +} diff --git a/src/Orleans.Runtime/Hosting/DefaultSiloServices.cs b/src/Orleans.Runtime/Hosting/DefaultSiloServices.cs index 0323580cad5..6ce3f76662a 100644 --- a/src/Orleans.Runtime/Hosting/DefaultSiloServices.cs +++ b/src/Orleans.Runtime/Hosting/DefaultSiloServices.cs @@ -508,21 +508,7 @@ static IProviderBuilder GetRequiredProvider(Dictionary<(string Kin } 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()) - { - if (string.Equals(attr.Target, "Silo")) - { - result[(attr.Kind, attr.Name)] = attr.Type; - } - } - } - - return result; - } + => ProviderRegistrationResolver.GetRegisteredProviders(ReferencedAssemblyProvider.GetRelevantAssemblies(), "Silo"); static void ApplySubsection(ISiloBuilder builder, IConfigurationSection cfg, Dictionary<(string Kind, string Name), Type> knownProviderTypes, string sectionName) { diff --git a/src/Orleans.Serialization/Configuration/ITypeManifestProvider.cs b/src/Orleans.Serialization/Configuration/ITypeManifestProvider.cs index 31f12422953..547f1ccc007 100644 --- a/src/Orleans.Serialization/Configuration/ITypeManifestProvider.cs +++ b/src/Orleans.Serialization/Configuration/ITypeManifestProvider.cs @@ -1,7 +1,21 @@ +using System; +using System.Collections.Generic; using Microsoft.Extensions.Options; namespace Orleans.Serialization.Configuration { + /// + /// Provides metadata for configuration-driven Orleans providers. + /// + public interface IProviderMetadataProvider + { + /// + /// Adds known providers to . + /// + /// The provider registrations, keyed by target, kind, and name. + void ConfigureProviders(IDictionary<(string Target, string Kind, string Name), Type> providers); + } + /// /// Provides type manifest information. /// diff --git a/src/Orleans.Serialization/Configuration/TypeManifestProviderAttribute.cs b/src/Orleans.Serialization/Configuration/TypeManifestProviderAttribute.cs index c44561f5745..e3bcdf5e784 100644 --- a/src/Orleans.Serialization/Configuration/TypeManifestProviderAttribute.cs +++ b/src/Orleans.Serialization/Configuration/TypeManifestProviderAttribute.cs @@ -1,4 +1,7 @@ using System; +#if NET5_0_OR_GREATER +using System.Diagnostics.CodeAnalysis; +#endif namespace Orleans.Serialization.Configuration { @@ -12,7 +15,11 @@ public sealed class TypeManifestProviderAttribute : Attribute /// Initializes a new instance of the class. /// /// The metadata provider type. - public TypeManifestProviderAttribute(Type providerType) + public TypeManifestProviderAttribute( +#if NET5_0_OR_GREATER + [DynamicallyAccessedMembers(DynamicallyAccessedMemberTypes.PublicConstructors)] +#endif + Type providerType) { if (providerType is null) { @@ -30,6 +37,9 @@ public TypeManifestProviderAttribute(Type providerType) /// /// Gets the manifest provider type. /// +#if NET5_0_OR_GREATER + [DynamicallyAccessedMembers(DynamicallyAccessedMemberTypes.PublicConstructors)] +#endif public Type ProviderType { get; } } } \ No newline at end of file diff --git a/src/api/Orleans.Serialization/Orleans.Serialization.cs b/src/api/Orleans.Serialization/Orleans.Serialization.cs index 8166790b309..3dbdc6ed09e 100644 --- a/src/api/Orleans.Serialization/Orleans.Serialization.cs +++ b/src/api/Orleans.Serialization/Orleans.Serialization.cs @@ -3335,6 +3335,11 @@ public void WriteField(ref Buffers.Writer writer, namespace Orleans.Serialization.Configuration { + public partial interface IProviderMetadataProvider + { + void ConfigureProviders(System.Collections.Generic.IDictionary<(string Target, string Kind, string Name), System.Type> providers); + } + public partial interface ITypeManifestProvider : Microsoft.Extensions.Options.IConfigureOptions { } diff --git a/test/Orleans.CodeGenerator.Tests/IncrementalCachingTests.cs b/test/Orleans.CodeGenerator.Tests/IncrementalCachingTests.cs index 1d186caf839..15bb3b8b6bf 100644 --- a/test/Orleans.CodeGenerator.Tests/IncrementalCachingTests.cs +++ b/test/Orleans.CodeGenerator.Tests/IncrementalCachingTests.cs @@ -111,6 +111,84 @@ public class UnrelatedClass AssertGeneratedSourcesIdentical(result1, result2); } + [Fact] + public async Task ProviderRegistration_UnrelatedChange_DoesNotTriggerMetadataRegeneration() + { + const string originalCode = """ + using Orleans; + + [assembly: RegisterProvider("Test", "Clustering", "Client", typeof(TestProject.Provider))] + + namespace TestProject; + + public sealed class Provider + { + } + """; + + const string modifiedCode = originalCode + """ + + public sealed class UnrelatedClass + { + public int Value { get; set; } + } + """; + + var compilation = await CreateCompilation(originalCode); + var newCompilation = ReplaceSource(compilation, modifiedCode); + var (result1, result2) = await RunTwice(compilation, newCompilation); + + AssertTrackedStepsCachedOrUnchanged( + result2, + OrleansSerializationSourceGenerator.ReferenceAssemblyDataTrackingName); + AssertTrackedStepsCached( + result2, + OrleansSerializationSourceGenerator.MetadataAggregateTrackingName, + OrleansSerializationSourceGenerator.MetadataOutputsTrackingName); + AssertGeneratedSourcesIdentical(result1, result2); + } + + [Fact] + public async Task AddedProviderRegistration_InvalidatesMetadataPipeline() + { + const string originalCode = """ + namespace TestProject; + + public sealed class Provider + { + } + """; + + const string modifiedCode = """ + using Orleans; + + [assembly: RegisterProvider("Test", "Clustering", "Client", typeof(TestProject.Provider))] + + namespace TestProject; + + public sealed class Provider + { + } + """; + + var compilation = await CreateCompilation(originalCode); + var newCompilation = ReplaceSource(compilation, modifiedCode); + var (result1, result2) = await RunTwice(compilation, newCompilation); + + AssertTrackedStepModifiedOrNew(result2, OrleansSerializationSourceGenerator.ReferenceAssemblyDataTrackingName); + AssertTrackedStepModifiedOrNew(result2, OrleansSerializationSourceGenerator.MetadataAggregateTrackingName); + AssertTrackedStepModifiedOrNew(result2, OrleansSerializationSourceGenerator.MetadataOutputsTrackingName); + + const string generatedProviderType = "typeof(global::TestProject.Provider)"; + var originalGeneratedSource = ConcatenateGeneratedSources(result1); + var modifiedGeneratedSource = ConcatenateGeneratedSources(result2); + Assert.DoesNotContain(generatedProviderType, originalGeneratedSource, StringComparison.Ordinal); + Assert.Contains(generatedProviderType, modifiedGeneratedSource, StringComparison.Ordinal); + Assert.Contains("global::Orleans.Serialization.Configuration.IProviderMetadataProvider", modifiedGeneratedSource, StringComparison.Ordinal); + Assert.Contains("public void ConfigureProviders(", modifiedGeneratedSource, StringComparison.Ordinal); + Assert.DoesNotContain("config.RegisteredProviders", modifiedGeneratedSource, StringComparison.Ordinal); + } + [Fact] public async Task AddingNewSerializableType_TriggersRegeneration() { @@ -838,6 +916,26 @@ private static void AssertTrackedStepsCachedOrUnchanged(GeneratorRunResult resul } } + private static void AssertTrackedStepsCached(GeneratorRunResult result, params string[] stepNames) + { + var trackedSteps = result.TrackedSteps; + Assert.NotEmpty(trackedSteps); + + foreach (var stepName in stepNames) + { + Assert.True(trackedSteps.TryGetValue(stepName, out var steps), $"Missing tracked step '{stepName}'."); + Assert.NotEmpty(steps); + + foreach (var step in steps) + { + foreach (var (_, reason) in step.Outputs) + { + Assert.Equal(IncrementalStepRunReason.Cached, reason); + } + } + } + } + private static void AssertTrackedStepModifiedOrNew(GeneratorRunResult result, string stepName) { var trackedSteps = result.TrackedSteps; diff --git a/test/Orleans.CodeGenerator.Tests/IncrementalModelEqualityTests.cs b/test/Orleans.CodeGenerator.Tests/IncrementalModelEqualityTests.cs index 80271f5cf7c..d0649fbaedb 100644 --- a/test/Orleans.CodeGenerator.Tests/IncrementalModelEqualityTests.cs +++ b/test/Orleans.CodeGenerator.Tests/IncrementalModelEqualityTests.cs @@ -246,6 +246,7 @@ private static ReferenceAssemblyModel CreateReferenceAssemblyModel( ImmutableArray referencedSerializableTypes = default, ImmutableArray referencedProxyInterfaces = default, ImmutableArray registeredCodecs = default, + ImmutableArray registeredProviders = default, ImmutableArray interfaceImplementations = default) => new( "TestAssembly", applicationParts, @@ -255,5 +256,6 @@ private static ReferenceAssemblyModel CreateReferenceAssemblyModel( referencedSerializableTypes, referencedProxyInterfaces, registeredCodecs, + registeredProviders, interfaceImplementations); } diff --git a/test/Orleans.CodeGenerator.Tests/ModelExtractorTests.cs b/test/Orleans.CodeGenerator.Tests/ModelExtractorTests.cs index 413834d55b8..b26b2247f2e 100644 --- a/test/Orleans.CodeGenerator.Tests/ModelExtractorTests.cs +++ b/test/Orleans.CodeGenerator.Tests/ModelExtractorTests.cs @@ -399,6 +399,11 @@ public async Task ExtractReferenceAssemblyData_CollectsCrossAssemblyMetadataAndD Assert.Contains("global::LibraryB.CopierType", registeredCodecTypes); Assert.Contains("global::LibraryB.ConverterType", registeredCodecTypes); Assert.Contains("global::LibraryB.SerializerType", registeredCodecTypes); + Assert.Equal( + [ + new RegisteredProviderModel("Client", "Clustering", "Consumer", new TypeRef("global::ConsumerProject.ConsumerMarker")), + ], + model.RegisteredProviders); Assert.Contains(model.InterfaceImplementations, implementation => implementation.ImplementationType.SyntaxString == "global::LibraryB.GeneratedInterfaceImplementation"); } @@ -415,6 +420,87 @@ public async Task ExtractReferenceAssemblyData_IsStableWhenReferenceOrderChanges Assert.Equal(modelA.GetHashCode(), modelB.GetHashCode()); } + [Fact] + public async Task ExtractReferenceAssemblyData_DoesNotEmitPartialMetadataForAliasOnlyProvider() + { + const string providerCode = """ + using Orleans; + + [assembly: RegisterProvider("Aliased", "Clustering", "Client", typeof(AliasedLibrary.Provider))] + + namespace AliasedLibrary; + + public sealed class Provider + { + } + """; + + var providerCompilation = await CreateCompilation(providerCode, "AliasedLibrary"); + var aliasedReference = providerCompilation.ToMetadataReference().WithAliases(["AliasOnly"]); + var consumerCompilation = await CreateCompilation( + """ + extern alias AliasOnly; + using Orleans; + + [assembly: RegisterProvider( + "Local", + "Clustering", + "Client", + typeof(ConsumerProject.LocalProvider))] + [assembly: RegisterProvider( + "Aliased", + "Clustering", + "Client", + typeof(ConsumerProject.ProviderContainer.NestedProvider))] + + namespace ConsumerProject; + + public sealed class LocalProvider + { + } + + public sealed class ProviderContainer + { + public sealed class NestedProvider + { + } + } + """, + "ConsumerProject", + aliasedReference); + + var model = ModelExtractor.ExtractReferenceAssemblyData(consumerCompilation, new CodeGeneratorOptions(), default); + + Assert.Empty(model.RegisteredProviders); + } + + [Fact] + public async Task ExtractReferenceAssemblyData_DoesNotEmitPartialMetadataForFileLocalProvider() + { + const string code = """ + using Orleans; + + [assembly: RegisterProvider("Local", "Clustering", "Client", typeof(TestProject.LocalProvider))] + [assembly: RegisterProvider("FileLocal", "Clustering", "Client", typeof(FileLocalProvider))] + + file sealed class FileLocalProvider + { + } + + namespace TestProject + { + public sealed class LocalProvider + { + } + } + """; + + var compilation = await CreateCompilation(code); + var model = ModelExtractor.ExtractReferenceAssemblyData(compilation, new CodeGeneratorOptions(), default); + + Assert.Empty(model.RegisteredProviders); + } + [Fact] public async Task ExtractProxyInterfaceModel_InheritedGenerateMethodSerializers_FallsBackToInheritedAttribute() { @@ -619,6 +705,11 @@ private static async Task CreateReferenceExtractionCompilatio using Orleans; using System.Threading.Tasks; + [assembly: RegisterProvider("LibraryB", "Clustering", "Client", typeof(LibraryB.BetaType))] + [assembly: RegisterProvider("Generic", "Clustering", "Client", typeof(LibraryB.GenericProvider))] + [assembly: RegisterProvider("Hidden", "Clustering", "Client", typeof(LibraryB.HiddenProvider))] + [assembly: RegisterProvider("Consumer", "Clustering", "Client", typeof(LibraryB.BetaType))] + namespace LibraryB; [Id(200)] @@ -658,6 +749,14 @@ public sealed class GeneratedInterfaceImplementation : IGeneratedInterface { public Task Ping() => Task.CompletedTask; } + + internal sealed class HiddenProvider + { + } + + public sealed class GenericProvider + { + } """; const string libraryACode = """ @@ -667,6 +766,7 @@ public sealed class GeneratedInterfaceImplementation : IGeneratedInterface [assembly: ApplicationPart("Zeta.Part")] [assembly: ApplicationPart("Alpha.Part")] [assembly: GenerateCodeForDeclaringAssembly(typeof(LibraryB.BetaType))] + [assembly: RegisterProvider("LibraryA", "Clustering", "Client", typeof(LibraryA.AlphaType))] namespace LibraryA; @@ -681,6 +781,7 @@ public sealed class AlphaType using Orleans; [assembly: GenerateCodeForDeclaringAssembly(typeof(LibraryA.AlphaType))] + [assembly: RegisterProvider("Consumer", "Clustering", "Client", typeof(ConsumerProject.ConsumerMarker))] namespace ConsumerProject; diff --git a/test/Orleans.Core.Tests/ProviderRegistrationResolverTests.cs b/test/Orleans.Core.Tests/ProviderRegistrationResolverTests.cs new file mode 100644 index 00000000000..604cc53e782 --- /dev/null +++ b/test/Orleans.Core.Tests/ProviderRegistrationResolverTests.cs @@ -0,0 +1,219 @@ +#nullable enable +using System; +using System.CodeDom.Compiler; +using System.Collections.Generic; +using System.Reflection; +using System.Reflection.Emit; +using System.Threading; +using Orleans; +using Orleans.Serialization.Configuration; +using Xunit; + +namespace NonSilo.Tests +{ + /// + /// Tests for provider discovery performed by , the shared seam used by + /// both client (DefaultClientServices) and silo (DefaultSiloServices) configuration. + /// These verify that generated provider metadata is consumed, that assemblies produced by older code generators + /// (which lack ) still work via the + /// fallback, and that custom type manifest providers are never activated during discovery outside of DI. + /// + [TestCategory("BVT")] + [TestCategory("Providers")] + public class ProviderRegistrationResolverTests + { + /// + /// Verifies that generated provider metadata (a type marked as generated by OrleansCodeGen which implements + /// ) is consumed by discovery for both the client and silo targets, + /// with each target only receiving the registrations declared for it. + /// + [Fact] + public void GeneratedProviderMetadata_IsConsumedByClientAndSiloConfiguration() + { + var assembly = CreateAssemblyWithAttributes( + TypeManifestProviderAttribute(typeof(GeneratedMetadataProvider))); + + var clientProviders = ProviderRegistrationResolver.GetRegisteredProviders(new[] { assembly }, "Client"); + var siloProviders = ProviderRegistrationResolver.GetRegisteredProviders(new[] { assembly }, "Silo"); + + // Client discovery receives only the client-targeted registration from the generated metadata. + Assert.Equal(typeof(ClientClusteringProviderBuilder), clientProviders[("Clustering", "UnitTestClustering")]); + Assert.False(clientProviders.ContainsKey(("GrainStorage", "UnitTestStorage"))); + + // Silo discovery receives only the silo-targeted registration from the same generated metadata. + Assert.Equal(typeof(SiloStorageProviderBuilder), siloProviders[("GrainStorage", "UnitTestStorage")]); + Assert.False(siloProviders.ContainsKey(("Clustering", "UnitTestClustering"))); + } + + /// + /// Verifies that an assembly produced by an older code generator - one whose generated type manifest provider + /// does not implement - still contributes its provider registrations + /// through the legacy fallback for both client and silo targets. + /// + [Fact] + public void LegacyAssemblyWithoutProviderMetadata_UsesRegisterProviderAttributeFallback() + { + OldGeneratorMetadataProvider.InstantiationCount = 0; + + var assembly = CreateAssemblyWithAttributes( + // An older generator still emits a generated type manifest provider, but it predates + // IProviderMetadataProvider so it does not implement it. + TypeManifestProviderAttribute(typeof(OldGeneratorMetadataProvider)), + RegisterProviderAttribute("LegacyClustering", "Clustering", "Client", typeof(LegacyClientProviderBuilder)), + RegisterProviderAttribute("LegacyStorage", "GrainStorage", "Silo", typeof(LegacySiloProviderBuilder))); + + var clientProviders = ProviderRegistrationResolver.GetRegisteredProviders(new[] { assembly }, "Client"); + var siloProviders = ProviderRegistrationResolver.GetRegisteredProviders(new[] { assembly }, "Silo"); + + // The client-targeted legacy registration is discovered via the RegisterProviderAttribute fallback. + Assert.Equal(typeof(LegacyClientProviderBuilder), clientProviders[("Clustering", "LegacyClustering")]); + Assert.False(clientProviders.ContainsKey(("GrainStorage", "LegacyStorage"))); + + // The silo-targeted legacy registration is discovered via the same fallback path. + Assert.Equal(typeof(LegacySiloProviderBuilder), siloProviders[("GrainStorage", "LegacyStorage")]); + Assert.False(siloProviders.ContainsKey(("Clustering", "LegacyClustering"))); + Assert.Equal(0, OldGeneratorMetadataProvider.InstantiationCount); + } + + /// + /// Verifies that a custom type manifest provider - one which is not marked as generated by OrleansCodeGen - is + /// never instantiated by provider discovery. Only OrleansCodeGen-generated providers are activated outside of + /// dependency injection; custom providers are activated by the serialization configuration pipeline via DI. + /// + [Fact] + public void CustomTypeManifestProvider_IsNotActivatedByProviderDiscovery() + { + CustomManifestProvider.InstantiationCount = 0; + + var assembly = CreateAssemblyWithAttributes( + TypeManifestProviderAttribute(typeof(CustomManifestProvider))); + + var clientProviders = ProviderRegistrationResolver.GetRegisteredProviders(new[] { assembly }, "Client"); + var siloProviders = ProviderRegistrationResolver.GetRegisteredProviders(new[] { assembly }, "Silo"); + + // The custom provider is not OrleansCodeGen-generated, so discovery must not activate it. + Assert.Equal(0, CustomManifestProvider.InstantiationCount); + + // And because it was never activated, it contributes no registrations to either target. + Assert.Empty(clientProviders); + Assert.Empty(siloProviders); + } + + /// + /// Verifies that when an assembly carries both modern generated provider metadata (implementing + /// ) and a legacy , the generated + /// metadata takes precedence and the legacy fallback is suppressed. This guards the behavior that the fallback + /// is only applied for assemblies built by older generators. + /// + [Fact] + public void GeneratedProviderMetadata_SuppressesRegisterProviderAttributeFallback() + { + var assembly = CreateAssemblyWithAttributes( + TypeManifestProviderAttribute(typeof(GeneratedMetadataProvider)), + // A stale RegisterProviderAttribute that must be ignored because generated metadata is present. + RegisterProviderAttribute("StaleClustering", "Clustering", "Client", typeof(LegacyClientProviderBuilder))); + + var clientProviders = ProviderRegistrationResolver.GetRegisteredProviders(new[] { assembly }, "Client"); + + // Only the generated metadata registration is present. + Assert.Equal(typeof(ClientClusteringProviderBuilder), clientProviders[("Clustering", "UnitTestClustering")]); + + // The legacy RegisterProviderAttribute registration is suppressed and never contributes. + Assert.False(clientProviders.ContainsKey(("Clustering", "StaleClustering"))); + Assert.Single(clientProviders); + } + + private static CustomAttributeBuilder TypeManifestProviderAttribute(Type providerType) + { + var ctor = typeof(TypeManifestProviderAttribute).GetConstructor(new[] { typeof(Type) })!; + return new CustomAttributeBuilder(ctor, new object[] { providerType }); + } + + private static CustomAttributeBuilder RegisterProviderAttribute(string name, string kind, string target, Type type) + { + var ctor = typeof(RegisterProviderAttribute).GetConstructor( + new[] { typeof(string), typeof(string), typeof(string), typeof(Type) })!; + return new CustomAttributeBuilder(ctor, new object[] { name, kind, target, type }); + } + + private static Assembly CreateAssemblyWithAttributes(params CustomAttributeBuilder[] attributes) + { + var name = new AssemblyName("ProviderDiscoveryTest_" + Guid.NewGuid().ToString("N")); + var assemblyBuilder = AssemblyBuilder.DefineDynamicAssembly(name, AssemblyBuilderAccess.Run); + foreach (var attribute in attributes) + { + assemblyBuilder.SetCustomAttribute(attribute); + } + + assemblyBuilder.DefineDynamicModule("Main"); + return assemblyBuilder; + } + + // Simulates a generated metadata provider: marked as generated by OrleansCodeGen and implementing + // IProviderMetadataProvider, so discovery activates it and reads its registrations. + [GeneratedCode("OrleansCodeGen", "1.0.0")] + public sealed class GeneratedMetadataProvider : TypeManifestProviderBase, IProviderMetadataProvider + { + public void ConfigureProviders(IDictionary<(string Target, string Kind, string Name), Type> providers) + { + providers[("Client", "Clustering", "UnitTestClustering")] = typeof(ClientClusteringProviderBuilder); + providers[("Silo", "GrainStorage", "UnitTestStorage")] = typeof(SiloStorageProviderBuilder); + } + + protected override void ConfigureInner(TypeManifestOptions options) + { + } + } + + // Simulates a generated metadata provider from an older generator: marked as generated by OrleansCodeGen but + // predating IProviderMetadataProvider, so discovery must fall back to RegisterProviderAttribute. + [GeneratedCode("OrleansCodeGen", "1.0.0")] + public sealed class OldGeneratorMetadataProvider : TypeManifestProviderBase + { + public static int InstantiationCount; + + public OldGeneratorMetadataProvider() + { + Interlocked.Increment(ref InstantiationCount); + } + + protected override void ConfigureInner(TypeManifestOptions options) + { + } + } + + // Simulates a user-authored custom type manifest provider. It is a valid ITypeManifestProvider (so it can be + // referenced by TypeManifestProviderAttribute) but is not OrleansCodeGen-generated, so provider discovery must + // never construct it. The instantiation counter proves whether the constructor ran. + public sealed class CustomManifestProvider : TypeManifestProviderBase + { + public static int InstantiationCount; + + public CustomManifestProvider() + { + Interlocked.Increment(ref InstantiationCount); + } + + protected override void ConfigureInner(TypeManifestOptions options) + { + } + } + + // Marker types used only as provider registration values; discovery stores the Type without activating it. + public sealed class ClientClusteringProviderBuilder + { + } + + public sealed class SiloStorageProviderBuilder + { + } + + public sealed class LegacyClientProviderBuilder + { + } + + public sealed class LegacySiloProviderBuilder + { + } + } +}