diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs index c1bbaa06d..05cef3c9f 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs @@ -171,22 +171,30 @@ private static bool IsInTaskReturningMethodOrDelegate(SyntaxNodeAnalysisContext // We want to scan invocations that occur inside Task and Task-returning delegates or methods. // That is: methods that either are or could be made async. IMethodSymbol? methodSymbol = null; - AnonymousFunctionExpressionSyntax? anonymousFunc = context.Node.FirstAncestorOrSelf(); - if (anonymousFunc is object) + for (SyntaxNode? focusedNode = context.Node; focusedNode is not null; focusedNode = focusedNode.Parent) { - SymbolInfo symbolInfo = context.SemanticModel.GetSymbolInfo(anonymousFunc, context.CancellationToken); - methodSymbol = symbolInfo.Symbol as IMethodSymbol; - } - else - { - MethodDeclarationSyntax? methodDecl = context.Node.FirstAncestorOrSelf(); - if (methodDecl is object) + switch (focusedNode) { - methodSymbol = context.SemanticModel.GetDeclaredSymbol(methodDecl, context.CancellationToken); + case AnonymousFunctionExpressionSyntax anonFunc: + SymbolInfo symbolInfo = context.SemanticModel.GetSymbolInfo(anonFunc, context.CancellationToken); + methodSymbol = symbolInfo.Symbol as IMethodSymbol; + break; + case LocalFunctionStatementSyntax localFunc: + methodSymbol = context.SemanticModel.GetDeclaredSymbol(localFunc, context.CancellationToken) as IMethodSymbol; + break; + case MethodDeclarationSyntax methodDecl: + methodSymbol = context.SemanticModel.GetDeclaredSymbol(methodDecl, context.CancellationToken); + break; + default: + // We want to continue iteration of the for loop. + continue; } + + // We encountered one of our case statements, so whether or not we have a methodSymbol, we shouldn't look further. + break; } - return methodSymbol.HasAsyncCompatibleReturnType(); + return methodSymbol?.HasAsyncCompatibleReturnType() is true; } private static bool InspectMemberAccess(SyntaxNodeAnalysisContext context, ExpressionSyntax memberName, IEnumerable problematicMethods) diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs index b665feb5d..b4ff2da62 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs @@ -1314,6 +1314,31 @@ Task MethodAsync() await CSVerify.VerifyAnalyzerAsync(test); } + [Fact] + public async Task DoNotRaiseInSyncLocalFunctionInsideAsyncMethod() + { + string test = """ + using System.Threading.Tasks; + + class SomeClass { + Task Foo() + { + return Task.CompletedTask; + + void CompletionHandler() + { + this.Bar(); + } + } + + void Bar() {} + Task BarAsync() => Task.CompletedTask; + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + private DiagnosticResult CreateDiagnostic(int line, int column, int length, string methodName) => CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithSpan(line, column, line, column + length).WithArguments(methodName); diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD110ObserveResultOfAsyncCallsAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD110ObserveResultOfAsyncCallsAnalyzerTests.cs index 5ee67c6fb..bae3ec25c 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD110ObserveResultOfAsyncCallsAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD110ObserveResultOfAsyncCallsAnalyzerTests.cs @@ -429,6 +429,21 @@ async Task Foo(Test? tester) await CSVerify.VerifyAnalyzerAsync(test); } + [Fact] + public async Task NullCoalescing_ProducesNoDiagnostic() + { + string test = """ + using System.Threading.Tasks; + + class Tree { + static Task ShakeTreeAsync(Tree? tree) => tree?.ShakeAsync() ?? Task.CompletedTask; + Task ShakeAsync() => Task.CompletedTask; + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + [Fact] public async Task TaskInFinalizer() { @@ -470,7 +485,4 @@ class Class1 await CSVerify.VerifyAnalyzerAsync(test); } - - private DiagnosticResult CreateDiagnostic(int line, int column, int length) - => CSVerify.Diagnostic().WithSpan(line, column, line, column + length); }