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
Original file line number Diff line number Diff line change
@@ -1,6 +1,9 @@
using System.Collections.Concurrent;
using System.Collections.Immutable;
using System.Runtime.CompilerServices;
using System.Text;
using Microsoft.CodeAnalysis;
using Microsoft.CodeAnalysis.CSharp;
using Microsoft.CodeAnalysis.CSharp.Syntax;
using Microsoft.CodeAnalysis.Text;
using TUnit.Core.SourceGenerator.Extensions;
Expand All @@ -25,114 +28,168 @@ public void Initialize(IncrementalGeneratorInitializationContext context)
return !string.Equals(value, "false", StringComparison.OrdinalIgnoreCase);
});

var testClasses = context.SyntaxProvider
// Only classes that can contribute a static data-source property are inspected semantically:
// classes declaring an attributed static property, classes with a base list (which may inherit
// such properties, including from other assemblies), and partial classes (whose other parts may
// declare either). Every other class only walks to System.Object without finding anything.
var classChains = context.SyntaxProvider
.CreateSyntaxProvider(
predicate: static (s, _) => s is ClassDeclarationSyntax,
transform: static (ctx, _) => ctx)
predicate: static (node, _) => IsCandidateClass(node),
transform: static (ctx, ct) => GetStaticPropertyChain(ctx, ct))
.Where(static chain => chain.Length > 0);

var testClasses = classChains
.Collect()
.Combine(enabledProvider)
.Select((classesProviderPair, _) =>
.Select(static (classesProviderPair, _) =>
ParseStaticPropertyInitializers(classesProviderPair.Left, classesProviderPair.Right))
.WithTrackingName(ParseStaticProperties);

context.RegisterSourceOutput(testClasses, GenerateStaticPropertyInitialization);
}

private static EquatableArray<PropertyWithDataSourceModel> ParseStaticPropertyInitializers(ImmutableArray<GeneratorSyntaxContext> classesContext, bool enabledProvider)
/// <summary>
/// The static data-source properties declared directly on one type of an inheritance chain.
/// </summary>
/// <param name="TypeKey">Identity of the type, matching <see cref="SymbolEqualityComparer.Default"/>.</param>
/// <param name="Properties">The type's public static data-source properties, in member order.</param>
internal sealed record StaticPropertyTypeSegment(string TypeKey, EquatableArray<StaticPropertyEntry> Properties);

internal sealed record StaticPropertyEntry(string Name, PropertyWithDataSourceModel Model);

// Base types are shared by many classes. Cache each type's segment per compilation so chain
// walks read the members and attributes of a common base (or BCL) type only once.
private static readonly ConditionalWeakTable<Compilation, ConcurrentDictionary<INamedTypeSymbol, StaticPropertyTypeSegment>> SegmentCaches = new();

private static bool IsCandidateClass(SyntaxNode node)
{
if (!enabledProvider)
if (node is not ClassDeclarationSyntax classDeclaration)
{
return EquatableArray<PropertyWithDataSourceModel>.Empty;
return false;
}

// Use a dictionary to deduplicate static properties by their declaring type and name
// This prevents duplicate initialization when derived classes inherit static properties
var uniqueStaticProperties = new Dictionary<(INamedTypeSymbol DeclaringType, string Name), PropertyWithDataSource>(SymbolEqualityComparer.Default.ToTupleComparer());
var visitedTypes = new HashSet<INamedTypeSymbol>(SymbolEqualityComparer.Default);
var properties = new List<PropertyWithDataSource>();
if (classDeclaration.BaseList is not null || classDeclaration.Modifiers.Any(SyntaxKind.PartialKeyword))
{
return true;
}

foreach (var context in classesContext)
foreach (var member in classDeclaration.Members)
{
if (context.SemanticModel.GetDeclaredSymbol(context.Node) is not INamedTypeSymbol typeSymbol)
if (member is PropertyDeclarationSyntax { AttributeLists.Count: > 0 } property
&& property.Modifiers.Any(SyntaxKind.StaticKeyword))
{
continue;
return true;
}
}

return false;
}

