File: Microsoft.NetCore.Analyzers\Runtime\UseStringEqualsOverStringCompare.Fixer.cs
Web Access
Project: src\sdk\src\Microsoft.CodeAnalysis.NetAnalyzers\src\Microsoft.CodeAnalysis.NetAnalyzers\Microsoft.CodeAnalysis.NetAnalyzers.csproj (Microsoft.CodeAnalysis.NetAnalyzers)
// Licensed to the .NET Foundation under one or more agreements.
// The .NET Foundation licenses this file to you under the MIT license.

using System;
using System.Collections.Generic;
using System.Collections.Immutable;
using System.Composition;
using System.Linq;
using System.Threading;
using System.Threading.Tasks;
using Analyzer.Utilities;
using Analyzer.Utilities.Extensions;
using Microsoft.CodeAnalysis;
using Microsoft.CodeAnalysis.CodeFixes;
using Microsoft.CodeAnalysis.Editing;
using Microsoft.CodeAnalysis.NetAnalyzers;
using Microsoft.CodeAnalysis.Operations;
using Microsoft.CodeAnalysis.Text;

using Resx = Microsoft.NetCore.Analyzers.MicrosoftNetCoreAnalyzersResources;
using RequiredSymbols = Microsoft.NetCore.Analyzers.Runtime.UseStringEqualsOverStringCompare.RequiredSymbols;

namespace Microsoft.NetCore.Analyzers.Runtime
{
    [ExportCodeFixProvider(LanguageNames.CSharp, LanguageNames.VisualBasic), Shared]
    public sealed class UseStringEqualsOverStringCompareFixer : SyntaxEditorBasedCodeFixProvider
    {
        public override ImmutableArray<string> FixableDiagnosticIds { get; } = ImmutableArray.Create(UseStringEqualsOverStringCompare.RuleId);

        public override async Task RegisterCodeFixesAsync(CodeFixContext context)
        {
            var semanticModel = await context.Document.GetRequiredSemanticModelAsync(context.CancellationToken).ConfigureAwait(false);
            var root = await context.Document.GetRequiredSyntaxRootAsync(context.CancellationToken).ConfigureAwait(false);

            if (GetViolation(root, context.Span, semanticModel, context.CancellationToken) is null)
            {
                return;
            }

            RegisterCodeFix(context, Resx.UseStringEqualsOverStringCompareCodeFixTitle, nameof(Resx.UseStringEqualsOverStringCompareCodeFixTitle));
        }

        protected override async Task ApplyFixAsync(Document document, Diagnostic diagnostic, SyntaxEditor editor, CancellationToken cancellationToken)
        {
            var semanticModel = await document.GetRequiredSemanticModelAsync(cancellationToken).ConfigureAwait(false);

            if (GetViolation(editor.OriginalRoot, diagnostic.Location.SourceSpan, semanticModel, cancellationToken) is not (IOperation violation, OperationReplacer replacer))
            {
                return;
            }

            //  The replacement is built out of the reported node's own descendants, so track them: a nested
            //  violation may already have been rewritten by the time this fix runs.
            foreach (var argument in replacer.GetArgumentSyntaxes(violation))
            {
                editor.TrackNode(argument);
            }

            editor.ReplaceNode(violation.Syntax, (currentNode, generator) =>
                replacer.CreateReplacementExpression(violation, generator, original => currentNode.GetCurrentNode(original) ?? original));
        }

        private static (IOperation Violation, OperationReplacer Replacer)? GetViolation(SyntaxNode root, TextSpan span, SemanticModel semanticModel, CancellationToken cancellationToken)
        {
            if (!RequiredSymbols.TryGetSymbols(semanticModel.Compilation, out var symbols))
            {
                return null;
            }

            var node = root.FindNode(span, getInnermostNodeForTie: true);
            var violation = semanticModel.GetOperation(node, cancellationToken);
            if (violation is not (IBinaryOperation or IInvocationOperation))
            {
                return null;
            }

            var replacer = GetOperationReplacers(symbols).FirstOrDefault(x => x.IsMatch(violation));

            return replacer is not null ? (violation, replacer) : null;
        }

        private static ImmutableArray<OperationReplacer> GetOperationReplacers(RequiredSymbols symbols)
        {
            return ImmutableArray.Create<OperationReplacer>(
                new StringStringCaseReplacer(symbols),
                new StringStringBoolReplacer(symbols),
                new StringStringStringComparisonReplacer(symbols),
                new OrdinalStringStringCaseReplacer(symbols));
        }

