mirror of
https://github.com/space-wizards/RobustToolbox.git
synced 2026-09-15 06:42:25 +02:00
186 lines
7.7 KiB
C#
186 lines
7.7 KiB
C#
using System.Collections.Immutable;
|
|
using System.Text;
|
|
using Microsoft.CodeAnalysis;
|
|
using Microsoft.CodeAnalysis.CSharp.Syntax;
|
|
using Robust.Roslyn.Shared;
|
|
using Robust.Roslyn.Shared.Helpers;
|
|
|
|
namespace Robust.Analyzers.Generators;
|
|
|
|
[Generator(LanguageNames.CSharp)]
|
|
public sealed class HasDependenciesGenerator : IIncrementalGenerator
|
|
{
|
|
private const string DependencyAttributeName = "Robust.Shared.IoC.DependencyAttribute";
|
|
private const string IHasDependenciesName = "Robust.Shared.IoC.IHasDependencies";
|
|
|
|
public void Initialize(IncrementalGeneratorInitializationContext context)
|
|
{
|
|
var fields = context.SyntaxProvider.ForAttributeWithMetadataName(
|
|
DependencyAttributeName,
|
|
static (node, _) => node is VariableDeclaratorSyntax,
|
|
static (syntaxContext, token) =>
|
|
{
|
|
var field = (IFieldSymbol)syntaxContext.TargetSymbol;
|
|
var fieldType = (INamedTypeSymbol)field.Type.WithNullableAnnotation(NullableAnnotation.NotAnnotated);
|
|
var owningType = (INamedTypeSymbol)field.ContainingSymbol;
|
|
|
|
var declarationSyntax = (TypeDeclarationSyntax)owningType.DeclaringSyntaxReferences[0]
|
|
.GetSyntax(token);
|
|
|
|
var partialTypeInfo = PartialTypeInfo.FromSymbol(owningType, declarationSyntax);
|
|
|
|
return (partialTypeInfo, FieldInfo: new FieldInfo(field.Name, fieldType.ToDisplayString(), field.IsReadOnly));
|
|
});
|
|
|
|
var grouped = fields
|
|
.Where(p => p.partialTypeInfo.IsValid)
|
|
.Collect()
|
|
.SelectMany(static (array, _) =>
|
|
{
|
|
return array.GroupBy(info => info.partialTypeInfo,
|
|
PartialTypeInfo.WithoutLocationComparer.Instance)
|
|
.Select(group => (group.Key, group.Select(e => e.FieldInfo).AsEquatableArray()));
|
|
});
|
|
|
|
var hasDependencyParents = grouped
|
|
.Collect()
|
|
.Combine(context.CompilationProvider)
|
|
.Select(static (a, cancel) =>
|
|
{
|
|
var (groups, compilation) = a;
|
|
|
|
var hasDependencyParents = new List<PartialTypeInfo>();
|
|
|
|
var ourAssemblyTypes = groups
|
|
.Where(g => g.Item2.All(static x => !x.IsReadOnly))
|
|
.Select(x => x.Key)
|
|
.ToDictionary<PartialTypeInfo, INamedTypeSymbol, PartialTypeInfo>(
|
|
x =>
|
|
{
|
|
var val = compilation.GetTypeByMetadataName(x.GetMetadataName());
|
|
if (val == null)
|
|
throw new InvalidOperationException();
|
|
|
|
return val.OriginalDefinition;
|
|
},
|
|
static x => x,
|
|
SymbolEqualityComparer.Default);
|
|
|
|
var hasDependencies = compilation.GetTypeByMetadataName(IHasDependenciesName);
|
|
if (hasDependencies == null && ourAssemblyTypes.Count != 0)
|
|
throw new InvalidOperationException();
|
|
|
|
foreach (var kvp in ourAssemblyTypes)
|
|
{
|
|
cancel.ThrowIfCancellationRequested();
|
|
|
|
var typeInfo = kvp.Key;
|
|
|
|
if (typeInfo.AllInterfaces.Contains(hasDependencies, SymbolEqualityComparer.Default))
|
|
{
|
|
hasDependencyParents.Add(kvp.Value);
|
|
continue;
|
|
}
|
|
|
|
for (
|
|
var ti = typeInfo.BaseType;
|
|
ti != null && SymbolEqualityComparer.Default.Equals(ti.ContainingAssembly, compilation.Assembly);
|
|
ti = ti.BaseType)
|
|
{
|
|
if (ourAssemblyTypes.ContainsKey(ti.OriginalDefinition))
|
|
{
|
|
hasDependencyParents.Add(kvp.Value);
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
|
|
return hasDependencyParents.ToImmutableArray();
|
|
});
|
|
|
|
context.RegisterImplementationSourceOutput(
|
|
grouped.Combine(hasDependencyParents),
|
|
static (productionContext, tuple) =>
|
|
{
|
|
var ((typeInfo, fields), hasParentList) = tuple;
|
|
|
|
if (fields.Any(a => a.IsReadOnly))
|
|
return;
|
|
|
|
var hasParent = hasParentList.Contains(typeInfo);
|
|
|
|
var sb = new IndentWriter(new StringBuilder());
|
|
|
|
sb.AppendLine("// <auto-generated />");
|
|
sb.AppendLine();
|
|
|
|
typeInfo.WriteHeader(ref sb, "[global::Robust.Shared.IoC.HasDependenciesGeneratedAttribute]");
|
|
|
|
if (!hasParent)
|
|
{
|
|
sb.AppendLine($" : global::{IHasDependenciesName}");
|
|
}
|
|
else
|
|
{
|
|
sb.AppendLine();
|
|
}
|
|
|
|
sb.AppendOpeningBrace(); // {
|
|
|
|
if (!hasParent && typeInfo.IsSealed)
|
|
{
|
|
// Explicit impl only
|
|
sb.AppendLineIndented("[global::Robust.Shared.Analyzers.RobustAutoGenerated]");
|
|
sb.AppendLineIndented($"void global::{IHasDependenciesName}.Inject(global::Robust.Shared.IoC.IDependencyCollection dependencies)");
|
|
sb.AppendOpeningBrace(); // {
|
|
WriteInject(ref sb, fields, false, typeInfo.DisplayName);
|
|
sb.AppendClosingBrace(); // }
|
|
}
|
|
else
|
|
{
|
|
if (!hasParent)
|
|
{
|
|
// Explicit impl -> protected virtual methods
|
|
sb.AppendLineIndented("[global::Robust.Shared.Analyzers.RobustAutoGenerated]");
|
|
sb.AppendLineIndented($"void global::{IHasDependenciesName}.Inject(global::Robust.Shared.IoC.IDependencyCollection dependencies)");
|
|
sb.AppendOpeningBrace(); // {
|
|
sb.AppendLineIndented("InjectImpl(dependencies);");
|
|
sb.AppendClosingBrace(); // }
|
|
sb.AppendLine();
|
|
}
|
|
|
|
// Protected virtual/override methods
|
|
sb.AppendLineIndented("[global::Robust.Shared.Analyzers.RobustAutoGenerated]");
|
|
sb.AppendLineIndented("[global::System.ComponentModel.EditorBrowsable(global::System.ComponentModel.EditorBrowsableState.Never)]");
|
|
sb.AppendLineIndented($"protected {(hasParent ? "override" : "virtual")} void InjectImpl(global::Robust.Shared.IoC.IDependencyCollection dependencies)");
|
|
sb.AppendOpeningBrace(); // {
|
|
WriteInject(ref sb, fields, hasParent, typeInfo.DisplayName);
|
|
sb.AppendClosingBrace(); // }
|
|
}
|
|
|
|
sb.AppendClosingBrace(); // }
|
|
|
|
typeInfo.WriteFooter(ref sb);
|
|
|
|
productionContext.AddSource(typeInfo.GetGeneratedFileName(), sb.ToString());
|
|
});
|
|
}
|
|
|
|
private static void WriteInject(ref IndentWriter sb, EquatableArray<FieldInfo> fields, bool isOverride, string typeName)
|
|
{
|
|
for (var i = 0; i < fields.Length; i++)
|
|
{
|
|
var field = fields[i];
|
|
sb.AppendLineIndented($"{field.Name} = dependencies.ResolveInject<global::{field.TypeName}>(typeof({typeName}));");
|
|
}
|
|
|
|
if (isOverride)
|
|
{
|
|
sb.AppendLine();
|
|
sb.AppendLineIndented("base.InjectImpl(dependencies);");
|
|
}
|
|
}
|
|
|
|
private readonly record struct FieldInfo(string Name, string TypeName, bool IsReadOnly);
|
|
}
|