private static EquatableArray<StaticPropertyTypeSegment> GetStaticPropertyChain(GeneratorSyntaxContext context, CancellationToken cancellationToken)
{
if (context.SemanticModel.GetDeclaredSymbol(context.Node, cancellationToken) is not INamedTypeSymbol typeSymbol)
{
return EquatableArray<StaticPropertyTypeSegment>.Empty;
}

// Skip open generic types - we can't generate code for types with unbound type parameters
// The initialization will happen in the consuming assembly that provides concrete type arguments
if (typeSymbol.IsGenericType && typeSymbol.TypeArguments.Any(t => t.TypeKind == TypeKind.TypeParameter))
{
return EquatableArray<StaticPropertyTypeSegment>.Empty;
}

var cache = SegmentCaches.GetValue(context.SemanticModel.Compilation,
static _ => new ConcurrentDictionary<INamedTypeSymbol, StaticPropertyTypeSegment>(SymbolEqualityComparer.IncludeNullability));

// Walk inheritance hierarchy to include base class static properties
var segments = new List<StaticPropertyTypeSegment>();
for (var currentType = typeSymbol; currentType != null; currentType = currentType.BaseType)
{
cancellationToken.ThrowIfCancellationRequested();

// Skip open generic types - we can't generate code for types with unbound type parameters
// The initialization will happen in the consuming assembly that provides concrete type arguments
if (typeSymbol.IsGenericType && typeSymbol.TypeArguments.Any(t => t.TypeKind == TypeKind.TypeParameter))
if (!cache.TryGetValue(currentType, out var segment))
{
continue;
segment = cache.GetOrAdd(currentType, static type => CreateSegment(type));
}

// Check if this type has any static properties with data source attributes
foreach (var prop in GetStaticPropertyDataSources(typeSymbol, visitedTypes, properties))
segments.Add(segment);
}

return segments.ToEquatableArray();
}

private static StaticPropertyTypeSegment CreateSegment(INamedTypeSymbol type)
{
List<StaticPropertyEntry>? properties = null;

foreach (var member in type.GetMembers())
{
if (member is IPropertySymbol { DeclaredAccessibility: Accessibility.Public, SetMethod.DeclaredAccessibility: Accessibility.Public, IsStatic: true } property) // Only static properties for session initialization
{
// Static properties belong to their declaring type, not derived types
// Only add if we haven't seen this exact property before
var key = (prop.Property.ContainingType, prop.Property.Name);
if (!uniqueStaticProperties.ContainsKey(key))
var dataSourceAttr = property.GetAttributes()
.FirstOrDefault(a => DataSourceAttributeHelper.IsDataSourceAttribute(a.AttributeClass));

if (dataSourceAttr != null)
{
uniqueStaticProperties[key] = prop;
properties ??= [];
properties.Add(new StaticPropertyEntry(
property.Name,
ToPropertyWithDataSourceModel(new PropertyWithDataSource
{
Property = property,
DataSourceAttribute = dataSourceAttr
})));
}
}
}

return uniqueStaticProperties.Values.Select(ToPropertyWithDataSourceModel).ToEquatableArray();
return new StaticPropertyTypeSegment(
type.ToDisplayString(SymbolDisplayFormat.FullyQualifiedFormat),
properties is null ? EquatableArray<StaticPropertyEntry>.Empty : properties.ToEquatableArray());
}