        /// <summary>
        /// Base class for an object that generate the replacement code for a reported violation.
        /// </summary>
        private abstract class OperationReplacer
        {
            protected OperationReplacer(RequiredSymbols symbols)
            {
                Symbols = symbols;
            }

            protected RequiredSymbols Symbols { get; }

            /// <summary>
            /// Indicates whether the current <see cref="OperationReplacer"/> applies to the specified violation.
            /// </summary>
            /// <param name="violation">The <see cref="IBinaryOperation"/> or <see cref="IInvocationOperation"/> at the location reported by the analyzer.</param>
            /// <returns>True if the current <see cref="OperationReplacer"/> applies to the specified violation.</returns>
            public abstract bool IsMatch(IOperation violation);

            /// <summary>
            /// Creates a replacement node for a violation that the current <see cref="OperationReplacer"/> applies to.
            /// Asserts if the current <see cref="OperationReplacer"/> does not apply to the specified violation.
            /// </summary>
            /// <param name="violation">The <see cref="IBinaryOperation"/> or <see cref="IInvocationOperation"/> obtained at the location reported by the analyzer.
            /// <see cref="IsMatch(IOperation)"/> must return <see langword="true"/> for this operation.</param>
            /// <param name="generator"></param>
            /// <param name="current">Maps a descendant of the violation onto its current form in the tree being edited.</param>
            /// <returns></returns>
            public abstract SyntaxNode CreateReplacementExpression(IOperation violation, SyntaxGenerator generator, Func<SyntaxNode, SyntaxNode> current);

            /// <summary>
            /// Gets the syntax nodes that <see cref="CreateReplacementExpression"/> carries over from the violation.
            /// </summary>
            public IEnumerable<SyntaxNode> GetArgumentSyntaxes(IOperation violation)
                => GetInvocation(violation).Arguments.Select(x => x.Value.Syntax);

            protected SyntaxNode CreateEqualsMemberAccess(SyntaxGenerator generator)
            {
                var stringTypeExpression = generator.TypeExpressionForStaticMemberAccess(Symbols.StringType);
                return generator.MemberAccessExpression(stringTypeExpression, nameof(string.Equals));
            }

            protected IInvocationOperation GetInvocation(IOperation violation)
            {
                var result = violation switch
                {
                    IBinaryOperation b => UseStringEqualsOverStringCompare.GetInvocationFromEqualityCheckWithLiteralZero(b),
                    IInvocationOperation i => UseStringEqualsOverStringCompare.GetInvocationFromEqualsCheckWithLiteralZero(i, Symbols.IntEquals),
                    _ => throw new NotSupportedException()
                };

                RoslynDebug.Assert(result is not null);

                return result;
            }

            protected static SyntaxNode InvertIfNotEquals(SyntaxNode stringEqualsInvocationExpression, IOperation equalsOrNotEqualsOperation, SyntaxGenerator generator)
            {
                if (equalsOrNotEqualsOperation is IBinaryOperation b)
                {
                    return b.OperatorKind is BinaryOperatorKind.NotEquals
                        ? generator.LogicalNotExpression(stringEqualsInvocationExpression)
                        : stringEqualsInvocationExpression;
                }

                if (equalsOrNotEqualsOperation is IInvocationOperation i)
                {
                    return i.Instance?.Parent is IUnaryOperation { OperatorKind: UnaryOperatorKind.Not }
                        ? generator.LogicalNotExpression(stringEqualsInvocationExpression)
                        : stringEqualsInvocationExpression;
                }

                throw new NotSupportedException();
            }
        }

        /// <summary>
        /// Replaces <see cref="string.Compare(string, string)"/> violations.
        /// </summary>
        private sealed class StringStringCaseReplacer : OperationReplacer
        {
            public StringStringCaseReplacer(RequiredSymbols symbols)
                : base(symbols)
            { }

            public override bool IsMatch(IOperation violation) => UseStringEqualsOverStringCompare.IsStringStringCase(violation, Symbols);

            public override SyntaxNode CreateReplacementExpression(IOperation violation, SyntaxGenerator generator, Func<SyntaxNode, SyntaxNode> current)
            {
                RoslynDebug.Assert(IsMatch(violation));

                var compareInvocation = GetInvocation(violation);
                var equalsInvocationSyntax = generator.InvocationExpression(
                    CreateEqualsMemberAccess(generator),
                    compareInvocation.Arguments.GetArgumentsInParameterOrder().Select(x => current(x.Value.Syntax)));

                return InvertIfNotEquals(equalsInvocationSyntax, violation, generator);
            }
        }

