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
1 change: 1 addition & 0 deletions src/Meziantou.Analyzer.CodeFixers/Directory.Build.props
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
<Compile Include="$(MSBuildThisFileDirectory)\..\Meziantou.Analyzer\Internals\SymbolExtensions.cs" LinkBase="Internals" />
<Compile Include="$(MSBuildThisFileDirectory)\..\Meziantou.Analyzer\Internals\OverloadFinder.cs" LinkBase="Internals" />
<Compile Include="$(MSBuildThisFileDirectory)\..\Meziantou.Analyzer\Internals\OverloadOptions.cs" LinkBase="Internals" />
<Compile Include="$(MSBuildThisFileDirectory)\..\Meziantou.Analyzer\Internals\OverloadLookupCache.cs" LinkBase="Internals" />
<Compile Include="$(MSBuildThisFileDirectory)\..\Meziantou.Analyzer\Internals\OverloadParameterType.cs" LinkBase="Internals" />
<Compile Include="$(MSBuildThisFileDirectory)\..\Meziantou.Analyzer\Rules\DoNotUseBlockingCallInAsyncContextData.cs" LinkBase="Rules" />
<Compile Include="$(MSBuildThisFileDirectory)\..\Meziantou.Analyzer\Rules\OptimizeLinqUsageData.cs" LinkBase="Rules" />
Expand Down
12 changes: 10 additions & 2 deletions src/Meziantou.Analyzer/Internals/OverloadFinder.cs
Original file line number Diff line number Diff line change
Expand Up @@ -671,10 +671,10 @@ static void AddSymbols(IEnumerable<ISymbol> symbols, List<ISymbol> results, Hash
var semanticModel = compilation.GetSemanticModel(options.SyntaxNode.SyntaxTree);
var position = options.SyntaxNode.GetLocation().SourceSpan.End;

AddSymbols(semanticModel.LookupSymbols(position, methodSymbol.ContainingType, methodName, includeReducedExtensionMethods: true), results, knownSymbols);
AddSymbols(LookupSymbols(semanticModel, options, position, methodSymbol.ContainingType, methodName, includeReducedExtensionMethods: true), results, knownSymbols);
if (reducedReceiverType is not null)
{
AddSymbols(semanticModel.LookupSymbols(position, reducedReceiverType, methodName, includeReducedExtensionMethods: false), results, knownSymbols);
AddSymbols(LookupSymbols(semanticModel, options, position, reducedReceiverType, methodName, includeReducedExtensionMethods: false), results, knownSymbols);
AddSymbols(reducedReceiverType.GetMembers(methodName), results, knownSymbols);
}

Expand All @@ -695,6 +695,14 @@ static void AddSymbols(IEnumerable<ISymbol> symbols, List<ISymbol> results, Hash
return results;
}

private static ImmutableArray<ISymbol> LookupSymbols(SemanticModel semanticModel, OverloadOptions options, int position, ITypeSymbol container, string name, bool includeReducedExtensionMethods)
{
if (options.LookupCache is null)
return semanticModel.LookupSymbols(position, container, name, includeReducedExtensionMethods);

return options.LookupCache.LookupSymbols(semanticModel, options.SyntaxNode!, position, container, name, includeReducedExtensionMethods);
}

/// <summary>
/// Adds the extension methods that apply to the receiver of <paramref name="methodSymbol"/>, are accessible at
/// <paramref name="position"/>, and are not in scope because their namespace is not imported. They are added after the
Expand Down
57 changes: 57 additions & 0 deletions src/Meziantou.Analyzer/Internals/OverloadLookupCache.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,57 @@
using System.Collections.Concurrent;
using System.Runtime.CompilerServices;
using Microsoft.CodeAnalysis.CSharp.Syntax;

namespace Meziantou.Analyzer.Internals;

/// <summary>
/// Caches the member lookups done by <see cref="OverloadFinder"/>. A cache is meant to be created when the analysis of a named
/// type starts (<c>RegisterSymbolStartAction</c>) and shared by the actions of that type, so it is released with the type.
/// It is thread-safe because the members of a type can be analyzed concurrently.
/// </summary>
internal sealed class OverloadLookupCache
{
private readonly ConcurrentDictionary<LookupKey, ImmutableArray<ISymbol>> _cache = new(LookupKeyComparer.Instance);

// A member lookup only depends on the file and enclosing type declaration (accessibility, imported namespaces), the container and the name
public ImmutableArray<ISymbol> LookupSymbols(SemanticModel semanticModel, SyntaxNode node, int position, ITypeSymbol container, string name, bool includeReducedExtensionMethods)
{
var scopeStart = node.FirstAncestorOrSelf<BaseTypeDeclarationSyntax>()?.SpanStart ?? -1;
var key = new LookupKey(node.SyntaxTree, scopeStart, container, name, includeReducedExtensionMethods);
if (!_cache.TryGetValue(key, out var symbols))
{
symbols = semanticModel.LookupSymbols(position, container, name, includeReducedExtensionMethods);
_ = _cache.TryAdd(key, symbols);
}

return symbols;
}

private readonly record struct LookupKey(SyntaxTree Tree, int ScopeStart, ITypeSymbol Container, string Name, bool IncludeReducedExtensionMethods);

private sealed class LookupKeyComparer : IEqualityComparer<LookupKey>
{
public static LookupKeyComparer Instance { get; } = new();

public bool Equals(LookupKey x, LookupKey y)
{
return ReferenceEquals(x.Tree, y.Tree)
&& x.ScopeStart == y.ScopeStart
&& SymbolEqualityComparer.Default.Equals(x.Container, y.Container)
&& string.Equals(x.Name, y.Name, StringComparison.Ordinal)
&& x.IncludeReducedExtensionMethods == y.IncludeReducedExtensionMethods;
}

public int GetHashCode(LookupKey obj)
{
unchecked
{
var hashCode = RuntimeHelpers.GetHashCode(obj.Tree);
hashCode = (hashCode * 397) ^ obj.ScopeStart;
hashCode = (hashCode * 397) ^ SymbolEqualityComparer.Default.GetHashCode(obj.Container);
hashCode = (hashCode * 397) ^ StringComparer.Ordinal.GetHashCode(obj.Name);
return (hashCode * 397) ^ (obj.IncludeReducedExtensionMethods ? 1 : 0);
}
}
}
}
3 changes: 2 additions & 1 deletion src/Meziantou.Analyzer/Internals/OverloadOptions.cs
Original file line number Diff line number Diff line change
Expand Up @@ -12,4 +12,5 @@ internal record struct OverloadOptions(
bool AllowInterfaceConversions = true,
bool AllowBaseTypeConversions = true,
Func<IMethodSymbol, bool>? ShouldCheckMethod = null,
bool IncludeExtensionMethodsFromNotImportedNamespaces = false);
bool IncludeExtensionMethodsFromNotImportedNamespaces = false,
OverloadLookupCache? LookupCache = null);
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,11 @@ public override void Initialize(AnalysisContext context)
var analyzerContext = new Context(ctx.Compilation);
if (analyzerContext.IsValid)
{
ctx.RegisterOperationAction(analyzerContext.AnalyzeInvocation, OperationKind.Invocation);
ctx.RegisterSymbolStartAction(symbolContext =>
{
var lookupCache = new OverloadLookupCache();
symbolContext.RegisterOperationAction(operationContext => analyzerContext.AnalyzeInvocation(operationContext, lookupCache), OperationKind.Invocation);
}, SymbolKind.NamedType);
ctx.RegisterOperationAction(analyzerContext.AnalyzePropertyReference, OperationKind.PropertyReference);
ctx.RegisterOperationAction(analyzerContext.AnalyzeUsing, OperationKind.Using);
ctx.RegisterOperationAction(analyzerContext.AnalyzeUsingDeclaration, OperationKind.UsingDeclaration);
Expand Down Expand Up @@ -205,7 +209,7 @@ public Context(Compilation compilation)

public bool IsValid => TaskSymbol is not null && TaskOfTSymbol is not null && TaskAwaiterSymbol is not null;

internal void AnalyzeInvocation(OperationAnalysisContext context)
internal void AnalyzeInvocation(OperationAnalysisContext context, OverloadLookupCache lookupCache)
{
var operation = (IInvocationOperation)context.Operation;

Expand Down Expand Up @@ -240,7 +244,7 @@ internal void AnalyzeInvocation(OperationAnalysisContext context)
var sqliteSpecialCasesEnabled = IsSqliteSpecialCasesEnabled(context, operation);
var includeExtensionMethodsFromNotImportedNamespaces = context.Options.GetConfigurationValue(operation, IncludeExtensionMethodsFromNotImportedNamespacesInAsyncContextConfiguration)
|| context.Options.GetConfigurationValue(operation, IncludeExtensionMethodsFromNotImportedNamespacesConfiguration);
var result = FindAsyncEquivalent(operation, sqliteSpecialCasesEnabled, includeExtensionMethodsFromNotImportedNamespaces, context.CancellationToken, out var diagnosticMessage);
var result = FindAsyncEquivalent(operation, sqliteSpecialCasesEnabled, includeExtensionMethodsFromNotImportedNamespaces, lookupCache, context.CancellationToken, out var diagnosticMessage);
if (diagnosticMessage is not null)
{
ReportDiagnosticIfNeeded(context, diagnosticMessage.CreateProperties(), operation, diagnosticMessage.DiagnosticMessage, asyncContextKind, requiresNamespaceImport: diagnosticMessage.NamespaceToImport is not null);
Expand All @@ -256,7 +260,7 @@ internal void AnalyzeInvocation(OperationAnalysisContext context)
/// Searches for an async equivalent of the called method, visible from the call site. <paramref name="data"/>
/// is set when the result is <see cref="AsyncEquivalentSearchResult.Found"/>.
/// </summary>
private AsyncEquivalentSearchResult FindAsyncEquivalent(IInvocationOperation operation, bool sqliteSpecialCasesEnabled, bool includeExtensionMethodsFromNotImportedNamespaces, CancellationToken cancellationToken, out DiagnosticData? data)
private AsyncEquivalentSearchResult FindAsyncEquivalent(IInvocationOperation operation, bool sqliteSpecialCasesEnabled, bool includeExtensionMethodsFromNotImportedNamespaces, OverloadLookupCache lookupCache, CancellationToken cancellationToken, out DiagnosticData? data)
{
data = null;
var targetMethod = operation.TargetMethod;
Expand Down Expand Up @@ -366,10 +370,10 @@ private AsyncEquivalentSearchResult FindAsyncEquivalent(IInvocationOperation ope
// as the code fix must add a using directive to call them
IMethodSymbol? notImportedAsyncEquivalentMethod = null;
string? namespaceToImport = null;
var asyncEquivalentMethod = FindPotentialAsyncEquivalent(operation, targetMethod, targetMethod.Name, includeExtensionMethodsFromNotImportedNamespaces, ref notImportedAsyncEquivalentMethod, ref namespaceToImport);
var asyncEquivalentMethod = FindPotentialAsyncEquivalent(operation, targetMethod, targetMethod.Name, includeExtensionMethodsFromNotImportedNamespaces, lookupCache, ref notImportedAsyncEquivalentMethod, ref namespaceToImport);
if (asyncEquivalentMethod is null && !targetMethod.Name.EndsWith("Async", StringComparison.Ordinal))
{
asyncEquivalentMethod = FindPotentialAsyncEquivalent(operation, targetMethod, targetMethod.Name + "Async", includeExtensionMethodsFromNotImportedNamespaces, ref notImportedAsyncEquivalentMethod, ref namespaceToImport);
asyncEquivalentMethod = FindPotentialAsyncEquivalent(operation, targetMethod, targetMethod.Name + "Async", includeExtensionMethodsFromNotImportedNamespaces, lookupCache, ref notImportedAsyncEquivalentMethod, ref namespaceToImport);
}

if (asyncEquivalentMethod is null)
Expand Down Expand Up @@ -519,13 +523,14 @@ private bool IsSqliteSpecialCaseMethod(IInvocationOperation operation, Cancellat
/// is set, the first async equivalent that requires to import its namespace is set to <paramref name="notImportedMethod"/>, and its
/// namespace to <paramref name="namespaceToImport"/>, if they are not already set.
/// </summary>
private IMethodSymbol? FindPotentialAsyncEquivalent(IInvocationOperation operation, IMethodSymbol targetMethod, string methodName, bool includeExtensionMethodsFromNotImportedNamespaces, ref IMethodSymbol? notImportedMethod, ref string? namespaceToImport)
private IMethodSymbol? FindPotentialAsyncEquivalent(IInvocationOperation operation, IMethodSymbol targetMethod, string methodName, bool includeExtensionMethodsFromNotImportedNamespaces, OverloadLookupCache lookupCache, ref IMethodSymbol? notImportedMethod, ref string? namespaceToImport)
{
var options = new OverloadOptions(
AllowOptionalParameters: false,
IncludeExtensionsMethods: true,
SyntaxNode: operation.Syntax,
IncludeExtensionMethodsFromNotImportedNamespaces: includeExtensionMethodsFromNotImportedNamespaces);
IncludeExtensionMethodsFromNotImportedNamespaces: includeExtensionMethodsFromNotImportedNamespaces,
LookupCache: lookupCache);

// When the method name is the same as the original method, and the original is non-generic
// while a candidate is generic, the compiler will always prefer the non-generic original
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -63,8 +63,12 @@ public override void Initialize(AnalysisContext context)
if (analyzerContext.CancellationTokenSymbol is null)
return;

ctx.RegisterOperationAction(analyzerContext.AnalyzeInvocation, OperationKind.Invocation);
ctx.RegisterOperationAction(analyzerContext.AnalyzeLoop, OperationKind.Loop);
ctx.RegisterSymbolStartAction(symbolContext =>
{
var lookupCache = new OverloadLookupCache();
symbolContext.RegisterOperationAction(context => analyzerContext.AnalyzeInvocation(context, lookupCache), OperationKind.Invocation);
symbolContext.RegisterOperationAction(context => analyzerContext.AnalyzeLoop(context, lookupCache), OperationKind.Loop);
}, SymbolKind.NamedType);
});
}

Expand Down Expand Up @@ -132,7 +136,7 @@ private bool HasExplicitCancellationTokenArgument(IInvocationOperation operation
/// <param name="NamespaceToImport">The namespace to import to call the overload, when it is an extension method declared in a namespace that is not imported.</param>
private sealed record AdditionalParameterInfo(int ParameterIndex, string? Name, bool HasEnumeratorCancellationAttribute, string? NamespaceToImport = null);

private bool HasAnOverloadWithCancellationToken(OperationAnalysisContext context, IInvocationOperation operation, [NotNullWhen(true)] out AdditionalParameterInfo? parameterInfo)
private bool HasAnOverloadWithCancellationToken(OperationAnalysisContext context, IInvocationOperation operation, OverloadLookupCache lookupCache, [NotNullWhen(true)] out AdditionalParameterInfo? parameterInfo)
{
parameterInfo = default;
var method = operation.TargetMethod;
Expand All @@ -145,7 +149,7 @@ private bool HasAnOverloadWithCancellationToken(OperationAnalysisContext context
var allowOptionalParameters = context.Options.GetConfigurationValue(operation, AllowOverloadsWithOptionalParametersConfiguration);
var includeExtensionMethodsFromNotImportedNamespaces = context.Options.GetConfigurationValue(operation, IncludeExtensionMethodsFromNotImportedNamespacesConfiguration)
|| context.Options.GetConfigurationValue(operation, IncludeExtensionMethodsFromNotImportedNamespacesWhenACancellationTokenIsAvailableConfiguration);
var overload = _overloadFinder.FindOverloadWithAdditionalParameterOfType(operation, new OverloadOptions(AllowOptionalParameters: allowOptionalParameters, IncludeExtensionMethodsFromNotImportedNamespaces: includeExtensionMethodsFromNotImportedNamespaces), [CancellationTokenSymbol]);
var overload = _overloadFinder.FindOverloadWithAdditionalParameterOfType(operation, new OverloadOptions(AllowOptionalParameters: allowOptionalParameters, IncludeExtensionMethodsFromNotImportedNamespaces: includeExtensionMethodsFromNotImportedNamespaces, LookupCache: lookupCache), [CancellationTokenSymbol]);
if (overload is not null)
{
var namespaceToImport = includeExtensionMethodsFromNotImportedNamespaces ? _overloadFinder.GetNamespaceToImport(overload, operation.Syntax) : null;
Expand Down Expand Up @@ -188,13 +192,13 @@ private bool HasEnumerableCancellationAttribute(IParameterSymbol? parameterSymbo
return parameterSymbol.HasAttribute(EnumeratorCancellationAttributeSymbol, inherits: false);
}

public void AnalyzeInvocation(OperationAnalysisContext context)
public void AnalyzeInvocation(OperationAnalysisContext context, OverloadLookupCache lookupCache)
{
var operation = (IInvocationOperation)context.Operation;
if (HasExplicitCancellationTokenArgument(operation))
return;

if (!HasAnOverloadWithCancellationToken(context, operation, out var parameterInfo))
if (!HasAnOverloadWithCancellationToken(context, operation, lookupCache, out var parameterInfo))
return;

if (IsExcluded(operation.TargetMethod))
Expand Down Expand Up @@ -231,7 +235,7 @@ private static bool IsOverloadIncluded(OperationAnalysisContext context, IOperat
return context.Options.GetConfigurationValue(operation, configuration);
}

public void AnalyzeLoop(OperationAnalysisContext context)
public void AnalyzeLoop(OperationAnalysisContext context, OverloadLookupCache lookupCache)
{
if (context.Operation is not IForEachLoopOperation op)
return;
Expand Down Expand Up @@ -261,7 +265,7 @@ public void AnalyzeLoop(OperationAnalysisContext context)
return;

// Already handled by AnalyzeInvocation
if (HasAnOverloadWithCancellationToken(context, invocation, out var invocationParameterInfo) &&
if (HasAnOverloadWithCancellationToken(context, invocation, lookupCache, out var invocationParameterInfo) &&
(invocationParameterInfo.NamespaceToImport is null || IsOverloadIncluded(context, invocation, invocationParameterInfo, hasAvailableCancellationTokens: _cancellationTokenFinder.FindPaths(invocation, context.CancellationToken).Length > 0)))
{
return;
Expand Down
Loading
Loading