private static List<PropertyWithDataSource> GetStaticPropertyDataSources(
INamedTypeSymbol typeSymbol,
HashSet<INamedTypeSymbol> visitedType,
List<PropertyWithDataSource> properties
)
private static EquatableArray<PropertyWithDataSourceModel> ParseStaticPropertyInitializers(ImmutableArray<EquatableArray<StaticPropertyTypeSegment>> classChains, bool enabledProvider)
{
properties.Clear();
if (!enabledProvider)
{
return EquatableArray<PropertyWithDataSourceModel>.Empty;
}

// Walk inheritance hierarchy to include base class static properties
var currentType = typeSymbol;
while (currentType != null)
// Use a set to deduplicate static properties by their declaring type and name
// This prevents duplicate initialization when derived classes inherit static properties
var uniqueStaticProperties = new HashSet<(string DeclaringType, string Name)>();
var walkPropertyNames = new HashSet<string>();
var result = new List<PropertyWithDataSourceModel>();

foreach (var chain in classChains)
{
if (!visitedType.Add(currentType))
{
break;
}
walkPropertyNames.Clear();

foreach (var member in currentType.GetMembers())
// Every chain is walked in full, even through types an earlier chain reached: a property
// hidden by a derived type in one walk must still be added when its declaring type's own
// chain is walked, whatever the order of the chains. uniqueStaticProperties dedupes.
foreach (var segment in chain)
{
if (member is IPropertySymbol { DeclaredAccessibility: Accessibility.Public, SetMethod.DeclaredAccessibility: Accessibility.Public, IsStatic: true } property) // Only static properties for session initialization
foreach (var property in segment.Properties)
{
var dataSourceAttr = property.GetAttributes()
.FirstOrDefault(a => DataSourceAttributeHelper.IsDataSourceAttribute(a.AttributeClass));
// Check if we already have this property (in case of overrides)
if (!walkPropertyNames.Add(property.Name))
{
continue;
}

if (dataSourceAttr != null)
// Static properties belong to their declaring type, not derived types
// Only add if we haven't seen this exact property before
if (uniqueStaticProperties.Add((segment.TypeKey, property.Name)))
{
// Check if we already have this property (in case of overrides)
bool newProperty = true;
foreach (var p in properties)
{
if (p.Property.Name == property.Name)
{
newProperty = false;
break;
}
}

if (newProperty)
{
properties.Add(new PropertyWithDataSource
{
Property = property,
DataSourceAttribute = dataSourceAttr
});
}
result.Add(property.Model);
}
}
}
currentType = currentType.BaseType;
}

return properties;
return result.ToEquatableArray();
}

private static PropertyWithDataSourceModel ToPropertyWithDataSourceModel(PropertyWithDataSource staticProperty)
Expand Down
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
using System.Collections.Concurrent;
using Microsoft.CodeAnalysis;
using Microsoft.CodeAnalysis.CSharp.Syntax;
using TUnit.Core.SourceGenerator.CodeGenerators.Helpers;
Expand All @@ -13,6 +14,23 @@ public class AttributeWriter(Compilation compilation)
private readonly Dictionary<INamedTypeSymbol, string> _argumentFreeAttributeInitializerCache = new(SymbolEqualityComparer.Default);
private readonly Dictionary<INamedTypeSymbol, bool> _tunitRelatedCache = new(SymbolEqualityComparer.Default);

// Compilation.GetSemanticModel builds a fresh model (with empty binder caches) on every call.
// This writer is per-compilation, so share one model per tree across all attribute arguments.
private readonly ConcurrentDictionary<SyntaxTree, SemanticModel> _semanticModelCache = new();

/// <summary>
/// Gets a semantic model for a syntax tree owned by this writer's compilation, reusing one model per tree.
/// </summary>
public SemanticModel GetSemanticModel(SyntaxTree syntaxTree)
{
if (_semanticModelCache.TryGetValue(syntaxTree, out var semanticModel))
{
return semanticModel;
}

return _semanticModelCache.GetOrAdd(syntaxTree, tree => compilation.GetSemanticModel(tree));
}