        /// <summary>
        /// Replaces <see cref="string.Compare(string, string, bool)"/> violations.
        /// </summary>
        private sealed class StringStringBoolReplacer : OperationReplacer
        {
            public StringStringBoolReplacer(RequiredSymbols symbols)
                : base(symbols)
            { }

            public override bool IsMatch(IOperation violation) => UseStringEqualsOverStringCompare.IsStringStringBoolCase(violation, Symbols);

            public override SyntaxNode CreateReplacementExpression(IOperation violation, SyntaxGenerator generator, Func<SyntaxNode, SyntaxNode> current)
            {
                RoslynDebug.Assert(IsMatch(violation));

                var compareInvocation = GetInvocation(violation);

                //  We know that the 'ignoreCase' argument in 'string.Compare(string, string, bool)' is a boolean literal
                //  because we've asserted that 'IsMatch' returns true.
                var ignoreCaseLiteral = (ILiteralOperation)compareInvocation.Arguments.GetArgumentForParameterAtIndex(2).Value;

                //  If the violation contains a call to 'string.Compare(x, y, true)' then we
                //  replace it with a call to 'string.Equals(x, y, StringComparison.CurrentCultureIgnoreCase)'.
                //  If the violation contains a call to 'string.Compare(x, y, false)' then we
                //  replace it with a call to 'string.Equals(x, y, StringComparison.CurrentCulture)'. 
                var stringComparisonEnumMemberName = ignoreCaseLiteral.ConstantValue.Value is true ?
                    nameof(StringComparison.CurrentCultureIgnoreCase) :
                    nameof(StringComparison.CurrentCulture);
                var stringComparisonMemberAccessSyntax = generator.MemberAccessExpression(
                    generator.TypeExpressionForStaticMemberAccess(Symbols.StringComparisonType),
                    stringComparisonEnumMemberName);

                var equalsInvocationSyntax = generator.InvocationExpression(
                    CreateEqualsMemberAccess(generator),
                    current(compareInvocation.Arguments.GetArgumentForParameterAtIndex(0).Value.Syntax),
                    current(compareInvocation.Arguments.GetArgumentForParameterAtIndex(1).Value.Syntax),
                    stringComparisonMemberAccessSyntax);

                return InvertIfNotEquals(equalsInvocationSyntax, violation, generator);
            }
        }

        /// <summary>
        /// Replaces <see cref="string.Compare(string, string, StringComparison)"/> violations.
        /// </summary>
        private sealed class StringStringStringComparisonReplacer : OperationReplacer
        {
            public StringStringStringComparisonReplacer(RequiredSymbols symbols)
                : base(symbols)
            { }

            public override bool IsMatch(IOperation violation) => UseStringEqualsOverStringCompare.IsStringStringStringComparisonCase(violation, Symbols);

            public override SyntaxNode CreateReplacementExpression(IOperation violation, SyntaxGenerator generator, Func<SyntaxNode, SyntaxNode> current)
            {
                RoslynDebug.Assert(IsMatch(violation));

                var invocation = GetInvocation(violation);
                var equalsInvocationSyntax = generator.InvocationExpression(
                    CreateEqualsMemberAccess(generator),
                    invocation.Arguments.GetArgumentsInParameterOrder().Select(x => current(x.Value.Syntax)));

                return InvertIfNotEquals(equalsInvocationSyntax, violation, generator);
            }
        }

        /// <summary>
        /// Replaces <see cref="string.CompareOrdinal(string, string)"/> violations.
        /// </summary>
        private sealed class OrdinalStringStringCaseReplacer : OperationReplacer
        {
            public OrdinalStringStringCaseReplacer(RequiredSymbols symbols)
                : base(symbols)
            { }

            public override bool IsMatch(IOperation violation) => UseStringEqualsOverStringCompare.IsOrdinalStringStringCase(violation, Symbols);

            public override SyntaxNode CreateReplacementExpression(IOperation violation, SyntaxGenerator generator, Func<SyntaxNode, SyntaxNode> current)
            {
                RoslynDebug.Assert(IsMatch(violation));

                var compareInvocation = GetInvocation(violation);
                var equalsInvocationSyntax = generator.InvocationExpression(
                    CreateEqualsMemberAccess(generator),
                    compareInvocation.Arguments.GetArgumentsInParameterOrder().Select(x => current(x.Value.Syntax)));

                return InvertIfNotEquals(equalsInvocationSyntax, violation, generator);
            }
        }
    }
}