Sitelet https://github.com/MessagePack-CSharp/MessagePack-CSharp/pull/1968/files
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
26 changes: 16 additions & 10 deletions src/MessagePack.SourceGenerator/CodeAnalysis/AnalyzerOptions.cs
Original file line number Diff line number Diff line change
Expand Up @@ -178,12 +178,11 @@ public record ResolverOptions
/// <summary>
/// Gets the name to use for the resolver.
/// </summary>
public string Name { get; init; } = "GeneratedMessagePackResolver";

/// <summary>
/// Gets the namespace the source generated resolver will be emitted into.
/// </summary>
public string? Namespace { get; init; } = "MessagePack";
public QualifiedNamedTypeName Name { get; init; } = new(TypeKind.Class)
{
Container = new NamespaceTypeContainer("MessagePack"),
Name = "GeneratedMessagePackResolver",
};
}

/// <summary>
Expand Down Expand Up @@ -211,11 +210,17 @@ public record GeneratorOptions
/// <param name="FormattableTypes">The type arguments that appear in each implemented <c>IMessagePackFormatter</c> interface. When generic, these should be the full name of their type definitions.</param>
public record FormatterDescriptor(QualifiedNamedTypeName Name, string? InstanceProvidingMember, QualifiedTypeName InstanceTypeName, ImmutableHashSet<FormattableType> FormattableTypes)
{
/// <summary>
/// Creates a descriptor for a formatter, if the given type implements at least one <c>IMessagePackFormatter</c> interface.
/// </summary>
/// <param name="type">The type symbol for the formatter.</param>
/// <param name="formatter">Receives the formatter descriptor, if applicable.</param>
/// <returns><see langword="true"/> if <paramref name="type"/> represents a formatter.</returns>
public static bool TryCreate(INamedTypeSymbol type, [NotNullWhen(true)] out FormatterDescriptor? formatter)
{
var formattedTypes =
AnalyzerUtilities.SearchTypeForFormatterImplementations(type)
.Select(i => new FormattableType(i))
.Select(i => new FormattableType(i, type.ContainingAssembly))
.ToImmutableHashSet();
if (formattedTypes.IsEmpty)
{
Expand Down Expand Up @@ -268,10 +273,11 @@ public virtual bool Equals(FormatterDescriptor? other)
/// Describes a formattable type.
/// </summary>
/// <param name="Name">The name of the formattable type.</param>
public record FormattableType(QualifiedTypeName Name)
/// <param name="IsFormatterInSameAssembly"><see langword="true" /> if the formatter and the formatted types are declared in the same assembly.</param>
public record FormattableType(QualifiedTypeName Name, bool IsFormatterInSameAssembly)
{
public FormattableType(ITypeSymbol type)
: this(QualifiedTypeName.Create(type))
public FormattableType(ITypeSymbol type, IAssemblySymbol? formatterAssembly)
: this(QualifiedTypeName.Create(type), SymbolEqualityComparer.Default.Equals(type.ContainingAssembly, formatterAssembly))
{
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@ public record CustomFormatterRegisterInfo : ResolverRegisterInfo
{
public required FormatterDescriptor CustomFormatter { get; init; }

public required FormattableType FormattableDataType { get; init; }

public override string GetFormatterInstanceForResolver()
{
return this.CustomFormatter.InstanceProvidingMember == ".ctor"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -223,10 +223,13 @@ public abstract record TypeContainer : IComparable<TypeContainer?>
public abstract int CompareTo(TypeContainer? other);
}

[DebuggerDisplay($"{{{nameof(NestingType)},nq}}+")]
public record NestingTypeContainer(QualifiedNamedTypeName NestingType) : TypeContainer
{
public static NestingTypeContainer? From(QualifiedNamedTypeName? nesting) => nesting is null ? null : new(nesting);

public override string ToString() => $"{this.NestingType}+";

public override int CompareTo(TypeContainer? other)
{
return other is NestingTypeContainer nestingOther ? this.NestingType.CompareTo(nestingOther.NestingType) :
Expand All @@ -235,10 +238,13 @@ public override int CompareTo(TypeContainer? other)
}
}

[DebuggerDisplay($"{{{nameof(Namespace)},nq}}.")]
public record NamespaceTypeContainer(string Namespace) : TypeContainer
{
public static NamespaceTypeContainer? From(string? ns) => ns is null ? null : new(ns);

public override string ToString() => $"{this.Namespace}.";

public override int CompareTo(TypeContainer? other)
{
return other is NamespaceTypeContainer nsOther ? this.Namespace.CompareTo(nsOther.Namespace) :
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -119,17 +119,11 @@ public static ResolverRegisterInfo CreateArray(IArrayTypeSymbol dataType, Resolv
};
}

QualifiedNamedTypeName generatedResolverName = new(TypeKind.Class)
{
Container = NamespaceTypeContainer.From(resolverOptions.Namespace),
Name = resolverOptions.Name,
};

// Each namespace of the data type also becomes a nesting type of the formatter.
if (dataType is QualifiedNamedTypeName { Container: NamespaceTypeContainer { Namespace: string ns } })
{
string[] namespaces = ns.Split('.');
QualifiedNamedTypeName? partialClassAsNamespaceStep = generatedResolverName;
QualifiedNamedTypeName? partialClassAsNamespaceStep = resolverOptions.Name;
for (int i = 0; i < namespaces.Length; i++)
{
partialClassAsNamespaceStep = new(TypeKind.Class)
Expand All @@ -143,7 +137,7 @@ public static ResolverRegisterInfo CreateArray(IArrayTypeSymbol dataType, Resolv
return partialClassAsNamespaceStep;
}

return generatedResolverName;
return resolverOptions.Name;
}

public override int GetHashCode() => this.DataType.GetHashCode();
Expand Down
4 changes: 2 additions & 2 deletions src/MessagePack.SourceGenerator/CodeAnalysis/TypeCollector.cs
Original file line number Diff line number Diff line change
Expand Up @@ -280,15 +280,15 @@ private bool CollectCore(ITypeSymbol typeSymbol)
return result;
}

FormattableType formattableType = new(typeSymbol);
FormattableType formattableType = new(typeSymbol, null);
if (formattableType.Name is QualifiedNamedTypeName { Name: string name } && EmbeddedTypes.Contains(name))
{
result = true;
this.alreadyCollected.Add(typeSymbol, result);
return result;
}

if (this.options.AssumedFormattableTypes.Contains(formattableType))
if (this.options.AssumedFormattableTypes.Contains(formattableType) || this.options.AssumedFormattableTypes.Contains(formattableType with { IsFormatterInSameAssembly = true }))
{
result = true;
this.alreadyCollected.Add(typeSymbol, result);
Expand Down
35 changes: 26 additions & 9 deletions src/MessagePack.SourceGenerator/CompositeResolverGenerator.cs
Original file line number Diff line number Diff line change
Expand Up @@ -16,15 +16,18 @@ public void Initialize(IncrementalGeneratorInitializationContext context)
$"{AttributeNamespace}.{CompositeResolverAttributeName}",
predicate: static (node, ct) => true,
transform: static (context, ct) => (
ResolverNamespace: context.TargetSymbol.ContainingNamespace.GetFullNamespaceName() ?? string.Empty,
ResolverName: context.TargetSymbol.Name,
ResolverName: new QualifiedNamedTypeName((INamedTypeSymbol)context.TargetSymbol),
Attribute: context.Attributes.Single()));

IncrementalValueProvider<AnalyzerOptions> options = GeneratorUtilities.GetAnalyzerOption(context);

var resolvers = attributeData.Combine(context.CompilationProvider).Select((leftRight, cancellationToken) =>
{
var source = leftRight.Left;
var compilation = leftRight.Right;

bool includeLocalFormatters = source.Attribute.NamedArguments.FirstOrDefault(kv => kv.Key == CompositeResolverAttributeIncludeLocalFormattersPropertyName).Value.Value is true;

if (source.Attribute.ConstructorArguments.Length > 0 && source.Attribute.ConstructorArguments[0].Kind == TypedConstantKind.Array)
{
// Get the semantic model we'll use for accessibility checks.
Expand All @@ -38,22 +41,36 @@ public void Initialize(IncrementalGeneratorInitializationContext context)
semanticModel,
source.Attribute.ConstructorArguments[0].Values.Select(tc => tc.Value as INamedTypeSymbol)).ToArray();

return (source.ResolverName, source.ResolverNamespace, ResolverCreationExpressions: resolverCreationExpressions, FormatterCreationExpressions: formatterCreationExpressions);
return (source.ResolverName, ResolverCreationExpressions: resolverCreationExpressions, FormatterCreationExpressions: formatterCreationExpressions, IncludeLocalFormatters: includeLocalFormatters);
}
else
{
return (source.ResolverName, source.ResolverNamespace, ResolverCreationExpressions: Array.Empty<string>(), FormatterCreationExpressions: Array.Empty<string>());
return (source.ResolverName, ResolverCreationExpressions: Array.Empty<string>(), FormatterCreationExpressions: Array.Empty<string>(), IncludeLocalFormatters: includeLocalFormatters);
}
});

context.RegisterSourceOutput(resolvers, (context, source) =>
context.RegisterSourceOutput(resolvers.Combine(options), (context, source) =>
{
AnalyzerOptions options = source.Right;

string[] formatterCreationExpressions = source.Left.FormatterCreationExpressions;
string[] resolverCreationExpressions = source.Left.ResolverCreationExpressions;
if (source.Left.IncludeLocalFormatters)
{
HashSet<string> allFormatters = new(formatterCreationExpressions, StringComparer.Ordinal);
allFormatters.UnionWith(options.KnownFormatters.Select(f => f.InstanceExpression));
formatterCreationExpressions = allFormatters.ToArray();

HashSet<string> allResolvers = new(resolverCreationExpressions, StringComparer.Ordinal);
allResolvers.Add($"{options.Generator.Resolver.Name.GetQualifiedName()}.Instance");
resolverCreationExpressions = allResolvers.ToArray();
}

CompositeResolverTemplate generator = new()
{
ResolverName = source.ResolverName,
ResolverNamespace = source.ResolverNamespace,
ResolverInstanceExpressions = source.ResolverCreationExpressions,
FormatterInstanceExpressions = source.FormatterCreationExpressions,
ResolverName = source.Left.ResolverName,
ResolverInstanceExpressions = resolverCreationExpressions,
FormatterInstanceExpressions = formatterCreationExpressions,
};
context.AddSource(generator.FileName, generator.TransformText());
});
Expand Down
73 changes: 73 additions & 0 deletions src/MessagePack.SourceGenerator/GeneratorUtilities.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
// Copyright (c) All contributors. All rights reserved.
// Licensed under the MIT license. See LICENSE file in the project root for full license information.

using System;
using System.Collections.Generic;
using System.Collections.Immutable;
using System.Text;
using Microsoft.CodeAnalysis;
using Microsoft.CodeAnalysis.CSharp.Syntax;
using static MessagePack.SourceGenerator.Constants;

namespace MessagePack.SourceGenerator;

internal static class GeneratorUtilities
{
internal static IncrementalValueProvider<AnalyzerOptions> GetAnalyzerOption(IncrementalGeneratorInitializationContext context)
{
// Search for [assembly: MessagePackKnownFormatter(typeof(SomeFormatter))]
var customFormatters = context.SyntaxProvider.ForAttributeWithMetadataName(
$"{AttributeNamespace}.{MessagePackKnownFormatterAttributeName}",
predicate: static (node, ct) => true,
transform: static (context, ct) => AnalyzerUtilities.ParseKnownFormatterAttribute(context.Attributes, ct)).Collect();

// Search for [assembly: MessagePackAssumedFormattable(typeof(SomeCustomType))]
var customFormattedTypes = context.SyntaxProvider.ForAttributeWithMetadataName(
$"{AttributeNamespace}.{MessagePackAssumedFormattableAttributeName}",
predicate: static (node, ct) => true,
transform: static (context, ct) => AnalyzerUtilities.ParseAssumedFormattableAttribute(context.Attributes, ct)).SelectMany((a, ct) => a).Collect();

// Search for all implementations of IMessagePackFormatter<T> in the compilation.
var customFormattersInThisCompilation = context.SyntaxProvider.CreateSyntaxProvider(
predicate: static (node, ct) => node is ClassDeclarationSyntax { BaseList.Types.Count: > 0 },
transform: (ctxt, ct) =>
{
return ctxt.SemanticModel.GetDeclaredSymbol(ctxt.Node, ct) is INamedTypeSymbol symbol
&& FormatterDescriptor.TryCreate(symbol, out FormatterDescriptor? formatter)
&& !formatter.ExcludeFromSourceGeneratedResolver
? formatter
: null;
}).Collect();

// Search for a [GeneratedMessagePackResolver] attribute (presumably on a partial class).
var resolverOptions = context.SyntaxProvider.ForAttributeWithMetadataName(
$"{AttributeNamespace}.{GeneratedMessagePackResolverAttributeName}",
predicate: static (node, ct) => true,
transform: static (context, ct) => AnalyzerUtilities.ParseGeneratorAttribute(context.Attributes, context.TargetSymbol, ct)).Collect().Select((a, ct) => a.SingleOrDefault(ao => ao is not null));

// Assembly an aggregating AnalyzerOptions object from the attributes and interface implementations that we've found.
var options = resolverOptions.Combine(customFormattedTypes).Combine(customFormatters).Combine(customFormattersInThisCompilation)
.Select(static (input, ct) =>
{
AnalyzerOptions? options = input.Left.Left.Left ?? new() { IsGeneratingSource = true };
ImmutableArray<FormatterDescriptor?> formatterImplementations = input.Right;

ImmutableArray<FormattableType> formattableTypes = input.Left.Left.Right;
ImmutableHashSet<FormatterDescriptor> formatterTypes = input.Left.Right.Aggregate(
ImmutableHashSet<FormatterDescriptor>.Empty,
(first, second) => first.Union(second));

// Merge the formatters discovered through attributes (which need only reference formatters from other assemblies),
// with formatters discovered in the project being compiled.
formatterTypes = formatterImplementations.Aggregate(
formatterTypes,
(first, second) => second is not null ? first.Add(second) : first);

options = options.WithFormatterTypes(formattableTypes, formatterTypes);

return options;
});

return options;
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,7 @@ private static void GenerateResolver(IGeneratorContext context, FullModel model)
.. model.EnumInfos,
.. model.UnionInfos,
.. model.ObjectInfos,
.. model.CustomFormatterInfos,
.. model.CustomFormatterInfos.Where(fi => fi.FormattableDataType.IsFormatterInSameAssembly),
];
ResolverTemplate resolverTemplate = new(options, registerInfos);
sb.Append(FileHeader);
Expand Down
54 changes: 2 additions & 52 deletions src/MessagePack.SourceGenerator/MessagePackGenerator.cs
Original file line number Diff line number Diff line change
Expand Up @@ -14,58 +14,7 @@ public partial class MessagePackGenerator : IIncrementalGenerator
{
public void Initialize(IncrementalGeneratorInitializationContext context)
{
// Search for [assembly: MessagePackKnownFormatter(typeof(SomeFormatter))]
var customFormatters = context.SyntaxProvider.ForAttributeWithMetadataName(
$"{AttributeNamespace}.{MessagePackKnownFormatterAttributeName}",
predicate: static (node, ct) => true,
transform: static (context, ct) => AnalyzerUtilities.ParseKnownFormatterAttribute(context.Attributes, ct)).Collect();

// Search for [assembly: MessagePackAssumedFormattable(typeof(SomeCustomType))]
var customFormattedTypes = context.SyntaxProvider.ForAttributeWithMetadataName(
$"{AttributeNamespace}.{MessagePackAssumedFormattableAttributeName}",
predicate: static (node, ct) => true,
transform: static (context, ct) => AnalyzerUtilities.ParseAssumedFormattableAttribute(context.Attributes, ct)).SelectMany((a, ct) => a).Collect();

// Search for all implementations of IMessagePackFormatter<T> in the compilation.
var customFormattersInThisCompilation = context.SyntaxProvider.CreateSyntaxProvider(
predicate: static (node, ct) => node is ClassDeclarationSyntax { BaseList.Types.Count: > 0 },
transform: (ctxt, ct) =>
{
return ctxt.SemanticModel.GetDeclaredSymbol(ctxt.Node, ct) is INamedTypeSymbol symbol
&& FormatterDescriptor.TryCreate(symbol, out FormatterDescriptor? formatter)
&& !formatter.ExcludeFromSourceGeneratedResolver
? formatter
: null;
}).Collect();

// Search for an [GeneratedMessagePackResolver] attribute (presumably on a partial class).
var resolverOptions = context.SyntaxProvider.ForAttributeWithMetadataName(
$"{AttributeNamespace}.{GeneratedMessagePackResolverAttributeName}",
predicate: static (node, ct) => true,
transform: static (context, ct) => AnalyzerUtilities.ParseGeneratorAttribute(context.Attributes, context.TargetSymbol, ct)).Collect().Select((a, ct) => a.SingleOrDefault(ao => ao is not null));

// Assembly an aggregating AnalyzerOptions object from the attributes and intefrace implementations that we've found.
var options = resolverOptions.Combine(customFormattedTypes).Combine(customFormatters).Combine(customFormattersInThisCompilation)
.Select(static (input, ct) =>
{
AnalyzerOptions? options = input.Left.Left.Left ?? new() { IsGeneratingSource = true };
ImmutableArray<FormatterDescriptor?> formatterImplementations = input.Right;

ImmutableArray<FormattableType> formattableTypes = input.Left.Left.Right;
ImmutableHashSet<FormatterDescriptor> formatterTypes = input.Left.Right.Aggregate(
ImmutableHashSet<FormatterDescriptor>.Empty,
(first, second) => first.Union(second));

// Merge the formatters discovered through attributes (which need only reference formatters from other assemblies),
// with formatters discovered in the project being compiled.
formatterTypes = formatterImplementations.Aggregate(
formatterTypes,
(first, second) => second is not null ? first.Add(second) : first);

options = options.WithFormatterTypes(formattableTypes, formatterTypes);

return options;
});
IncrementalValueProvider<AnalyzerOptions> options = GeneratorUtilities.GetAnalyzerOption(context);

var messagePackObjectTypes = context.SyntaxProvider.ForAttributeWithMetadataName(
$"{AttributeNamespace}.{MessagePackObjectAttributeName}",
Expand Down Expand Up @@ -122,6 +71,7 @@ from formatted in known.FormattableTypes
{
Formatter = known.Name,
DataType = formatted.Name,
FormattableDataType = formatted,
CustomFormatter = known,
});
modelPerType.Add(FullModel.Empty with { CustomFormatterInfos = customFormatterInfos, Options = options });
Expand Down
Loading