public void WriteAttributes(ICodeWriter sourceCodeWriter,
IEnumerable<AttributeData> attributeDatas)
{
Expand Down Expand Up @@ -89,7 +107,7 @@ private string GetAttributeObjectInitializer(AttributeData attributeData, Attrib
{
if (!_argumentFreeAttributeInitializerCache.TryGetValue(attributeClass, out var argumentFreeInitializer))
{
argumentFreeInitializer = GetAttributeObjectInitializerInner(compilation, attributeData, syntax);
argumentFreeInitializer = GetAttributeObjectInitializerInner(attributeData, syntax);
_argumentFreeAttributeInitializerCache.Add(attributeClass, argumentFreeInitializer);
}

Expand All @@ -101,12 +119,12 @@ private string GetAttributeObjectInitializer(AttributeData attributeData, Attrib
return initializer;
}

initializer = GetAttributeObjectInitializerInner(compilation, attributeData, syntax);
initializer = GetAttributeObjectInitializerInner(attributeData, syntax);
_attributeObjectInitializerCache.Add(attributeData, initializer);
return initializer;
}

private static string GetAttributeObjectInitializerInner(Compilation compilation, AttributeData attributeData, AttributeSyntax syntax)
private string GetAttributeObjectInitializerInner(AttributeData attributeData, AttributeSyntax syntax)
{
var sourceCodeWriter = new CodeWriter("", includeHeader: false);

Expand All @@ -118,9 +136,9 @@ private static string GetAttributeObjectInitializerInner(Compilation compilation

var attributeName = attributeData.AttributeClass!.GloballyQualified();

var formattedConstructorArgs = string.Join(", ", constructorArgs.Select(x => FormatConstructorArgument(compilation, x)));
var formattedConstructorArgs = string.Join(", ", constructorArgs.Select(x => FormatConstructorArgument(x)));

var formattedProperties = properties.Select(x => FormatProperty(compilation, x)).ToArray();
var formattedProperties = properties.Select(x => FormatProperty(x)).ToArray();

sourceCodeWriter.Append($"new {attributeName}({formattedConstructorArgs})");

Expand All @@ -143,19 +161,19 @@ private static string GetAttributeObjectInitializerInner(Compilation compilation
return sourceCodeWriter.ToString();
}

private static string FormatConstructorArgument(Compilation compilation, AttributeArgumentSyntax attributeArgumentSyntax)
private string FormatConstructorArgument(AttributeArgumentSyntax attributeArgumentSyntax)
{
if (attributeArgumentSyntax.NameColon is not null)
{
return $"{attributeArgumentSyntax.NameColon!.Name}: {attributeArgumentSyntax.Expression.Accept(new FullyQualifiedWithGlobalPrefixRewriter(compilation.GetSemanticModel(attributeArgumentSyntax.SyntaxTree)))!.ToFullString()}";
return $"{attributeArgumentSyntax.NameColon!.Name}: {attributeArgumentSyntax.Expression.Accept(new FullyQualifiedWithGlobalPrefixRewriter(GetSemanticModel(attributeArgumentSyntax.SyntaxTree)))!.ToFullString()}";
}

return attributeArgumentSyntax.Accept(new FullyQualifiedWithGlobalPrefixRewriter(compilation.GetSemanticModel(attributeArgumentSyntax.SyntaxTree)))!.ToFullString();
return attributeArgumentSyntax.Accept(new FullyQualifiedWithGlobalPrefixRewriter(GetSemanticModel(attributeArgumentSyntax.SyntaxTree)))!.ToFullString();
}

private static string FormatProperty(Compilation compilation, AttributeArgumentSyntax attributeArgumentSyntax)
private string FormatProperty(AttributeArgumentSyntax attributeArgumentSyntax)
{
return $"{attributeArgumentSyntax.NameEquals!.Name} = {attributeArgumentSyntax.Expression.Accept(new FullyQualifiedWithGlobalPrefixRewriter(compilation.GetSemanticModel(attributeArgumentSyntax.SyntaxTree)))!.ToFullString()}";
return $"{attributeArgumentSyntax.NameEquals!.Name} = {attributeArgumentSyntax.Expression.Accept(new FullyQualifiedWithGlobalPrefixRewriter(GetSemanticModel(attributeArgumentSyntax.SyntaxTree)))!.ToFullString()}";
}

public static void WriteAttributeWithoutSyntax(ICodeWriter sourceCodeWriter, AttributeData attributeData)
Expand Down
Loading
Loading