From 86425c36c2d8c56ccc6fba3071485481146eb955 Mon Sep 17 00:00:00 2001 From: meziantou Date: Mon, 11 Mar 2019 17:31:52 -0400 Subject: [PATCH] Add optimize Enumerable.Count() rule --- .../source.extension.vsixmanifest | 2 +- .../Internals/TypeSymbolExtensions.cs | 8 + src/Meziantou.Analyzer/RuleIdentifiers.cs | 3 +- .../Rules/OptimizeLinqUsageAnalyzer.cs | 238 ++++++++++++++++-- .../OptimizeLinqUsageAnalyzer.Count.Tests.cs | 183 ++++++++++++++ ...inqUsageAnalyzer.DuplicateOrderBy.Tests.cs | 24 ++ 6 files changed, 442 insertions(+), 16 deletions(-) create mode 100644 tests/Meziantou.Analyzer.Test/Rules/OptimizeLinqUsageAnalyzer.Count.Tests.cs diff --git a/src/Meziantou.Analyzer.Vsix/source.extension.vsixmanifest b/src/Meziantou.Analyzer.Vsix/source.extension.vsixmanifest index 12cc4d31f..afe4afdee 100644 --- a/src/Meziantou.Analyzer.Vsix/source.extension.vsixmanifest +++ b/src/Meziantou.Analyzer.Vsix/source.extension.vsixmanifest @@ -1,7 +1,7 @@ - + Meziantou.Analyzer A Roslyn analyzer to enforce some good pratices in C# diff --git a/src/Meziantou.Analyzer/Internals/TypeSymbolExtensions.cs b/src/Meziantou.Analyzer/Internals/TypeSymbolExtensions.cs index 8e5c5f1b0..3ef18451c 100644 --- a/src/Meziantou.Analyzer/Internals/TypeSymbolExtensions.cs +++ b/src/Meziantou.Analyzer/Internals/TypeSymbolExtensions.cs @@ -76,6 +76,14 @@ public static bool IsString(this ITypeSymbol symbol) return symbol.SpecialType == SpecialType.System_String; } + public static bool IsInt32(this ITypeSymbol symbol) + { + if (symbol == null) + return false; + + return symbol.SpecialType == SpecialType.System_Int32; + } + public static bool IsBoolean(this ITypeSymbol symbol) { if (symbol == null) diff --git a/src/Meziantou.Analyzer/RuleIdentifiers.cs b/src/Meziantou.Analyzer/RuleIdentifiers.cs index 822d4a541..b24a675be 100644 --- a/src/Meziantou.Analyzer/RuleIdentifiers.cs +++ b/src/Meziantou.Analyzer/RuleIdentifiers.cs @@ -33,7 +33,8 @@ internal static class RuleIdentifiers public const string DoNotRemoveOriginalExceptionFromThrowStatement = "MA0027"; public const string OptimizeStringBuilderUsage = "MA0028"; public const string OptimizeLinqUsage = "MA0029"; - public const string DuplicateOrderBy = "MA0030"; + public const string DuplicateEnumerable_OrderBy = "MA0030"; + public const string OptimizeEnumerable_Count = "MA0031"; public static string GetHelpUri(string idenfifier) { diff --git a/src/Meziantou.Analyzer/Rules/OptimizeLinqUsageAnalyzer.cs b/src/Meziantou.Analyzer/Rules/OptimizeLinqUsageAnalyzer.cs index 5ed830190..421b44b57 100644 --- a/src/Meziantou.Analyzer/Rules/OptimizeLinqUsageAnalyzer.cs +++ b/src/Meziantou.Analyzer/Rules/OptimizeLinqUsageAnalyzer.cs @@ -1,10 +1,10 @@ using System; -using System.Collections.Generic; using System.Collections.Immutable; using System.Linq; using Microsoft.CodeAnalysis; using Microsoft.CodeAnalysis.Diagnostics; using Microsoft.CodeAnalysis.Operations; +using static System.FormattableString; namespace Meziantou.Analyzer.Rules { @@ -32,16 +32,26 @@ public class OptimizeLinqUsageAnalyzer : DiagnosticAnalyzer helpLinkUri: RuleIdentifiers.GetHelpUri(RuleIdentifiers.OptimizeLinqUsage)); private static readonly DiagnosticDescriptor s_duplicateOrderByMethodsRule = new DiagnosticDescriptor( - RuleIdentifiers.DuplicateOrderBy, + RuleIdentifiers.DuplicateEnumerable_OrderBy, title: "Optimize LINQ usage", messageFormat: "Remove the first '{0}' method or use '{1}'", RuleCategories.Usage, DiagnosticSeverity.Info, isEnabledByDefault: true, description: "", - helpLinkUri: RuleIdentifiers.GetHelpUri(RuleIdentifiers.DuplicateOrderBy)); + helpLinkUri: RuleIdentifiers.GetHelpUri(RuleIdentifiers.DuplicateEnumerable_OrderBy)); - public override ImmutableArray SupportedDiagnostics => ImmutableArray.Create(s_listMethodsRule, s_combineLinqMethodsRule, s_duplicateOrderByMethodsRule); + private static readonly DiagnosticDescriptor s_optimizeCountRule = new DiagnosticDescriptor( + RuleIdentifiers.OptimizeEnumerable_Count, + title: "Optimize Enumerable.Count usage", + messageFormat: "{0}", + RuleCategories.Usage, + DiagnosticSeverity.Info, + isEnabledByDefault: true, + description: "", + helpLinkUri: RuleIdentifiers.GetHelpUri(RuleIdentifiers.OptimizeEnumerable_Count)); + + public override ImmutableArray SupportedDiagnostics => ImmutableArray.Create(s_listMethodsRule, s_combineLinqMethodsRule, s_duplicateOrderByMethodsRule, s_optimizeCountRule); public override void Initialize(AnalysisContext context) { @@ -70,15 +80,7 @@ private void Analyze(OperationAnalysisContext context) UseIndexerInsteadOfElementAt(context, operation); CombineWhereWithNextMethod(context, operation, enumerableSymbol); RemoveTwoConsecutiveOrderBy(context, operation, enumerableSymbol); - - // TODO Count() < 0 => false - // TODO Count() <= 0 => !Any() - // TODO Count() == 0 => !Any() - // TODO Count() > 0 => Any() - // TODO Count() < 10 => .Any() - // TODO Count() <= 10 => .Any() - // TODO Count() >= 10 => .Skip(9).Any() - // TODO Count() > 10 => .Skip(10).Any() + OptimizeCountUsage(context, operation); } private void UseCountPropertyInsteadOfMethod(OperationAnalysisContext context, IInvocationOperation operation) @@ -197,7 +199,9 @@ private void CombineWhereWithNextMethod(OperationAnalysisContext context, IInvoc private void RemoveTwoConsecutiveOrderBy(OperationAnalysisContext context, IInvocationOperation operation, ITypeSymbol enumerableSymbol) { if (string.Equals(operation.TargetMethod.Name, nameof(Enumerable.OrderBy), StringComparison.Ordinal) || - string.Equals(operation.TargetMethod.Name, nameof(Enumerable.OrderByDescending), StringComparison.Ordinal)) + string.Equals(operation.TargetMethod.Name, nameof(Enumerable.OrderByDescending), StringComparison.Ordinal) || + string.Equals(operation.TargetMethod.Name, nameof(Enumerable.ThenBy), StringComparison.Ordinal) || + string.Equals(operation.TargetMethod.Name, nameof(Enumerable.ThenByDescending), StringComparison.Ordinal)) { var parent = GetParentLinqOperation(operation); if (parent != null && parent.TargetMethod.ContainingType.IsEqualsTo(enumerableSymbol)) @@ -211,6 +215,193 @@ private void RemoveTwoConsecutiveOrderBy(OperationAnalysisContext context, IInvo } } + private void OptimizeCountUsage(OperationAnalysisContext context, IInvocationOperation operation) + { + if (!string.Equals(operation.TargetMethod.Name, nameof(Enumerable.Count), StringComparison.Ordinal)) + return; + + var binaryOperation = GetParentBinaryOperation(operation, out var countOperand); + if (binaryOperation == null) + return; + + if (!IsSupportedOperator(binaryOperation.OperatorKind)) + return; + + if (!binaryOperation.LeftOperand.Type.IsInt32() || !binaryOperation.RightOperand.Type.IsInt32()) + return; + + var opKind = NormalizeOperator(); + var otherOperand = binaryOperation.LeftOperand == countOperand ? binaryOperation.RightOperand : binaryOperation.LeftOperand; + if (otherOperand == null) + return; + + if (otherOperand.ConstantValue.HasValue && otherOperand.ConstantValue.Value is int value) + { + string message = null; + switch (opKind) + { + case BinaryOperatorKind.Equals: + if (value < 0) + { + // expr.Count() == -1 + message = "Expression is always false"; + } + else if (value == 0) + { + // expr.Count() == 0 + message = "Replace 'Count() == 0' with 'Any() == false'"; + } + else if (value == 1) + { + // expr.Count() == 1 + message = "Replace 'Count() == 1' with 'Any()'"; + } + else + { + // expr.Count() == 10 => expr.Skip(9).Any() + message = Invariant($"Replace 'Count() == {value}' with 'Skip({value - 1}).Any()'"); + } + + break; + + case BinaryOperatorKind.NotEquals: + if (value < 0) + { + // expr.Count() != -1 is always true + message = "Expression is always true"; + } + else if (value == 0) + { + // expr.Count() != 0 + message = "Replace 'Count() != 0' with 'Any()'"; + } + + break; + + case BinaryOperatorKind.LessThan: + if (value <= 0) + { + // expr.Count() < 0 + message = "Expression is always false"; + } + else if (value == 1) + { + // expr.Count() < 1 ==> expr.Count() == 0 + message = "Replace 'Count() < 1' with 'Any() == false'"; + } + else + { + // expr.Count() < 10 + message = Invariant($"Replace 'Count() < {value}' with 'Skip({value - 1}).Any() == false'"); + } + + break; + + case BinaryOperatorKind.LessThanOrEqual: + if (value < 0) + { + // expr.Count() <= -1 + message = "Expression is always false"; + } + else if (value == 0) + { + // expr.Count() <= 0 + message = "Replace 'Count() <= 0' with 'Any() == false'"; + } + else + { + // expr.Count() < 10 + message = Invariant($"Replace 'Count() <= {value}' with 'Skip({value}).Any() == false'"); + } + + break; + + case BinaryOperatorKind.GreaterThan: + if (value < 0) + { + // expr.Count() > -1 + message = "Expression is always true"; + } + else if (value == 0) + { + // expr.Count() > 0 + message = "Replace 'Count() > 0' with 'Any()'"; + } + else + { + // expr.Count() > 1 + message = Invariant($"Replace 'Count() > {value}' with 'Skip({value}).Any()'"); + } + + break; + + case BinaryOperatorKind.GreaterThanOrEqual: + if (value <= 0) + { + // expr.Count() >= 0 + message = "Expression is always true"; + } + else if (value == 1) + { + // expr.Count() >= 1 + message = "Replace 'Count() >= 1' with 'Any()'"; + } + else + { + // expr.Count() >= 2 + message = Invariant($"Replace 'Count() >= {value}' with 'Skip({value - 1}).Any()'"); + } + + break; + + } + + if (message != null) + { + context.ReportDiagnostic(Diagnostic.Create(s_optimizeCountRule, binaryOperation.Syntax.GetLocation(), message)); + } + + // TODO detect non constant values + } + + bool IsSupportedOperator(BinaryOperatorKind operatorKind) + { + switch (operatorKind) + { + case BinaryOperatorKind.Equals: + case BinaryOperatorKind.NotEquals: + case BinaryOperatorKind.LessThan: + case BinaryOperatorKind.LessThanOrEqual: + case BinaryOperatorKind.GreaterThanOrEqual: + case BinaryOperatorKind.GreaterThan: + return true; + default: + return false; + } + } + + BinaryOperatorKind NormalizeOperator() + { + bool isCountLeftOperand = binaryOperation.LeftOperand == countOperand; + switch (binaryOperation.OperatorKind) + { + case BinaryOperatorKind.LessThan: + return isCountLeftOperand ? BinaryOperatorKind.LessThan : BinaryOperatorKind.GreaterThan; + case BinaryOperatorKind.LessThanOrEqual: + return isCountLeftOperand ? BinaryOperatorKind.LessThanOrEqual : BinaryOperatorKind.GreaterThanOrEqual; + + case BinaryOperatorKind.GreaterThanOrEqual: + return isCountLeftOperand ? BinaryOperatorKind.GreaterThanOrEqual : BinaryOperatorKind.LessThanOrEqual; + + case BinaryOperatorKind.GreaterThan: + return isCountLeftOperand ? BinaryOperatorKind.GreaterThan : BinaryOperatorKind.LessThan; + + default: + return binaryOperation.OperatorKind; + } + } + } + private static ITypeSymbol GetActualType(IArgumentOperation argument) { var value = argument.Value; @@ -240,5 +431,24 @@ private static IInvocationOperation GetParentLinqOperation(IOperation op) return null; } + + private static IBinaryOperation GetParentBinaryOperation(IOperation op, out IOperation operand) + { + var parent = op.Parent; + if (parent is IConversionOperation) + { + op = parent; + parent = parent.Parent; + } + + if (parent is IBinaryOperation binaryOperation) + { + operand = op; + return binaryOperation; + } + + operand = null; + return null; + } } } diff --git a/tests/Meziantou.Analyzer.Test/Rules/OptimizeLinqUsageAnalyzer.Count.Tests.cs b/tests/Meziantou.Analyzer.Test/Rules/OptimizeLinqUsageAnalyzer.Count.Tests.cs new file mode 100644 index 000000000..c7004d846 --- /dev/null +++ b/tests/Meziantou.Analyzer.Test/Rules/OptimizeLinqUsageAnalyzer.Count.Tests.cs @@ -0,0 +1,183 @@ +using System.Collections.Generic; +using System.Linq; +using Meziantou.Analyzer.Rules; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.Diagnostics; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using TestHelper; + +namespace Meziantou.Analyzer.Test.Rules +{ + [TestClass] + public class OptimizeLinqUsageAnalyzerCountTests : CodeFixVerifier + { + protected override DiagnosticAnalyzer GetCSharpDiagnosticAnalyzer() => new OptimizeLinqUsageAnalyzer(); + protected override string ExpectedDiagnosticId => "MA0031"; + protected override DiagnosticSeverity ExpectedDiagnosticSeverity => DiagnosticSeverity.Info; + + [DataTestMethod] + [DataRow("Count() == -1", "Expression is always false")] + [DataRow("Count() == 0", "Replace 'Count() == 0' with 'Any() == false'")] + [DataRow("Count() == 1", "Replace 'Count() == 1' with 'Any()'")] + [DataRow("Count() == 2", "Replace 'Count() == 2' with 'Skip(1).Any()'")] + public void Count_Equals(string text, string expectedMessage) + { + var project = new ProjectBuilder() + .AddReference(typeof(IEnumerable<>)) + .AddReference(typeof(Enumerable)) + .WithSource(@"using System.Linq; +class Test +{ + public Test() + { + var enumerable = System.Linq.Enumerable.Empty(); + _ = enumerable." + text + @"; + } +} +"); + + VerifyDiagnostic(project, CreateDiagnosticResult(line: 7, column: 13, message: expectedMessage)); + } + + [DataTestMethod] + [DataRow("Count() != -2", "Expression is always true")] + [DataRow("Count() != 0", "Replace 'Count() != 0' with 'Any()'")] + public void Count_NotEquals(string text, string expectedMessage) + { + var project = new ProjectBuilder() + .AddReference(typeof(IEnumerable<>)) + .AddReference(typeof(Enumerable)) + .WithSource(@"using System.Linq; +class Test +{ + public Test() + { + var enumerable = System.Linq.Enumerable.Empty(); + _ = enumerable." + text + @"; + } +} +"); + + VerifyDiagnostic(project, CreateDiagnosticResult(line: 7, column: 13, message: expectedMessage)); + } + + [DataTestMethod] + [DataRow("Count() != 1")] + [DataRow("Count() != 2")] + public void Count_NotEquals_Valid(string text) + { + var project = new ProjectBuilder() + .AddReference(typeof(IEnumerable<>)) + .AddReference(typeof(Enumerable)) + .WithSource(@"using System.Linq; +class Test +{ + public Test() + { + var enumerable = System.Linq.Enumerable.Empty(); + _ = enumerable." + text + @"; + } +} +"); + + VerifyDiagnostic(project); + } + + [DataTestMethod] + [DataRow("Count() < -1", "Expression is always false")] + [DataRow("Count() < 0", "Expression is always false")] + [DataRow("Count() < 1", "Replace 'Count() < 1' with 'Any() == false'")] + [DataRow("Count() < 2", "Replace 'Count() < 2' with 'Skip(1).Any() == false'")] + public void Count_LessThan(string text, string expectedMessage) + { + var project = new ProjectBuilder() + .AddReference(typeof(IEnumerable<>)) + .AddReference(typeof(Enumerable)) + .WithSource(@"using System.Linq; +class Test +{ + public Test() + { + var enumerable = System.Linq.Enumerable.Empty(); + _ = enumerable." + text + @"; + } +} +"); + + VerifyDiagnostic(project, CreateDiagnosticResult(line: 7, column: 13, message: expectedMessage)); + } + + [DataTestMethod] + [DataRow("Count() <= -1", "Expression is always false")] + [DataRow("Count() <= 0", "Replace 'Count() <= 0' with 'Any() == false'")] + [DataRow("Count() <= 1", "Replace 'Count() <= 1' with 'Skip(1).Any() == false'")] + [DataRow("Count() <= 2", "Replace 'Count() <= 2' with 'Skip(2).Any() == false'")] + public void Count_LessThanOrEqual(string text, string expectedMessage) + { + var project = new ProjectBuilder() + .AddReference(typeof(IEnumerable<>)) + .AddReference(typeof(Enumerable)) + .WithSource(@"using System.Linq; +class Test +{ + public Test() + { + var enumerable = System.Linq.Enumerable.Empty(); + _ = enumerable." + text + @"; + } +} +"); + + VerifyDiagnostic(project, CreateDiagnosticResult(line: 7, column: 13, message: expectedMessage)); + } + + [DataTestMethod] + [DataRow("Count() > -1", "Expression is always true")] + [DataRow("Count() > 0", "Replace 'Count() > 0' with 'Any()'")] + [DataRow("Count() > 1", "Replace 'Count() > 1' with 'Skip(1).Any()'")] + [DataRow("Count() > 2", "Replace 'Count() > 2' with 'Skip(2).Any()'")] + public void Count_GreaterThan(string text, string expectedMessage) + { + var project = new ProjectBuilder() + .AddReference(typeof(IEnumerable<>)) + .AddReference(typeof(Enumerable)) + .WithSource(@"using System.Linq; +class Test +{ + public Test() + { + var enumerable = System.Linq.Enumerable.Empty(); + _ = enumerable." + text + @"; + } +} +"); + + VerifyDiagnostic(project, CreateDiagnosticResult(line: 7, column: 13, message: expectedMessage)); + } + + [DataTestMethod] + [DataRow("enumerable.Count() >= -1", "Expression is always true")] + [DataRow("-1 <= enumerable.Count()", "Expression is always true")] + [DataRow("enumerable.Count() >= 0", "Expression is always true")] + [DataRow("enumerable.Count() >= 1", "Replace 'Count() >= 1' with 'Any()'")] + [DataRow("enumerable.Count() >= 2", "Replace 'Count() >= 2' with 'Skip(1).Any()'")] + public void Count_GreaterThanOrEqual(string text, string expectedMessage) + { + var project = new ProjectBuilder() + .AddReference(typeof(IEnumerable<>)) + .AddReference(typeof(Enumerable)) + .WithSource(@"using System.Linq; +class Test +{ + public Test() + { + var enumerable = System.Linq.Enumerable.Empty(); + _ = " + text + @"; + } +} +"); + + VerifyDiagnostic(project, CreateDiagnosticResult(line: 7, column: 13, message: expectedMessage)); + } + } +} diff --git a/tests/Meziantou.Analyzer.Test/Rules/OptimizeLinqUsageAnalyzer.DuplicateOrderBy.Tests.cs b/tests/Meziantou.Analyzer.Test/Rules/OptimizeLinqUsageAnalyzer.DuplicateOrderBy.Tests.cs index 13482d672..73b099700 100644 --- a/tests/Meziantou.Analyzer.Test/Rules/OptimizeLinqUsageAnalyzer.DuplicateOrderBy.Tests.cs +++ b/tests/Meziantou.Analyzer.Test/Rules/OptimizeLinqUsageAnalyzer.DuplicateOrderBy.Tests.cs @@ -34,6 +34,30 @@ public Test() enumerable." + a + @"(x => x)." + b + @"(x => x); } } +"); + + VerifyDiagnostic(project, CreateDiagnosticResult(line: 7, column: 9, message: $"Remove the first '{a}' method or use '{expectedMethod}'")); + } + + [DataTestMethod] + [DataRow("ThenBy", "OrderBy", "ThenBy")] + [DataRow("ThenByDescending", "OrderBy", "ThenBy")] + [DataRow("ThenBy", "OrderByDescending", "ThenByDescending")] + [DataRow("ThenByDescending", "OrderByDescending", "ThenByDescending")] + public void ThenByFollowedByOrderBy(string a, string b, string expectedMethod) + { + var project = new ProjectBuilder() + .AddReference(typeof(IEnumerable<>)) + .AddReference(typeof(Enumerable)) + .WithSource(@"using System.Linq; +class Test +{ + public Test() + { + var enumerable = System.Linq.Enumerable.Empty(); + enumerable.OrderBy(x => x)." + a + @"(x => x)." + b + @"(x => x); + } +} "); VerifyDiagnostic(project, CreateDiagnosticResult(line: 7, column: 9, message: $"Remove the first '{a}' method or use '{expectedMethod}'"));