Skip to content
Merged
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
191 changes: 2 additions & 189 deletions src/TUnit.Core.SourceGenerator/CodeGenerationHelpers.cs
Original file line number Diff line number Diff line change
@@ -1,6 +1,4 @@
using System.Collections.Immutable;
using Microsoft.CodeAnalysis;
using Microsoft.CodeAnalysis.CSharp.Syntax;
using TUnit.Core.SourceGenerator.CodeGenerators.Helpers;
using TUnit.Core.SourceGenerator.Extensions;

Expand All @@ -14,191 +12,16 @@ internal static class CodeGenerationHelpers
/// <summary>
/// Generates direct instantiation code for attributes.
/// </summary>
public static string GenerateAttributeInstantiation(AttributeData attr, ImmutableArray<IParameterSymbol> targetParameters = default)
public static string GenerateAttributeInstantiation(AttributeData attr)
{
var typeName = attr.AttributeClass!.GloballyQualified();
using var writer = new CodeWriter("", includeHeader: false);
writer.SetIndentLevel(1);
writer.Append($"new {typeName}(");

// Try to get the original syntax for better precision with decimal literals
var syntax = attr.ApplicationSyntaxReference?.GetSyntax();
var syntaxArguments = syntax?.ChildNodes()
.OfType<AttributeArgumentListSyntax>()
.FirstOrDefault()
?.Arguments.Where(x => x.NameEquals == null).ToList();

if (attr.ConstructorArguments.Length > 0)
{
var argStrings = new List<string>();

// Determine if this is an Arguments attribute and get parameter types
ITypeSymbol[]? parameterTypes = null;
if (attr.AttributeClass?.Name == "ArgumentsAttribute" && !targetParameters.IsDefault)
{
parameterTypes = targetParameters.Select(p => p.Type).ToArray();
}

var syntaxIndex = 0;
for (var i = 0; i < attr.ConstructorArguments.Length; i++)
{
var arg = attr.ConstructorArguments[i];

// Check if this is a params array parameter
if (i == attr.ConstructorArguments.Length - 1 && IsParamsArrayArgument(attr))
{
if (arg.Kind == TypedConstantKind.Array)
{
if (!arg.Values.IsDefault)
{
var elementIndex = 0;
var elements = arg.Values.Select(v =>
{
var paramType = parameterTypes != null && elementIndex < parameterTypes.Length
? parameterTypes[elementIndex]
: null;

// Check if the parameter type is decimal or nullable decimal
var underlyingType = paramType?.GetNullableUnderlyingType() ?? paramType;
var isDecimalType = underlyingType?.SpecialType == SpecialType.System_Decimal;

// For decimal parameters with syntax available, use the original text
if (isDecimalType &&
syntaxArguments != null && syntaxIndex < syntaxArguments.Count)
{
var syntaxExpression = syntaxArguments[syntaxIndex].Expression;
var originalText = syntaxExpression.ToString();
syntaxIndex++;

// Skip special handling for null values
if (originalText == "null")
{
elementIndex++;
return "null";
}

// Check if it's a string literal (starts and ends with quotes)
if (originalText.StartsWith("\"") && originalText.EndsWith("\""))
{
// For string literals, let the normal processing handle it (will use decimal.Parse)
syntaxIndex--; // Back up so normal processing can handle it
elementIndex++;
return TypedConstantParser.GetRawTypedConstantValue(v, paramType);
}

// Check if it's a constant reference (identifier) rather than a literal
// Identifiers don't contain dots, parentheses, or other operators
if (syntaxExpression is NameSyntax)
{
// For constant references, use the actual value from TypedConstant
elementIndex++;
return TypedConstantParser.GetRawTypedConstantValue(v, paramType);
}

// For numeric literals, remove any suffix and add 'm' for decimal
originalText = originalText.TrimEnd('d', 'D', 'f', 'F', 'l', 'L', 'u', 'U', 'm', 'M');
return $"{originalText}m";
}

syntaxIndex++;
elementIndex++;
return TypedConstantParser.GetRawTypedConstantValue(v, paramType);
});
argStrings.AddRange(elements);
}
}
else
{
var paramType = parameterTypes != null && i < parameterTypes.Length ? parameterTypes[i] : null;

// Check if the parameter type is decimal or nullable decimal
var underlyingType = paramType?.GetNullableUnderlyingType() ?? paramType;
var isDecimalType = underlyingType?.SpecialType == SpecialType.System_Decimal;

// For decimal parameters with syntax available, use the original text
if (isDecimalType &&
syntaxArguments != null && syntaxIndex < syntaxArguments.Count)
{
var syntaxExpression = syntaxArguments[syntaxIndex].Expression;
var originalText = syntaxExpression.ToString();
syntaxIndex++;

// Skip special handling for null values
if (originalText == "null")
{
argStrings.Add("null");
}
// Check if it's a string literal (starts and ends with quotes)
else if (originalText.StartsWith("\"") && originalText.EndsWith("\""))
{
// For string literals, let the normal processing handle it (will use decimal.Parse)
syntaxIndex--; // Back up so normal processing can handle it
argStrings.Add(TypedConstantParser.GetRawTypedConstantValue(arg, paramType));
}
// Check if it's a constant reference (identifier) rather than a literal
// Identifiers don't contain dots, parentheses, or other operators
else if (syntaxExpression is NameSyntax)
{
// For constant references, use the actual value from TypedConstant
argStrings.Add(TypedConstantParser.GetRawTypedConstantValue(arg, paramType));
}
else
{
// For numeric literals, remove any suffix and add 'm' for decimal
originalText = originalText.TrimEnd('d', 'D', 'f', 'F', 'l', 'L', 'u', 'U', 'm', 'M');
argStrings.Add($"{originalText}m");
}
}
else
{
syntaxIndex++;
argStrings.Add(TypedConstantParser.GetRawTypedConstantValue(arg, paramType));
}
}
}
else
{
var paramType = parameterTypes != null && i < parameterTypes.Length ? parameterTypes[i] : null;

// For decimal parameters with syntax available, use the original text
if (paramType?.SpecialType == SpecialType.System_Decimal &&
syntaxArguments != null && syntaxIndex < syntaxArguments.Count)
{
var syntaxExpression = syntaxArguments[syntaxIndex].Expression;
var originalText = syntaxExpression.ToString();
syntaxIndex++;
// Check if it's a string literal (starts and ends with quotes)
if (originalText.StartsWith("\"") && originalText.EndsWith("\""))
{
// For string literals, let the normal processing handle it (will use decimal.Parse)
syntaxIndex--; // Back up so normal processing can handle it
argStrings.Add(TypedConstantParser.GetRawTypedConstantValue(arg, paramType));
}
// Check if it's a constant reference (identifier) rather than a literal
// Identifiers don't contain dots, parentheses, or other operators
else if (syntaxExpression is NameSyntax)
{
// For constant references, use the actual value from TypedConstant
argStrings.Add(TypedConstantParser.GetRawTypedConstantValue(arg, paramType));
}
else
{
// For numeric literals, remove any suffix and add 'm' for decimal
originalText = originalText.TrimEnd('d', 'D', 'f', 'F', 'l', 'L', 'u', 'U', 'm', 'M');
argStrings.Add($"{originalText}m");
}
}
else
{
if (syntaxArguments != null && syntaxIndex < syntaxArguments.Count)
{
syntaxIndex++;
}
argStrings.Add(TypedConstantParser.GetRawTypedConstantValue(arg, paramType));
}
}
}

var argStrings = attr.ConstructorArguments.Select(arg => TypedConstantParser.GetRawTypedConstantValue(arg));
writer.AppendJoin(", ", argStrings);
}

Expand All @@ -215,16 +38,6 @@ public static string GenerateAttributeInstantiation(AttributeData attr, Immutabl
return writer.ToString().Trim();
}

/// <summary>
/// Determines if an argument is for a params array parameter.
/// </summary>
private static bool IsParamsArrayArgument(AttributeData attr)
{
var typeName = attr.AttributeClass!.GloballyQualified();

return typeName is "global::TUnit.Core.ArgumentsAttribute" or "global::TUnit.Core.InlineDataAttribute";
}

/// <summary>
/// Determines if an attribute should be excluded from metadata.
/// </summary>
Expand Down
Loading