diff --git a/src/Meziantou.Analyzer.CodeFixers/Rules/UseAnOverloadThatHasMidpointRoundingFixer.cs b/src/Meziantou.Analyzer.CodeFixers/Rules/UseAnOverloadThatHasMidpointRoundingFixer.cs index 284dd4da7..617a68228 100644 --- a/src/Meziantou.Analyzer.CodeFixers/Rules/UseAnOverloadThatHasMidpointRoundingFixer.cs +++ b/src/Meziantou.Analyzer.CodeFixers/Rules/UseAnOverloadThatHasMidpointRoundingFixer.cs @@ -29,7 +29,7 @@ public override async Task RegisterCodeFixesAsync(CodeFixContext context) if (midpointRoundingSymbol is null) return; - if (!TryGetMidpointRoundingParameterInfo(semanticModel.Compilation, invocationOperation, midpointRoundingSymbol, out var parameterInfo)) + if (!TryGetFixInfo(semanticModel.Compilation, invocationOperation, invocationExpression, midpointRoundingSymbol, out var fixInfo)) return; foreach (var midpointRoundingMember in midpointRoundingSymbol.GetMembers().OfType()) @@ -44,61 +44,115 @@ public override async Task RegisterCodeFixesAsync(CodeFixContext context) var title = "Add MidpointRounding." + midpointRoundingMemberName; var codeAction = CodeAction.Create( title, - ct => AddMidpointRounding(context.Document, invocationExpression, parameterInfo, midpointRoundingSymbol, midpointRoundingMemberName, ct), + ct => AddMidpointRounding(context.Document, invocationExpression, fixInfo, midpointRoundingSymbol, midpointRoundingMemberName, ct), equivalenceKey: title); context.RegisterCodeFix(codeAction, context.Diagnostics); } } - private static bool TryGetMidpointRoundingParameterInfo(Compilation compilation, IInvocationOperation invocationOperation, INamedTypeSymbol midpointRoundingSymbol, out AdditionalParameterInfo parameterInfo) + private static bool TryGetFixInfo(Compilation compilation, IInvocationOperation invocationOperation, InvocationExpressionSyntax invocationExpression, INamedTypeSymbol midpointRoundingSymbol, [NotNullWhen(true)] out MidpointRoundingFixInfo? fixInfo) { + fixInfo = null; + var overloadFinder = new OverloadFinder(compilation); var overload = overloadFinder.FindOverloadWithAdditionalParameterOfType(invocationOperation, new OverloadOptions(IncludeObsoleteMembers: false, AllowOptionalParameters: true), [midpointRoundingSymbol]); if (overload is null) - { - parameterInfo = default; return false; - } + var parameterIndex = -1; for (var i = 0; i < overload.Parameters.Length; i++) { if (overload.Parameters[i].Type.IsEqualTo(midpointRoundingSymbol)) { - parameterInfo = new AdditionalParameterInfo(i, overload.Parameters[i].Name); - return true; + parameterIndex = i; + break; } } - parameterInfo = default; - return false; + if (parameterIndex < 0) + return false; + + var parameterName = overload.Parameters[parameterIndex].Name; + var arguments = invocationExpression.ArgumentList.Arguments; + + // All the arguments are positional, so they keep their bindings and the new argument can be added at the + // position of the parameter. When the position is after the last argument, some optional parameters are + // omitted, so the new argument must be named. + if (!arguments.Any(argument => argument.NameColon is not null)) + { + fixInfo = new MidpointRoundingFixInfo(invocationExpression.ArgumentList, parameterIndex <= arguments.Count ? parameterIndex : null, parameterName); + return true; + } + + // Some arguments are named, so they may not be in the parameter order. A positional argument cannot be added + // after them, and the parameters of the overload may not have the same names. The new argument is named and + // appended, and the existing named arguments are bound to the parameters of the overload. + var newArguments = arguments; + for (var i = 0; i < arguments.Count; i++) + { + var argument = arguments[i]; + var parameter = GetParameter(invocationOperation, argument); + if (parameter is null) + return false; + + var overloadParameterIndex = parameter.Ordinal < parameterIndex ? parameter.Ordinal : parameter.Ordinal + 1; + if (overloadParameterIndex >= overload.Parameters.Length) + return false; + + // The overload can declare the parameters in a different order, in which case the arguments cannot be reused + var overloadParameter = overload.Parameters[overloadParameterIndex]; + if (!overloadParameter.Type.IsEqualTo(parameter.Type)) + return false; + + if (argument.NameColon is null) + { + // A positional argument binds to the parameter at the same position, so it must come before the new parameter + if (i >= parameterIndex) + return false; + } + else if (!string.Equals(argument.NameColon.Name.Identifier.ValueText, overloadParameter.Name, StringComparison.Ordinal)) + { + var newArgument = argument.WithNameColon(argument.NameColon.WithName(SyntaxFactory.IdentifierName(overloadParameter.Name))); + newArguments = newArguments.Replace(newArguments[i], newArgument); + } + } + + fixInfo = new MidpointRoundingFixInfo(invocationExpression.ArgumentList.WithArguments(newArguments), ArgumentIndex: null, parameterName); + return true; + + static IParameterSymbol? GetParameter(IInvocationOperation invocationOperation, ArgumentSyntax argument) + { + foreach (var argumentOperation in invocationOperation.Arguments) + { + if (argumentOperation.Syntax == argument) + return argumentOperation.Parameter; + } + + return null; + } } - private static async Task AddMidpointRounding(Document document, InvocationExpressionSyntax invocationExpression, AdditionalParameterInfo parameterInfo, INamedTypeSymbol midpointRoundingSymbol, string midpointRoundingMember, CancellationToken cancellationToken) + private static async Task AddMidpointRounding(Document document, InvocationExpressionSyntax invocationExpression, MidpointRoundingFixInfo fixInfo, INamedTypeSymbol midpointRoundingSymbol, string midpointRoundingMember, CancellationToken cancellationToken) { var editor = await DocumentEditor.CreateAsync(document, cancellationToken).ConfigureAwait(false); var generator = editor.Generator; var midpointRoundingExpression = generator.TypeMemberAccessExpression(midpointRoundingSymbol, midpointRoundingMember, addImport: true); - var newArgument = (ArgumentSyntax)generator.Argument(midpointRoundingExpression); + var arguments = fixInfo.ArgumentList.Arguments; + var newArguments = fixInfo.ArgumentIndex is int argumentIndex + ? arguments.Insert(argumentIndex, (ArgumentSyntax)generator.Argument(midpointRoundingExpression)) + : arguments.Add((ArgumentSyntax)generator.Argument(fixInfo.ParameterName, RefKind.None, midpointRoundingExpression)); - InvocationExpressionSyntax newInvocation; - if (parameterInfo.ParameterIndex > invocationExpression.ArgumentList.Arguments.Count) - { - var namedArgument = (ArgumentSyntax)generator.Argument(parameterInfo.ParameterName, RefKind.None, midpointRoundingExpression); - var newArguments = invocationExpression.ArgumentList.Arguments.Add(namedArgument); - newInvocation = invocationExpression.WithArgumentList(SyntaxFactory.ArgumentList(newArguments)); - } - else - { - var newArguments = invocationExpression.ArgumentList.Arguments.Insert(parameterInfo.ParameterIndex, newArgument); - newInvocation = invocationExpression.WithArgumentList(SyntaxFactory.ArgumentList(newArguments)); - } + var newInvocation = invocationExpression.WithArgumentList(fixInfo.ArgumentList.WithArguments(newArguments)); editor.ReplaceNode(invocationExpression, newInvocation); return editor.GetChangedDocument(); } - private readonly record struct AdditionalParameterInfo(int ParameterIndex, string? ParameterName); + /// The arguments of the invocation, bound to the parameters of the overload. + /// The index at which the positional argument must be added, or when the argument must be named and appended. + /// The name of the parameter of the overload. + private sealed record MidpointRoundingFixInfo(ArgumentListSyntax ArgumentList, int? ArgumentIndex, string ParameterName); } diff --git a/tests/Meziantou.Analyzer.Test/Rules/UseAnOverloadThatHasMidpointRoundingAnalyzerTests.cs b/tests/Meziantou.Analyzer.Test/Rules/UseAnOverloadThatHasMidpointRoundingAnalyzerTests.cs index 758395d84..7844f1e2e 100644 --- a/tests/Meziantou.Analyzer.Test/Rules/UseAnOverloadThatHasMidpointRoundingAnalyzerTests.cs +++ b/tests/Meziantou.Analyzer.Test/Rules/UseAnOverloadThatHasMidpointRoundingAnalyzerTests.cs @@ -77,6 +77,118 @@ void A() return test.RunAsync(); } + [Fact] + public Task MathRoundWithReorderedNamedArguments_CodeFix_UsesNamedArgument() + { + var test = CreateTest(); + test.TestCode = """ + class Test + { + void A() + { + _ = {|MA0193:System.Math.Round(digits: 2, value: 1.25)|}; + } + } + """; + test.CodeActionIndex = 0; + test.FixedCode = """ + class Test + { + void A() + { + _ = System.Math.Round(digits: 2, value: 1.25, mode: System.MidpointRounding.ToEven); + } + } + """; + + return test.RunAsync(); + } + + [Fact] + public Task MathRoundWithNamedArgumentRenamedInOverload_CodeFix_UsesOverloadParameterName() + { + var test = CreateTest(); + test.TestCode = """ + class Test + { + void A() + { + _ = {|MA0193:System.Math.Round(a: 2.5)|}; + } + } + """; + test.CodeActionIndex = 0; + test.FixedCode = """ + class Test + { + void A() + { + _ = System.Math.Round(value: 2.5, mode: System.MidpointRounding.ToEven); + } + } + """; + + return test.RunAsync(); + } + + [Fact] + public Task MathRoundWithTrailingNamedArgument_CodeFix_UsesNamedArgument() + { + var test = CreateTest(); + test.TestCode = """ + class Test + { + void A() + { + _ = {|MA0193:System.Math.Round(2.5, digits: 2)|}; + } + } + """; + test.CodeActionIndex = 0; + test.FixedCode = """ + class Test + { + void A() + { + _ = System.Math.Round(2.5, digits: 2, mode: System.MidpointRounding.ToEven); + } + } + """; + + return test.RunAsync(); + } + + [Fact] + public Task IFloatingPointRoundWithReorderedNamedArguments_CodeFix_UsesNamedArgument() + { + var test = CreateTest(); + test.TestCode = """ + using System.Numerics; + + class Test + { + static T Round(T value) where T : IFloatingPoint + { + return {|MA0193:T.Round(digits: 2, x: value)|}; + } + } + """; + test.CodeActionIndex = 0; + test.FixedCode = """ + using System.Numerics; + + class Test + { + static T Round(T value) where T : IFloatingPoint + { + return T.Round(digits: 2, x: value, mode: System.MidpointRounding.ToEven); + } + } + """; + + return test.RunAsync(); + } + [Fact] public Task MathFRoundWithoutMode_ReportDiagnostic() {