diff --git a/docfx/analyzers/VSTHRD002.md b/docfx/analyzers/VSTHRD002.md index 6a05a94ca..69b021b93 100644 --- a/docfx/analyzers/VSTHRD002.md +++ b/docfx/analyzers/VSTHRD002.md @@ -40,6 +40,14 @@ void DoSomething() } ``` +Accessing a task's result is not reported when the analyzer can prove the task has completed. +Recognized proofs include awaiting the task (directly or through `Task.WhenAll`), guarding the +access with a completion property such as `IsCompletedSuccessfully`, and awaiting the task in +the negative branch of such a guard. + +VSTHRD002 can also report project-specific synchronous blocking methods configured in +`vs-threading.SyncBlockingMethods.txt`. See [Analyzer Configuration](configuration.md#additional-synchronous-blocking-methods-for-vsthrd002). + Refer to [Asynchronous and multithreaded programming within VS using the JoinableTaskFactory][1] for more information. [1]: https://devblogs.microsoft.com/premier-developer/asynchronous-and-multithreaded-programming-within-vs-using-the-joinabletaskfactory/ diff --git a/docfx/analyzers/configuration.md b/docfx/analyzers/configuration.md index c21d4c7ab..13b3eede3 100644 --- a/docfx/analyzers/configuration.md +++ b/docfx/analyzers/configuration.md @@ -105,6 +105,19 @@ excluded from VSTHRD103 analysis by specifying them in a configuration file. **Generic sample:** ``[Microsoft.EntityFrameworkCore.DbSet`1]::Add`` +## Additional synchronous blocking methods for VSTHRD002 + +Projects that wrap synchronous waits in their own APIs can configure those methods to be +reported by VSTHRD002. Instance, static, and extension methods are supported. Because the +analyzer cannot infer an asynchronous equivalent for a configured method, it does not offer +the "use await instead" code fix for these diagnostics. + +**Filename:** `vs-threading.SyncBlockingMethods.txt` + +**Line format:** `[Namespace.TypeName]::MethodName` + +**Sample:** `[Contoso.Threading.TaskExtensions]::WaitSynchronously` + ## Types that require the Async suffix VSTHRD200 requires methods returning `Task`, `ValueTask`, and other async-focused types diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index bf159a90e..99c2f6b1f 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -4,12 +4,15 @@ using System; using System.Collections.Generic; using System.Collections.Immutable; +using System.Diagnostics.CodeAnalysis; using System.Linq; using System.Text; +using System.Threading.Tasks; using Microsoft.CodeAnalysis; using Microsoft.CodeAnalysis.CSharp; using Microsoft.CodeAnalysis.CSharp.Syntax; using Microsoft.CodeAnalysis.Diagnostics; +using Microsoft.CodeAnalysis.Operations; namespace Microsoft.VisualStudio.Threading.Analyzers; @@ -28,6 +31,164 @@ internal static class CSharpCommonInterest SyntaxKind.AddAccessorDeclaration, SyntaxKind.RemoveAccessorDeclaration); + /// + /// Gets a symbol and ref locals that definitely or potentially alias it at the specified syntax node. + /// + internal static (ImmutableHashSet Definite, ImmutableHashSet Potential) GetSymbolAndRefAliases( + SyntaxNodeAnalysisContext context, + SyntaxNode node, + ISymbol symbol, + SyntaxNode? aliasSearchRoot = null, + bool includeAllCandidates = false) + { + SyntaxNode searchRoot = aliasSearchRoot ?? node.AncestorsAndSelf().FirstOrDefault( + ancestor => ancestor is AnonymousFunctionExpressionSyntax + or LocalFunctionStatementSyntax + or BaseMethodDeclarationSyntax + or AccessorDeclarationSyntax) + ?? node.FirstAncestorOrSelf()?.Parent + ?? node; + ITypeSymbol? trackedType = symbol switch + { + ILocalSymbol local => local.Type, + IParameterSymbol parameter => parameter.Type, + IFieldSymbol field => field.Type, + _ => null, + }; + + var refTargets = new Dictionary>(SymbolEqualityComparer.Default); + var potentialOnlyRefLocals = new HashSet(SymbolEqualityComparer.Default); + bool DescendIntoChildren(SyntaxNode child) => + child == searchRoot || child is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; + + HashSet GetRefTargets(ISymbol candidate) + { + if (refTargets.TryGetValue(candidate, out HashSet? targets)) + { + return new HashSet(targets, SymbolEqualityComparer.Default); + } + + return new HashSet(SymbolEqualityComparer.Default) { candidate }; + } + + bool DefinitelyPrecedesNode(SyntaxNode candidate) + { + StatementSyntax? candidateStatement = candidate.FirstAncestorOrSelf(); + if (candidateStatement?.Parent is BlockSyntax block) + { + StatementSyntax? nodeStatement = node.AncestorsAndSelf().OfType().FirstOrDefault(statement => statement.Parent == block); + return nodeStatement is object && block.Statements.IndexOf(candidateStatement) < block.Statements.IndexOf(nodeStatement); + } + + GlobalStatementSyntax? candidateGlobalStatement = candidate.FirstAncestorOrSelf(); + GlobalStatementSyntax? nodeGlobalStatement = node.FirstAncestorOrSelf(); + return candidateGlobalStatement?.Parent is CompilationUnitSyntax compilationUnit + && nodeGlobalStatement?.Parent == compilationUnit + && compilationUnit.Members.IndexOf(candidateGlobalStatement) < compilationUnit.Members.IndexOf(nodeGlobalStatement); + } + + foreach (SyntaxNode candidate in searchRoot.DescendantNodes(DescendIntoChildren) + .Where(candidate => (includeAllCandidates || candidate.SpanStart < node.SpanStart) + && candidate is VariableDeclaratorSyntax or AssignmentExpressionSyntax) + .OrderBy(candidate => candidate.SpanStart)) + { + if (candidate is VariableDeclaratorSyntax variable + && variable.Initializer is not null + && context.SemanticModel.GetDeclaredSymbol(variable, context.CancellationToken) is ILocalSymbol { RefKind: not RefKind.None } local) + { + ExpressionSyntax initializer = variable.Initializer.Value is RefExpressionSyntax refInitializer + ? refInitializer.Expression + : variable.Initializer.Value; + if (context.SemanticModel.GetSymbolInfo(UnwrapParentheses(initializer), context.CancellationToken).Symbol is ISymbol initializedFrom) + { + if (initializedFrom is IMethodSymbol refReturningMethod + && (refReturningMethod.ReturnsByRef || refReturningMethod.ReturnsByRefReadonly) + && SymbolEqualityComparer.Default.Equals(refReturningMethod.ReturnType, trackedType)) + { + refTargets[local] = new HashSet(SymbolEqualityComparer.Default) { symbol }; + potentialOnlyRefLocals.Add(local); + } + else + { + refTargets[local] = GetRefTargets(initializedFrom); + if (potentialOnlyRefLocals.Contains(initializedFrom)) + { + potentialOnlyRefLocals.Add(local); + } + } + } + else if (MayAliasTaskStorage(context, initializer, ImmutableHashSet.Create(SymbolEqualityComparer.Default, symbol))) + { + refTargets[local] = new HashSet(SymbolEqualityComparer.Default) { symbol }; + potentialOnlyRefLocals.Add(local); + } + } + else if (candidate is AssignmentExpressionSyntax { Right: RefExpressionSyntax refAssignment } assignment + && context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol is ILocalSymbol { RefKind: not RefKind.None } reboundLocal) + { + ExpressionSyntax assignedExpression = UnwrapParentheses(refAssignment.Expression); + ISymbol? assignedFrom = context.SemanticModel.GetSymbolInfo(assignedExpression, context.CancellationToken).Symbol; + HashSet? assignedTargets = assignedFrom is object + ? GetRefTargets(assignedFrom) + : MayAliasTaskStorage(context, assignedExpression, ImmutableHashSet.Create(SymbolEqualityComparer.Default, symbol)) + ? new HashSet(SymbolEqualityComparer.Default) { symbol } + : null; + if (assignedTargets is object) + { + if ((!includeAllCandidates && DefinitelyPrecedesNode(candidate)) + || !refTargets.TryGetValue(reboundLocal, out HashSet? existingTargets)) + { + refTargets[reboundLocal] = assignedTargets; + if (assignedFrom is null || potentialOnlyRefLocals.Contains(assignedFrom)) + { + potentialOnlyRefLocals.Add(reboundLocal); + } + else + { + potentialOnlyRefLocals.Remove(reboundLocal); + } + } + else + { + existingTargets.UnionWith(assignedTargets); + if (assignedFrom is null || potentialOnlyRefLocals.Contains(assignedFrom)) + { + potentialOnlyRefLocals.Add(reboundLocal); + } + } + } + } + } + + HashSet symbolTargets = GetRefTargets(symbol); + ImmutableHashSet.Builder definiteSymbols = ImmutableHashSet.CreateBuilder(SymbolEqualityComparer.Default); + ImmutableHashSet.Builder potentialSymbols = ImmutableHashSet.CreateBuilder(SymbolEqualityComparer.Default); + definiteSymbols.Add(symbol); + potentialSymbols.Add(symbol); + potentialSymbols.UnionWith(symbolTargets); + if (symbolTargets.Count == 1) + { + definiteSymbols.UnionWith(symbolTargets); + } + + foreach (KeyValuePair> refTarget in refTargets) + { + if (symbolTargets.Count == 1 + && refTarget.Value.SetEquals(symbolTargets) + && !potentialOnlyRefLocals.Contains(refTarget.Key)) + { + definiteSymbols.Add(refTarget.Key); + } + + if (refTarget.Value.Overlaps(symbolTargets)) + { + potentialSymbols.Add(refTarget.Key); + } + } + + return (definiteSymbols.ToImmutable(), potentialSymbols.ToImmutable()); + } + /// /// This is an explicit rule to ignore the code that was generated by Xaml2CS. /// @@ -94,13 +255,17 @@ internal static void InspectMemberAccess( } ITypeSymbol? typeReceiver = context.SemanticModel.GetTypeInfo(memberAccessSyntax.Expression).Type; - if (typeReceiver is object) + ISymbol? accessedSymbol = memberAccessSyntax.Parent is InvocationExpressionSyntax invocation + ? context.SemanticModel.GetSymbolInfo(invocation, context.CancellationToken).Symbol + : context.SemanticModel.GetSymbolInfo(memberAccessSyntax, context.CancellationToken).Symbol; + if (typeReceiver is object && accessedSymbol is object) { foreach (CommonInterest.SyncBlockingMethod item in problematicMethods) { if (memberAccessSyntax.Name.Identifier.Text == item.Method.Name && typeReceiver.Name == item.Method.ContainingType.Name && - typeReceiver.BelongsToNamespace(item.Method.ContainingType.Namespace)) + typeReceiver.BelongsToNamespace(item.Method.ContainingType.Namespace) && + IsBuiltInBlockingMember(context, accessedSymbol, item.Method)) { if (HasTaskCompleted(context, memberAccessSyntax)) { @@ -114,44 +279,193 @@ internal static void InspectMemberAccess( } } - private static SyntaxNode? GetEnclosingBlock(SyntaxNode? node) + /// + /// Inspects a conditionally accessed member for configured or built-in synchronous blocking behavior. + /// + /// The syntax analysis context. + /// The member binding to inspect. + /// The expression receiving the conditional access. + /// The complete conditional access expression. + /// The diagnostic descriptor to report. + /// The synchronous blocking members recognized by the analyzer. + internal static void InspectMemberBinding( + SyntaxNodeAnalysisContext context, + MemberBindingExpressionSyntax memberBinding, + ExpressionSyntax receiver, + SyntaxNode accessSyntax, + DiagnosticDescriptor descriptor, + IEnumerable problematicMethods) { - while (node is not null) + if (descriptor is null) + { + throw new ArgumentNullException(nameof(descriptor)); + } + + if (ShouldIgnoreContext(context) || CSharpUtils.IsWithinNameOf(context.Node as ExpressionSyntax)) + { + return; + } + + ITypeSymbol? receiverType = context.SemanticModel.GetTypeInfo(receiver, context.CancellationToken).Type; + ISymbol? accessedSymbol = memberBinding.Parent is InvocationExpressionSyntax invocation + ? context.SemanticModel.GetSymbolInfo(invocation, context.CancellationToken).Symbol + : context.SemanticModel.GetSymbolInfo(memberBinding, context.CancellationToken).Symbol; + if (receiverType is null || accessedSymbol is null) + { + return; + } + + foreach (CommonInterest.SyncBlockingMethod item in problematicMethods) { - if (node.IsKind(SyntaxKind.Block)) + if (memberBinding.Name.Identifier.ValueText == item.Method.Name + && receiverType.Name == item.Method.ContainingType.Name + && receiverType.BelongsToNamespace(item.Method.ContainingType.Namespace) + && IsBuiltInBlockingMember(context, accessedSymbol, item.Method)) { - return node; - } + if (HasTaskCompleted(context, receiver, accessSyntax)) + { + return; + } - node = node.Parent; + context.ReportDiagnostic(Diagnostic.Create(descriptor, memberBinding.Name.GetLocation())); + } } + } - return null; + /// + /// Gets the symbol represented by the normalized task-like receiver of a blocking member access. + /// + internal static ISymbol? GetTaskReceiverSymbol(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) + { + ExpressionSyntax receiver = GetTaskReceiver(context, memberAccessSyntax); + return context.SemanticModel.GetSymbolInfo(receiver, context.CancellationToken).Symbol; } - private static bool IsVariablePassedToInvocation(InvocationExpressionSyntax invocationExpr, string variableName, bool byRef) + /// + /// Determines whether an async alternative is applicable to the arguments of a synchronous invocation. + /// + internal static bool IsApplicableAsyncAlternative( + SyntaxNodeAnalysisContext context, + InvocationExpressionSyntax invocation, + IMethodSymbol candidateMethod) { - ArgumentListSyntax? argList = invocationExpr.ChildNodes().OfType().FirstOrDefault(); - if (argList is null) + SimpleNameSyntax? invokedName = invocation.Expression switch + { + MemberAccessExpressionSyntax memberAccess => memberAccess.Name, + MemberBindingExpressionSyntax memberBinding => memberBinding.Name, + SimpleNameSyntax simpleName => simpleName, + _ => null, + }; + if (invokedName is null) + { + return false; + } + + SyntaxToken newIdentifier = SyntaxFactory.Identifier( + invokedName.Identifier.LeadingTrivia, + candidateMethod.Name, + invokedName.Identifier.TrailingTrivia); + SimpleNameSyntax asyncName = (SimpleNameSyntax)invokedName.ReplaceToken(invokedName.Identifier, newIdentifier); + InvocationExpressionSyntax asyncInvocation = invocation.ReplaceNode(invokedName, asyncName); + + ExpressionSyntax speculativeExpression = asyncInvocation; + if (invocation.Expression is MemberBindingExpressionSyntax + && invocation.FirstAncestorOrSelf() is { } conditionalAccess) + { + speculativeExpression = conditionalAccess.ReplaceNode(invocation, asyncInvocation); + } + + ExpressionSyntax detachedSpeculativeExpression = SyntaxFactory.ParseExpression(speculativeExpression.ToString()); + SymbolInfo speculativeSymbolInfo = context.SemanticModel.GetSpeculativeSymbolInfo( + invocation.SpanStart, + detachedSpeculativeExpression, + SpeculativeBindingOption.BindAsExpression); + if (speculativeSymbolInfo.Symbol is not IMethodSymbol applicableMethod) { return false; } - foreach (ArgumentSyntax arg in argList.ChildNodes().OfType()) + IMethodSymbol applicableDefinition = (applicableMethod.ReducedFrom ?? applicableMethod).OriginalDefinition; + IMethodSymbol candidateDefinition = (candidateMethod.ReducedFrom ?? candidateMethod).OriginalDefinition; + return SymbolEqualityComparer.Default.Equals(applicableDefinition, candidateDefinition); + } + + /// + /// Determines whether a blocking member access has a receiver that is provably complete. + /// + internal static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) + { + ExpressionSyntax taskReceiver = GetTaskReceiver(context, memberAccessSyntax); + return HasTaskCompleted(context, taskReceiver, memberAccessSyntax); + } + + /// + /// Determines whether a task-like expression is provably complete at a syntax node. + /// + internal static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, ExpressionSyntax taskReceiver, SyntaxNode accessSyntax) + => HasTaskCompletedInContinuation(context, taskReceiver, accessSyntax) + || HasTaskCompletedCore(context, taskReceiver, accessSyntax); + + private static bool HasTaskCompletedInContinuation( + SyntaxNodeAnalysisContext context, + ExpressionSyntax taskReceiver, + SyntaxNode accessSyntax) + { + foreach (AnonymousFunctionExpressionSyntax anonymousFunction in accessSyntax.Ancestors().OfType()) { - // `byRef` includes `out` parameters because they are the same as `ref` except don't require initialization first. - if (byRef && !arg.RefKindKeyword.IsKind(SyntaxKind.RefKeyword) && !arg.RefKindKeyword.IsKind(SyntaxKind.OutKeyword)) + ExpressionSyntax callbackExpression = anonymousFunction; + while (callbackExpression.Parent is ParenthesizedExpressionSyntax or CastExpressionSyntax) + { + callbackExpression = (ExpressionSyntax)callbackExpression.Parent; + } + + if (callbackExpression.Parent is not ArgumentSyntax anonymousFunctionArgument + || anonymousFunctionArgument.Parent?.Parent is not InvocationExpressionSyntax continuationInvocation + || context.SemanticModel.GetOperation(continuationInvocation, context.CancellationToken) is not IInvocationOperation continuationOperation + || !continuationOperation.Arguments.Any(argument => argument.Parameter?.Ordinal == 0 + && argument.Syntax.Span.Contains(anonymousFunction.Span))) { continue; } - IdentifierNameSyntax identiferName = arg.ChildNodes().OfType().FirstOrDefault(); - if (identiferName is null) + if (continuationOperation.TargetMethod.Name != nameof(Task.ContinueWith) + || !Utils.IsTask(continuationOperation.TargetMethod.ContainingType)) { - return false; + continue; + } + + ParameterSyntax? firstParameter = anonymousFunction switch + { + SimpleLambdaExpressionSyntax lambda => lambda.Parameter, + ParenthesizedLambdaExpressionSyntax lambda => lambda.ParameterList.Parameters.FirstOrDefault(), + AnonymousMethodExpressionSyntax anonymousMethod => anonymousMethod.ParameterList?.Parameters.FirstOrDefault(), + _ => null, + }; + if (firstParameter is null + || context.SemanticModel.GetDeclaredSymbol(firstParameter, context.CancellationToken) is not IParameterSymbol completedTask) + { + continue; } - if (identiferName.Identifier.ValueText == variableName) + (ImmutableHashSet taskSymbols, ImmutableHashSet potentialTaskSymbols) = + GetSymbolAndRefAliases(context, accessSyntax, completedTask); + if (accessSyntax.Ancestors().TakeWhile(node => node != anonymousFunction) + .Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax)) + { + (ImmutableHashSet outerDefiniteAliases, ImmutableHashSet outerPotentialAliases) = GetSymbolAndRefAliases( + context, + accessSyntax, + completedTask, + anonymousFunction, + includeAllCandidates: true); + taskSymbols = taskSymbols.Union(outerDefiniteAliases); + potentialTaskSymbols = potentialTaskSymbols.Union(outerPotentialAliases); + } + + ISymbol? receiverSymbol = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(taskReceiver), context.CancellationToken).Symbol; + if (receiverSymbol is object + && (SymbolEqualityComparer.Default.Equals(receiverSymbol, completedTask) || taskSymbols.Contains(receiverSymbol)) + && !IsTaskReassignedInContinuation(context, anonymousFunction, accessSyntax, potentialTaskSymbols)) { return true; } @@ -160,132 +474,1032 @@ private static bool IsVariablePassedToInvocation(InvocationExpressionSyntax invo return false; } - private static bool IsTaskCompletedWithWhenAll(SyntaxNodeAnalysisContext context, InvocationExpressionSyntax invocationExpr, string taskVariableName) + private static bool IsTaskReassignedInContinuation( + SyntaxNodeAnalysisContext context, + AnonymousFunctionExpressionSyntax continuation, + SyntaxNode accessSyntax, + IImmutableSet taskSymbols) { - // We only care about awaited invocations, because an un-awaited Task.WhenAll will be an error. - if (invocationExpr.Parent is not AwaitExpressionSyntax) + bool accessIsNested = accessSyntax.Ancestors().TakeWhile(node => node != continuation) + .Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax); + int beforePosition = accessIsNested ? continuation.Span.End + 1 : accessSyntax.SpanStart; + + foreach (AssignmentExpressionSyntax assignment in continuation.DescendantNodes().OfType()) { - return false; + SyntaxNode? nestedFunction = assignment.Ancestors().TakeWhile(node => node != continuation) + .FirstOrDefault(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax); + bool isDeferredWrite = accessIsNested + || nestedFunction is LocalFunctionStatementSyntax + || nestedFunction?.SpanStart < accessSyntax.SpanStart; + if ((assignment.SpanStart < beforePosition || isDeferredWrite) + && IsAssignmentToTask(context, assignment.Left, taskSymbols)) + { + return true; + } } - IEnumerable? memberAccessList = invocationExpr.ChildNodes().OfType(); - if (memberAccessList.Count() != 1) + foreach (ArgumentSyntax argument in continuation.DescendantNodes().OfType()) { - return false; + SyntaxNode? nestedFunction = argument.Ancestors().TakeWhile(node => node != continuation) + .FirstOrDefault(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax); + bool isDeferredWrite = accessIsNested + || nestedFunction is LocalFunctionStatementSyntax + || nestedFunction?.SpanStart < accessSyntax.SpanStart; + if ((argument.SpanStart < beforePosition || isDeferredWrite) + && (argument.RefKindKeyword.IsKind(SyntaxKind.RefKeyword) || argument.RefKindKeyword.IsKind(SyntaxKind.OutKeyword)) + && MayAliasTaskStorage(context, argument.Expression, taskSymbols)) + { + return true; + } + } + + return false; + } + + private static ExpressionSyntax UnwrapParentheses(ExpressionSyntax expression) + { + while (expression is ParenthesizedExpressionSyntax parenthesized) + { + expression = parenthesized.Expression; } - MemberAccessExpressionSyntax? memberAccess = memberAccessList.First(); + return expression; + } + + private static bool IsBuiltInBlockingMember( + SyntaxNodeAnalysisContext context, + ISymbol accessedSymbol, + CommonInterest.QualifiedMember expectedMember) + { + if (accessedSymbol is not IMethodSymbol { ReducedFrom: not null } reducedMethod) + { + return true; + } - // Does the invocation have the expected `Task.WhenAll` syntax? This is cheaper to verify before looking up its semantic type. - bool correctSyntax = memberAccess.Expression is IdentifierNameSyntax { Identifier.ValueText: Types.Task.TypeName } - && memberAccess.Name is IdentifierNameSyntax { Identifier.ValueText: Types.Task.WhenAll }; + if (expectedMember.IsMatch(reducedMethod.ReducedFrom)) + { + return true; + } - if (!correctSyntax) + if (expectedMember.Name != nameof(Task.Wait) + || context.Compilation.GetTypeByMetadataName(Types.Task.FullName) is not INamedTypeSymbol taskType) { return false; } - // Is this `Task.WhenAll` invocation from the System.Threading.Tasks.Task type? - ITypeSymbol? classType = context.SemanticModel.GetTypeInfo(memberAccess.Expression).Type; - var correctType = classType?.Name == Types.Task.TypeName && classType.BelongsToNamespace(Types.Task.Namespace); - if (!correctType) + return reducedMethod.Parameters.IsEmpty + && reducedMethod.ReturnsVoid + && Utils.IsEqualToOrDerivedFrom(reducedMethod.ReceiverType, taskType); + } + + private static ExpressionSyntax GetTaskReceiver(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) + { + ExpressionSyntax receiver = UnwrapParentheses(memberAccessSyntax.Expression); + if (receiver is InvocationExpressionSyntax getAwaiterInvocation + && getAwaiterInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: "GetAwaiter" } getAwaiterAccess + && IsSupportedGetAwaiterInvocation(context, getAwaiterInvocation, getAwaiterAccess.Expression)) + { + receiver = UnwrapParentheses(getAwaiterAccess.Expression); + } + + if (receiver is InvocationExpressionSyntax configureAwaitInvocation + && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess + && IsSupportedConfigureAwaitInvocation(context, configureAwaitInvocation)) + { + receiver = UnwrapParentheses(configureAwaitAccess.Expression); + } + + return receiver; + } + + private static bool IsSupportedGetAwaiterInvocation( + SyntaxNodeAnalysisContext context, + InvocationExpressionSyntax invocation, + ExpressionSyntax receiver) + { + if (invocation.ArgumentList.Arguments.Count != 0 + || context.SemanticModel.GetSymbolInfo(invocation, context.CancellationToken).Symbol is not IMethodSymbol method + || method.ReducedFrom is object + || method.IsStatic + || !method.Parameters.IsEmpty) { return false; } - // Is the task variable passed as an argument to `Task.WhenAll`? - return IsVariablePassedToInvocation(invocationExpr, taskVariableName, byRef: false); + if (IsTaskLike(method.ContainingType)) + { + return true; + } + + receiver = UnwrapParentheses(receiver); + return receiver is InvocationExpressionSyntax configureAwaitInvocation + && IsSupportedConfigureAwaitInvocation(context, configureAwaitInvocation); } - private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) + private static bool IsSupportedConfigureAwaitInvocation(SyntaxNodeAnalysisContext context, InvocationExpressionSyntax invocation) + => invocation.ArgumentList.Arguments.Count == 1 + && context.SemanticModel.GetSymbolInfo(invocation, context.CancellationToken).Symbol is IMethodSymbol method + && method.ReducedFrom is null + && !method.IsStatic + && method.Parameters.Length == 1 + && IsTaskLike(method.ContainingType); + + private static bool IsTaskLike(ITypeSymbol? type) + => Utils.IsTask(type) + || (type?.Name == nameof(ValueTask) && type.BelongsToNamespace(Namespaces.SystemThreadingTasks)); + + private static bool HasTaskCompletedCore( + SyntaxNodeAnalysisContext context, + ExpressionSyntax taskReceiver, + SyntaxNode accessSyntax) { - SyntaxNode? enclosingBlock = GetEnclosingBlock(memberAccessSyntax); - if (enclosingBlock is null) + taskReceiver = UnwrapParentheses(taskReceiver); + ITypeSymbol? taskType = context.SemanticModel.GetTypeInfo(taskReceiver, context.CancellationToken).Type; + if (!IsTaskLike(taskType)) + { + return false; + } + + ISymbol? taskSymbol = context.SemanticModel.GetSymbolInfo(taskReceiver, context.CancellationToken).Symbol; + if (taskSymbol is IParameterSymbol { RefKind: not RefKind.None } + || taskSymbol is not ILocalSymbol and not IParameterSymbol) + { + return false; + } + + if (context.SemanticModel.GetEnclosingSymbol(accessSyntax.SpanStart, context.CancellationToken) is not IMethodSymbol enclosingMethod + || !SymbolEqualityComparer.Default.Equals(taskSymbol.ContainingSymbol, enclosingMethod)) + { + return false; + } + + (ImmutableHashSet taskSymbols, ImmutableHashSet potentialTaskSymbols) = + GetSymbolAndRefAliases(context, accessSyntax, taskSymbol); + if (taskSymbols.Any(symbol => symbol is IParameterSymbol { RefKind: not RefKind.None })) + { + return false; + } + + if (NestedFunctionMayReassignTask(context, accessSyntax, potentialTaskSymbols)) + { + return false; + } + + if (ContainsPotentialControlFlowBypass(accessSyntax)) + { + return false; + } + + if (IsWithinCompletedTaskBranch(context, accessSyntax, taskSymbols, potentialTaskSymbols)) + { + return Utils.IsTask(taskType) || !MayHaveConsumedValueTaskBefore(context, accessSyntax, taskSymbols); + } + + // Awaiting an IValueTaskSource-backed ValueTask consumes it, so a later Result access is not safe. + if (!Utils.IsTask(taskType)) { return false; } - // Get the task variable name from the problematic member access expression so that we can later try - // and determine if it has been used in a `Task.WhenAll` invocation. - // Examples: - // task1.Result; - // task2.GetAwaiter().GetResult(); - string? taskVariableName = null; - ExpressionSyntax parentExpr = memberAccessSyntax.Expression; - while (parentExpr is not null) + StatementSyntax? containingStatement = accessSyntax.FirstAncestorOrSelf(); + if (containingStatement is null) + { + ArrowExpressionClauseSyntax? arrowExpression = accessSyntax.FirstAncestorOrSelf(); + return arrowExpression is object + && TryGetAwaitExpression(context, arrowExpression, taskSymbols, accessSyntax.SpanStart, out AwaitExpressionSyntax? arrowPrecedingAwait) + && !MayReassignTask(context, arrowExpression, potentialTaskSymbols, arrowPrecedingAwait.Span.End, accessSyntax.SpanStart); + } + + if (TryGetAwaitExpression(context, containingStatement, taskSymbols, accessSyntax.SpanStart, out AwaitExpressionSyntax? precedingAwait) + && !MayReassignTask(context, containingStatement, potentialTaskSymbols, precedingAwait.Span.End, accessSyntax.SpanStart)) + { + return true; + } + + while (true) { - if (parentExpr is IdentifierNameSyntax identifierExpr) + if (MayReassignTask(context, containingStatement, potentialTaskSymbols, containingStatement.SpanStart - 1, accessSyntax.SpanStart)) { - taskVariableName = identifierExpr.Identifier.ValueText; - break; + return false; } - else if (parentExpr is MemberAccessExpressionSyntax memberAccessExpr) + + SyntaxList statements = containingStatement.Parent switch { - parentExpr = memberAccessExpr.Expression; + BlockSyntax block => block.Statements, + SwitchSectionSyntax switchSection => switchSection.Statements, + _ => default, + }; + int statementIndex = statements.IndexOf(containingStatement); + for (int i = statementIndex - 1; i >= 0; i--) + { + StatementSyntax statement = statements[i]; + if (StatementCompletesTask(context, statement, taskSymbols, potentialTaskSymbols)) + { + return true; + } + + if (MayReassignTask(context, statement, potentialTaskSymbols)) + { + return false; + } } - else if (parentExpr is InvocationExpressionSyntax invocExpr) + + StatementSyntax? outerStatement = containingStatement.Ancestors().OfType() + .FirstOrDefault(statement => statement.Parent is BlockSyntax or SwitchSectionSyntax); + if (outerStatement is null) { - parentExpr = invocExpr.Expression; + return false; } - else + + if (outerStatement is WhileStatementSyntax + or DoStatementSyntax + or ForStatementSyntax + or ForEachStatementSyntax + or ForEachVariableStatementSyntax + && MayReassignTask(context, outerStatement, potentialTaskSymbols)) { - break; + return false; } - } - if (taskVariableName is null) - { - return false; + if (MayReassignTask(context, outerStatement, potentialTaskSymbols, outerStatement.SpanStart - 1, containingStatement.SpanStart)) + { + return false; + } + + containingStatement = outerStatement; } + } - // Find all `Task.WhenAll` invocations that precede the problematic member access, which are also in the same enclosing block. - IEnumerable? taskWhenAllInvocationList = - from invoc in enclosingBlock.DescendantNodes().OfType() - where memberAccessSyntax.SpanStart > invoc.Span.End && - IsTaskCompletedWithWhenAll(context, invoc, taskVariableName) - select invoc; + private static bool MayHaveConsumedValueTaskBefore( + SyntaxNodeAnalysisContext context, + SyntaxNode accessSyntax, + IImmutableSet taskSymbols) + { + SyntaxNode searchRoot = accessSyntax.AncestorsAndSelf().FirstOrDefault( + ancestor => ancestor is AnonymousFunctionExpressionSyntax + or LocalFunctionStatementSyntax + or BaseMethodDeclarationSyntax + or AccessorDeclarationSyntax) + ?? accessSyntax.FirstAncestorOrSelf()?.Parent + ?? accessSyntax; + bool DescendIntoChildren(SyntaxNode child) => + child == searchRoot || child is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; - if (!taskWhenAllInvocationList.Any()) + bool IsDefinitelyExecutedBefore(SyntaxNode candidate, SyntaxNode consumption) { - return false; + StatementSyntax? consumptionStatement = consumption.FirstAncestorOrSelf(); + if (consumptionStatement?.Parent is not BlockSyntax block) + { + return false; + } + + StatementSyntax? candidateStatement = candidate.FirstAncestorOrSelf(); + while (candidateStatement is object && candidateStatement.Parent != block) + { + if (candidateStatement.Parent is not BlockSyntax containingBlock) + { + return false; + } + + candidateStatement = containingBlock; + } + + return candidateStatement is object + && block.Statements.IndexOf(candidateStatement) < block.Statements.IndexOf(consumptionStatement); } - // If a `Task.WhenAll` invocation precedes the problematic member access, and the task variable has not been - // invalidated in between, then we consider the task to be completed. - // Example: - // await Task.WhenAll(task1, task2, task3); - // task1 = Task.Run(...); // Invalidates `task1` - // DoSomething(ref task2); // Invalidates `task2` - // task1.Result; // Warn - // task2.Result; // Warn - // task3.Result; // No warning, task3 has not been invalidated in between WhenAll and this problematic member access - foreach (InvocationExpressionSyntax? taskWhenAllInvocation in taskWhenAllInvocationList) + ImmutableHashSet GetTaskAndCopySymbolsBefore(SyntaxNode consumption) { - // Has the task variable been assigned to a new task? - IEnumerable? assignmentList = - from assign in enclosingBlock.DescendantNodes().OfType() - where assign.SpanStart > taskWhenAllInvocation.Span.End && - assign.SpanStart < memberAccessSyntax.SpanStart && - ((IdentifierNameSyntax)assign.Left).Identifier.ValueText == taskVariableName - select assign; + var taskAndCopySymbols = new HashSet(taskSymbols, SymbolEqualityComparer.Default); + bool IsTaskOrCopy(ExpressionSyntax expression) + { + expression = UnwrapParentheses(expression); + ISymbol? expressionSymbol = context.SemanticModel.GetSymbolInfo(expression, context.CancellationToken).Symbol; + if (expressionSymbol is object && taskAndCopySymbols.Contains(expressionSymbol)) + { + return true; + } + + return expression is InvocationExpressionSyntax configureAwaitInvocation + && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess + && IsTaskOrCopy(configureAwaitAccess.Expression); + } - if (assignmentList.Any()) + IEnumerable copyOperations = searchRoot.DescendantNodes(DescendIntoChildren) + .Where(node => node.SpanStart < consumption.SpanStart + && node is VariableDeclaratorSyntax or AssignmentExpressionSyntax) + .OrderBy(node => node.SpanStart); + foreach (SyntaxNode copyOperation in copyOperations) { - return false; + if (copyOperation is VariableDeclaratorSyntax { Initializer: { } initializer } variable + && context.SemanticModel.GetDeclaredSymbol(variable, context.CancellationToken) is ILocalSymbol declaredLocal + && IsTaskOrCopy(initializer.Value)) + { + taskAndCopySymbols.Add(declaredLocal); + } + else if (copyOperation is AssignmentExpressionSyntax assignment + && context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol is ISymbol assignedSymbol + && assignedSymbol is ILocalSymbol or IParameterSymbol) + { + if (IsTaskOrCopy(assignment.Right)) + { + taskAndCopySymbols.Add(assignedSymbol); + } + else if (IsDefinitelyExecutedBefore(assignment, consumption)) + { + taskAndCopySymbols.Remove(assignedSymbol); + } + } } - // Has the task variable been passed by ref to a method? - // If so, we must assume the worst case that the method has assigned it to a new task. - IEnumerable? invocationList = - from invoc in enclosingBlock.DescendantNodes().OfType() - where invoc.SpanStart > taskWhenAllInvocation.Span.End && - invoc.SpanStart < memberAccessSyntax.SpanStart && - IsVariablePassedToInvocation(invoc, taskVariableName, byRef: true) - select invoc; + return taskAndCopySymbols.ToImmutableHashSet(SymbolEqualityComparer.Default); + } - return !invocationList.Any(); + foreach (AwaitExpressionSyntax awaitExpression in searchRoot.DescendantNodes(DescendIntoChildren) + .OfType() + .Where(awaitExpression => awaitExpression.SpanStart < accessSyntax.SpanStart) + .OrderBy(awaitExpression => awaitExpression.SpanStart)) + { + ImmutableHashSet taskAndCopySymbols = GetTaskAndCopySymbolsBefore(awaitExpression); + if (AwaitCompletesTask(context, awaitExpression, taskAndCopySymbols) + || AwaitMayConsumeValueTask(context, awaitExpression, taskAndCopySymbols)) + { + return true; + } } - return false; + foreach (MemberAccessExpressionSyntax blockingAccess in searchRoot.DescendantNodes(DescendIntoChildren) + .OfType() + .Where(memberAccess => memberAccess.SpanStart < accessSyntax.SpanStart + && IsValueTaskConsumption(memberAccess))) + { + ImmutableHashSet taskAndCopySymbols = GetTaskAndCopySymbolsBefore(blockingAccess); + if (IsOneOfSymbols(context, GetTaskReceiver(context, blockingAccess), taskAndCopySymbols)) + { + return true; + } + } + + foreach (LocalFunctionStatementSyntax localFunction in searchRoot.DescendantNodes() + .OfType() + .Where(localFunction => localFunction.SpanStart < accessSyntax.SpanStart)) + { + if (context.SemanticModel.GetDeclaredSymbol(localFunction, context.CancellationToken) is not IMethodSymbol localFunctionSymbol) + { + continue; + } + + if (searchRoot.DescendantNodes(DescendIntoChildren) + .OfType() + .Any(invocation => invocation.SpanStart < accessSyntax.SpanStart + && SymbolEqualityComparer.Default.Equals( + context.SemanticModel.GetSymbolInfo(invocation, context.CancellationToken).Symbol?.OriginalDefinition, + localFunctionSymbol.OriginalDefinition) + && NestedFunctionMayConsumeValueTask(localFunction, invocation))) + { + return true; + } + } + + foreach (VariableDeclaratorSyntax delegateVariable in searchRoot.DescendantNodes(DescendIntoChildren) + .OfType() + .Where(variable => variable.SpanStart < accessSyntax.SpanStart + && variable.Initializer?.Value is AnonymousFunctionExpressionSyntax)) + { + var anonymousFunction = (AnonymousFunctionExpressionSyntax)delegateVariable.Initializer!.Value; + if (context.SemanticModel.GetDeclaredSymbol(delegateVariable, context.CancellationToken) is not ILocalSymbol delegateSymbol) + { + continue; + } + + if (searchRoot.DescendantNodes(DescendIntoChildren) + .OfType() + .Any(invocation => invocation.SpanStart < accessSyntax.SpanStart + && SymbolEqualityComparer.Default.Equals( + context.SemanticModel.GetSymbolInfo(invocation.Expression, context.CancellationToken).Symbol, + delegateSymbol) + && NestedFunctionMayConsumeValueTask(anonymousFunction, invocation))) + { + return true; + } + } + + return false; + + bool IsValueTaskConsumption(MemberAccessExpressionSyntax memberAccess) + { + IOperation? operation = context.SemanticModel.GetOperation(memberAccess, context.CancellationToken); + for (IOperation? ancestor = operation; ancestor is object; ancestor = ancestor.Parent) + { + if (ancestor is INameOfOperation) + { + return false; + } + } + + return memberAccess.Name.Identifier.ValueText switch + { + nameof(Task.Result) => operation is IPropertyReferenceOperation, + "GetResult" => memberAccess.Parent is InvocationExpressionSyntax invocation + && invocation.Expression == memberAccess + && context.SemanticModel.GetOperation(invocation, context.CancellationToken) is IInvocationOperation, + _ => false, + }; + } + + bool NestedFunctionMayConsumeValueTask(SyntaxNode function, InvocationExpressionSyntax invocation) + { + bool DescendIntoFunction(SyntaxNode child) => + child == function || child is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; + var nestedTaskSymbols = new HashSet(GetTaskAndCopySymbolsBefore(invocation), SymbolEqualityComparer.Default); + foreach (SyntaxNode copyOperation in function.DescendantNodes(DescendIntoFunction) + .Where(node => node is VariableDeclaratorSyntax or AssignmentExpressionSyntax) + .OrderBy(node => node.SpanStart)) + { + ExpressionSyntax? source = copyOperation switch + { + VariableDeclaratorSyntax { Initializer.Value: { } initializer } => initializer, + AssignmentExpressionSyntax assignment => assignment.Right, + _ => null, + }; + ISymbol? target = copyOperation switch + { + VariableDeclaratorSyntax variable => context.SemanticModel.GetDeclaredSymbol(variable, context.CancellationToken), + AssignmentExpressionSyntax assignment => context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol, + _ => null, + }; + if (source is object + && target is ILocalSymbol or IParameterSymbol + && IsOneOfSymbols(context, source, nestedTaskSymbols.ToImmutableHashSet(SymbolEqualityComparer.Default))) + { + nestedTaskSymbols.Add(target); + } + } + + ImmutableHashSet symbols = nestedTaskSymbols.ToImmutableHashSet(SymbolEqualityComparer.Default); + return function.DescendantNodes(DescendIntoFunction).OfType() + .Any(awaitExpression => AwaitMayConsumeValueTask(context, awaitExpression, symbols)) + || function.DescendantNodes(DescendIntoFunction).OfType() + .Any(memberAccess => IsValueTaskConsumption(memberAccess) + && IsOneOfSymbols(context, GetTaskReceiver(context, memberAccess), symbols)); + } + } + + private static bool AwaitMayConsumeValueTask( + SyntaxNodeAnalysisContext context, + AwaitExpressionSyntax awaitExpression, + IImmutableSet taskSymbols) + { + ExpressionSyntax awaitedExpression = UnwrapParentheses(awaitExpression.Expression); + if (IsOneOfSymbols(context, awaitedExpression, taskSymbols)) + { + return true; + } + + return awaitedExpression is InvocationExpressionSyntax invocation + && ((invocation.Expression is MemberAccessExpressionSyntax memberAccess + && IsOneOfSymbols(context, memberAccess.Expression, taskSymbols)) + || (context.SemanticModel.GetOperation(invocation, context.CancellationToken) is IInvocationOperation invocationOperation + && invocationOperation.Arguments.Any(argument => argument.Syntax is ArgumentSyntax argumentSyntax + && IsOneOfSymbols(context, argumentSyntax.Expression, taskSymbols)))); + } + + private static bool ContainsPotentialControlFlowBypass(SyntaxNode accessSyntax) + { + SyntaxNode searchRoot = accessSyntax.AncestorsAndSelf().FirstOrDefault( + ancestor => ancestor is AnonymousFunctionExpressionSyntax + or LocalFunctionStatementSyntax + or BaseMethodDeclarationSyntax + or AccessorDeclarationSyntax) + ?? accessSyntax.FirstAncestorOrSelf()?.Parent + ?? accessSyntax; + bool DescendIntoChildren(SyntaxNode child) => + child == searchRoot || child is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; + + return searchRoot.DescendantNodes(DescendIntoChildren) + .Any(node => node.SpanStart < accessSyntax.SpanStart + && node is GotoStatementSyntax or LabeledStatementSyntax); + } + + private static bool IsWithinCompletedTaskBranch( + SyntaxNodeAnalysisContext context, + SyntaxNode accessSyntax, + IImmutableSet taskSymbols, + IImmutableSet potentialTaskSymbols) + { + foreach (IfStatementSyntax ifStatement in accessSyntax.Ancestors().OfType()) + { + IEnumerable nodesBetweenAccessAndCondition = accessSyntax.Ancestors().TakeWhile(node => node != ifStatement); + if (nodesBetweenAccessAndCondition.Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax)) + { + continue; + } + + if (nodesBetweenAccessAndCondition + .Where(node => node is WhileStatementSyntax + or DoStatementSyntax + or ForStatementSyntax + or ForEachStatementSyntax + or ForEachVariableStatementSyntax) + .Any(loop => MayReassignTask(context, loop, potentialTaskSymbols))) + { + continue; + } + + foreach (BinaryExpressionSyntax binary in accessSyntax.Ancestors() + .TakeWhile(node => node != ifStatement) + .OfType()) + { + bool leftProvesCompletion = binary.Right.FullSpan.Contains(accessSyntax.Span) + && ((binary.IsKind(SyntaxKind.LogicalAndExpression) + && ConditionProvesCompletion(context, binary.Left, taskSymbols, conditionValue: true)) + || (binary.IsKind(SyntaxKind.LogicalOrExpression) + && ConditionProvesCompletion(context, binary.Left, taskSymbols, conditionValue: false))); + if (leftProvesCompletion + && !MayReassignTask(context, binary, potentialTaskSymbols, binary.Left.SpanStart - 1, accessSyntax.SpanStart)) + { + return true; + } + } + + if (ifStatement.Statement.FullSpan.Contains(accessSyntax.Span) + && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: true) + && !MayReassignTask(context, ifStatement, potentialTaskSymbols, ifStatement.Condition.SpanStart - 1, accessSyntax.SpanStart)) + { + return true; + } + + if (ifStatement.Else?.Statement.FullSpan.Contains(accessSyntax.Span) is true + && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: false) + && !MayReassignTask(context, ifStatement, potentialTaskSymbols, ifStatement.Condition.SpanStart - 1, accessSyntax.SpanStart)) + { + return true; + } + } + + return false; + } + + private static bool ConditionProvesCompletion( + SyntaxNodeAnalysisContext context, + ExpressionSyntax condition, + IImmutableSet taskSymbols, + bool conditionValue) + { + condition = UnwrapParentheses(condition); + if (condition is PrefixUnaryExpressionSyntax { RawKind: (int)SyntaxKind.LogicalNotExpression } logicalNot) + { + return ConditionProvesCompletion(context, logicalNot.Operand, taskSymbols, !conditionValue); + } + + if (condition is BinaryExpressionSyntax binary) + { + if (conditionValue && binary.IsKind(SyntaxKind.LogicalAndExpression)) + { + return ConditionProvesCompletion(context, binary.Left, taskSymbols, conditionValue: true) + || ConditionProvesCompletion(context, binary.Right, taskSymbols, conditionValue: true); + } + + if (!conditionValue && binary.IsKind(SyntaxKind.LogicalOrExpression)) + { + return ConditionProvesCompletion(context, binary.Left, taskSymbols, conditionValue: false) + || ConditionProvesCompletion(context, binary.Right, taskSymbols, conditionValue: false); + } + + if (conditionValue && binary.IsKind(SyntaxKind.LogicalOrExpression)) + { + return ConditionProvesCompletion(context, binary.Left, taskSymbols, conditionValue: true) + && ConditionProvesCompletion(context, binary.Right, taskSymbols, conditionValue: true); + } + + if (!conditionValue && binary.IsKind(SyntaxKind.LogicalAndExpression)) + { + return ConditionProvesCompletion(context, binary.Left, taskSymbols, conditionValue: false) + && ConditionProvesCompletion(context, binary.Right, taskSymbols, conditionValue: false); + } + + bool equalityHolds = binary.IsKind(SyntaxKind.EqualsExpression) ? conditionValue + : binary.IsKind(SyntaxKind.NotEqualsExpression) ? !conditionValue + : false; + if ((binary.IsKind(SyntaxKind.EqualsExpression) || binary.IsKind(SyntaxKind.NotEqualsExpression)) + && IsRanToCompletionComparison(context, binary.Left, binary.Right, taskSymbols)) + { + return equalityHolds; + } + } + + if (condition is MemberAccessExpressionSyntax completedProperty + && IsOneOfSymbols(context, completedProperty.Expression, taskSymbols) + && completedProperty.Name.Identifier.ValueText is nameof(Task.IsCompleted) + or nameof(Task.IsCanceled) + or nameof(Task.IsFaulted) + or "IsCompletedSuccessfully") + { + return conditionValue; + } + + return false; + } + + private static bool IsRanToCompletionComparison( + SyntaxNodeAnalysisContext context, + ExpressionSyntax left, + ExpressionSyntax right, + IImmutableSet taskSymbols) + { + return (IsTaskStatus(context, left, taskSymbols) && IsRanToCompletion(context, right)) + || (IsTaskStatus(context, right, taskSymbols) && IsRanToCompletion(context, left)); + } + + private static bool IsTaskStatus(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, IImmutableSet taskSymbols) + => expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.Status) } statusAccess + && IsOneOfSymbols(context, statusAccess.Expression, taskSymbols); + + private static bool IsRanToCompletion(SyntaxNodeAnalysisContext context, ExpressionSyntax expression) + => context.SemanticModel.GetSymbolInfo(expression, context.CancellationToken).Symbol is IFieldSymbol status + && status.Name == nameof(TaskStatus.RanToCompletion) + && status.ContainingType.Name == nameof(TaskStatus) + && status.ContainingType.BelongsToNamespace(Namespaces.SystemThreadingTasks); + + private static bool StatementCompletesTask( + SyntaxNodeAnalysisContext context, + StatementSyntax statement, + IImmutableSet taskSymbols, + IImmutableSet potentialTaskSymbols) + { + if (TryGetAwaitExpression(context, statement, taskSymbols, out AwaitExpressionSyntax? awaitExpression) + && !MayReassignTask(context, statement, potentialTaskSymbols, awaitExpression.Span.End, statement.Span.End + 1)) + { + return true; + } + + if (statement is BlockSyntax block) + { + return StatementDefinitelyAwaitsTask(context, block, taskSymbols, potentialTaskSymbols); + } + + if (statement is not IfStatementSyntax ifStatement) + { + return false; + } + + if (MayReassignTask(context, ifStatement.Condition, potentialTaskSymbols)) + { + return false; + } + + if (ifStatement.Else is null) + { + return ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: false) + && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbols, potentialTaskSymbols); + } + + return (ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: true) + && !MayReassignTask(context, ifStatement.Statement, potentialTaskSymbols) + && StatementDefinitelyAwaitsTask(context, ifStatement.Else.Statement, taskSymbols, potentialTaskSymbols)) + || (ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: false) + && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbols, potentialTaskSymbols) + && !MayReassignTask(context, ifStatement.Else.Statement, potentialTaskSymbols)); + } + + private static bool StatementDefinitelyAwaitsTask( + SyntaxNodeAnalysisContext context, + StatementSyntax statement, + IImmutableSet taskSymbols, + IImmutableSet potentialTaskSymbols) + { + if (TryGetAwaitExpression(context, statement, taskSymbols, out AwaitExpressionSyntax? awaitExpression)) + { + return !MayReassignTask(context, statement, potentialTaskSymbols, awaitExpression.Span.End, statement.Span.End + 1); + } + + if (statement is IfStatementSyntax ifStatement) + { + if (StatementCompletesTask(context, ifStatement, taskSymbols, potentialTaskSymbols)) + { + return true; + } + + return ifStatement.Else is { } elseClause + && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbols, potentialTaskSymbols) + && StatementDefinitelyAwaitsTask(context, elseClause.Statement, taskSymbols, potentialTaskSymbols); + } + + if (statement is BlockSyntax block) + { + for (int i = block.Statements.Count - 1; i >= 0; i--) + { + if (StatementDefinitelyAwaitsTask(context, block.Statements[i], taskSymbols, potentialTaskSymbols)) + { + return true; + } + + if (MayReassignTask(context, block.Statements[i], potentialTaskSymbols)) + { + return false; + } + } + } + + return false; + } + + private static bool TryGetAwaitExpression( + SyntaxNodeAnalysisContext context, + StatementSyntax statement, + IImmutableSet taskSymbols, + [NotNullWhen(true)] out AwaitExpressionSyntax? awaitExpression) + => TryGetAwaitExpression(context, statement, taskSymbols, statement.Span.End + 1, out awaitExpression); + + private static bool TryGetAwaitExpression( + SyntaxNodeAnalysisContext context, + SyntaxNode node, + IImmutableSet taskSymbols, + int beforePosition, + [NotNullWhen(true)] out AwaitExpressionSyntax? awaitExpression) + { + static bool DescendIntoChildren(SyntaxNode node) => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; + + if (node is WhileStatementSyntax or DoStatementSyntax or ForStatementSyntax or ForEachStatementSyntax or ForEachVariableStatementSyntax) + { + awaitExpression = null; + return false; + } + + foreach (AwaitExpressionSyntax candidate in node.DescendantNodes(DescendIntoChildren).OfType().Reverse()) + { + if (candidate.Span.End >= beforePosition) + { + continue; + } + + IEnumerable ancestorsWithinStatement = candidate.Ancestors().TakeWhile(ancestor => ancestor != node); + bool isConditionallyExecuted = ancestorsWithinStatement.Any( + ancestor => ancestor is StatementSyntax + or WhenClauseSyntax + or CatchFilterClauseSyntax + || (ancestor is SwitchExpressionSyntax switchExpression + && !switchExpression.GoverningExpression.FullSpan.Contains(candidate.Span)) + || (ancestor is ConditionalAccessExpressionSyntax conditionalAccess + && !conditionalAccess.Expression.FullSpan.Contains(candidate.Span)) + || (ancestor is ConditionalExpressionSyntax conditional + && !conditional.Condition.FullSpan.Contains(candidate.Span)) + || (ancestor is BinaryExpressionSyntax binary + && (binary.IsKind(SyntaxKind.LogicalAndExpression) + || binary.IsKind(SyntaxKind.LogicalOrExpression) + || binary.IsKind(SyntaxKind.CoalesceExpression)) + && binary.Right.FullSpan.Contains(candidate.Span)) + || (ancestor is AssignmentExpressionSyntax assignment + && assignment.IsKind(SyntaxKind.CoalesceAssignmentExpression) + && assignment.Right.FullSpan.Contains(candidate.Span))); + if (!isConditionallyExecuted && AwaitCompletesTask(context, candidate, taskSymbols)) + { + awaitExpression = candidate; + return true; + } + } + + awaitExpression = null; + return false; + } + + private static bool AwaitCompletesTask( + SyntaxNodeAnalysisContext context, + AwaitExpressionSyntax awaitExpression, + IImmutableSet taskSymbols) + { + if (MayReassignTask(context, awaitExpression.Expression, taskSymbols)) + { + return false; + } + + ExpressionSyntax awaitedExpression = UnwrapParentheses(awaitExpression.Expression); + if (awaitedExpression is InvocationExpressionSyntax configureAwaitInvocation + && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess + && IsSupportedConfigureAwaitInvocation(context, configureAwaitInvocation)) + { + awaitedExpression = UnwrapParentheses(configureAwaitAccess.Expression); + } + + if (IsOneOfSymbols(context, awaitedExpression, taskSymbols)) + { + return true; + } + + if (awaitedExpression is InvocationExpressionSyntax whenAllInvocation + && context.SemanticModel.GetSymbolInfo(whenAllInvocation, context.CancellationToken).Symbol is IMethodSymbol whenAllMethod + && whenAllMethod.Name == nameof(Task.WhenAll) + && whenAllMethod.ContainingType.Name == nameof(Task) + && whenAllMethod.ContainingType.BelongsToNamespace(Namespaces.SystemThreadingTasks)) + { + return whenAllInvocation.ArgumentList.Arguments.Any( + argument => IsOneOfSymbols(context, argument.Expression, taskSymbols) + || (argument.Expression is ImplicitArrayCreationExpressionSyntax { Initializer: { } implicitArrayInitializer } + && implicitArrayInitializer.Expressions.Any(expression => IsOneOfSymbols(context, expression, taskSymbols))) + || (argument.Expression is ArrayCreationExpressionSyntax { Initializer: { } arrayInitializer } + && arrayInitializer.Expressions.Any(expression => IsOneOfSymbols(context, expression, taskSymbols))) + || LocalTaskCollectionContainsTrackedTask(context, argument.Expression, awaitExpression, taskSymbols)); + } + + return false; + } + + private static bool LocalTaskCollectionContainsTrackedTask( + SyntaxNodeAnalysisContext context, + ExpressionSyntax collectionExpression, + AwaitExpressionSyntax awaitExpression, + IImmutableSet taskSymbols) + { + if (context.SemanticModel.GetSymbolInfo(UnwrapParentheses(collectionExpression), context.CancellationToken).Symbol is not ILocalSymbol collection + || collection.DeclaringSyntaxReferences.SingleOrDefault()?.GetSyntax(context.CancellationToken) is not VariableDeclaratorSyntax { Initializer: { } initializer } + || initializer.SpanStart >= awaitExpression.SpanStart) + { + return false; + } + + bool InitializerContainsTrackedTask(ExpressionSyntax expression) + { + InitializerExpressionSyntax? arrayInitializer = expression switch + { + InitializerExpressionSyntax directInitializer => directInitializer, + ArrayCreationExpressionSyntax arrayCreation => arrayCreation.Initializer, + ImplicitArrayCreationExpressionSyntax implicitArrayCreation => implicitArrayCreation.Initializer, + _ => null, + }; + return arrayInitializer?.Expressions.Any(item => IsOneOfSymbols(context, item, taskSymbols)) is true; + } + + if (!InitializerContainsTrackedTask(initializer.Value)) + { + return false; + } + + SyntaxNode searchRoot = awaitExpression.AncestorsAndSelf().FirstOrDefault( + ancestor => ancestor is AnonymousFunctionExpressionSyntax + or LocalFunctionStatementSyntax + or BaseMethodDeclarationSyntax + or AccessorDeclarationSyntax) + ?? awaitExpression; + IImmutableSet collectionSymbol = ImmutableHashSet.Create(SymbolEqualityComparer.Default, collection); + if (MayReassignTask(context, searchRoot, taskSymbols, initializer.Span.End, awaitExpression.SpanStart) + || MayReassignTask(context, searchRoot, collectionSymbol, initializer.Span.End, awaitExpression.SpanStart) + || NestedFunctionMayReassignTask(context, awaitExpression, taskSymbols) + || NestedFunctionMayReassignTask(context, awaitExpression, collectionSymbol)) + { + return false; + } + + bool IsCollectionElement(ExpressionSyntax expression) + => UnwrapParentheses(expression) is ElementAccessExpressionSyntax elementAccess + && IsOneOfSymbols(context, elementAccess.Expression, collectionSymbol); + return !searchRoot.DescendantNodes() + .Where(node => node.SpanStart > initializer.Span.End && node.SpanStart < awaitExpression.SpanStart) + .Any(node => (node is AssignmentExpressionSyntax assignment && IsCollectionElement(assignment.Left)) + || (node is ArgumentSyntax argument + && (argument.RefKindKeyword.IsKind(SyntaxKind.RefKeyword) || argument.RefKindKeyword.IsKind(SyntaxKind.OutKeyword)) + && IsCollectionElement(argument.Expression))); + } + + private static bool MayReassignTask(SyntaxNodeAnalysisContext context, SyntaxNode node, IImmutableSet taskSymbols) + => MayReassignTask(context, node, taskSymbols, node.SpanStart - 1, node.Span.End + 1); + + private static bool MayReassignTask( + SyntaxNodeAnalysisContext context, + SyntaxNode node, + IImmutableSet taskSymbols, + int afterPosition, + int beforePosition) + { + bool DescendIntoChildren(SyntaxNode child) => + child == node || child is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; + + foreach (AssignmentExpressionSyntax assignment in node.DescendantNodesAndSelf(DescendIntoChildren).OfType()) + { + if (assignment.SpanStart > afterPosition + && assignment.SpanStart < beforePosition + && IsAssignmentToTask(context, assignment.Left, taskSymbols)) + { + return true; + } + } + + foreach (ArgumentSyntax argument in node.DescendantNodesAndSelf(DescendIntoChildren).OfType()) + { + if (argument.SpanStart > afterPosition + && argument.SpanStart < beforePosition + && context.SemanticModel.GetOperation(argument, context.CancellationToken) is IArgumentOperation { Parameter.RefKind: RefKind.Ref or RefKind.Out } + && MayAliasTaskStorage(context, argument.Expression, taskSymbols)) + { + return true; + } + } + + return false; + } + + private static bool NestedFunctionMayReassignTask( + SyntaxNodeAnalysisContext context, + SyntaxNode node, + IImmutableSet taskSymbols) + { + SyntaxNode? containingFunction = node.Ancestors().FirstOrDefault( + ancestor => ancestor is AnonymousFunctionExpressionSyntax + or LocalFunctionStatementSyntax + or BaseMethodDeclarationSyntax + or AccessorDeclarationSyntax) + ?? node.FirstAncestorOrSelf()?.Parent; + bool FunctionMayReassignTask(SyntaxNode nestedFunction, IImmutableSet containingTaskSymbols) + { + ImmutableHashSet.Builder nestedTaskSymbols = ImmutableHashSet.CreateBuilder(SymbolEqualityComparer.Default); + nestedTaskSymbols.UnionWith(containingTaskSymbols); + foreach (ISymbol taskSymbol in containingTaskSymbols) + { + nestedTaskSymbols.UnionWith(GetSymbolAndRefAliases( + context, + nestedFunction, + taskSymbol, + nestedFunction, + includeAllCandidates: true).Potential); + } + + IImmutableSet nestedSymbols = nestedTaskSymbols.ToImmutable(); + if (MayReassignTask(context, nestedFunction, nestedSymbols)) + { + return true; + } + + return nestedFunction.DescendantNodes( + descendant => descendant == nestedFunction + || descendant is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax) + .Where(descendant => descendant is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax) + .Any(descendant => FunctionMayReassignTask(descendant, nestedSymbols)); + } + + return containingFunction?.DescendantNodes() + .Where(descendant => descendant is LocalFunctionStatementSyntax + || (descendant is AnonymousFunctionExpressionSyntax && descendant.SpanStart < node.SpanStart)) + .Any(nestedFunction => FunctionMayReassignTask(nestedFunction, taskSymbols)) is true; + } + + private static bool IsAssignmentToTask(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, IImmutableSet taskSymbols) + { + expression = UnwrapParentheses(expression); + if (expression is TupleExpressionSyntax tuple) + { + return tuple.Arguments.Any(argument => IsAssignmentToTask(context, argument.Expression, taskSymbols)); + } + + return MayAliasTaskStorage(context, expression, taskSymbols); + } + + private static bool MayAliasTaskStorage( + SyntaxNodeAnalysisContext context, + ExpressionSyntax expression, + IImmutableSet taskSymbols) + { + expression = UnwrapParentheses(expression); + if (expression is RefExpressionSyntax refExpression) + { + return MayAliasTaskStorage(context, refExpression.Expression, taskSymbols); + } + + if (IsOneOfSymbols(context, expression, taskSymbols)) + { + return true; + } + + if (expression is ConditionalExpressionSyntax conditional) + { + return MayAliasTaskStorage(context, conditional.WhenTrue, taskSymbols) + || MayAliasTaskStorage(context, conditional.WhenFalse, taskSymbols); + } + + if (expression is InvocationExpressionSyntax invocation + && context.SemanticModel.GetOperation(invocation, context.CancellationToken) is IInvocationOperation invocationOperation + && (invocationOperation.TargetMethod.ReturnsByRef || invocationOperation.TargetMethod.ReturnsByRefReadonly)) + { + IMethodSymbol method = invocationOperation.TargetMethod; + if ((method.ReducedFrom ?? method).Parameters is [{ RefKind: not RefKind.None }, ..] + && invocation.Expression is MemberAccessExpressionSyntax memberAccess + && MayAliasTaskStorage(context, memberAccess.Expression, taskSymbols)) + { + return true; + } + + if (invocationOperation.Arguments.Any(argument => argument.Parameter?.RefKind != RefKind.None + && argument.Syntax is ArgumentSyntax argumentSyntax + && MayAliasTaskStorage(context, argumentSyntax.Expression, taskSymbols))) + { + return true; + } + } + + return false; + } + + private static bool IsOneOfSymbols(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, IImmutableSet symbols) + { + ISymbol? symbol = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(expression), context.CancellationToken).Symbol; + return symbol is object && symbols.Contains(symbol); } } diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpUtils.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpUtils.cs index 47292ea5b..94634928f 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpUtils.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpUtils.cs @@ -208,10 +208,9 @@ public static MemberAccessExpressionSyntax MemberAccess(IReadOnlyList qu /// public static bool IsWithinNameOf([NotNullWhen(true)] SyntaxNode? syntaxNode) { - InvocationExpressionSyntax? invocation = syntaxNode?.FirstAncestorOrSelf(); - return invocation is object - && (invocation.Expression as IdentifierNameSyntax)?.Identifier.Text == "nameof" - && invocation.ArgumentList.Arguments.Count == 1; + return syntaxNode?.AncestorsAndSelf().OfType().Any( + invocation => (invocation.Expression as IdentifierNameSyntax)?.Identifier.Text == "nameof" + && invocation.ArgumentList.Arguments.Count == 1) is true; } public override Location? GetLocationOfBaseTypeName(INamedTypeSymbol symbol, INamedTypeSymbol baseType, Compilation compilation, CancellationToken cancellationToken) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index 3d2ba5c1e..09c179b8e 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -5,6 +5,7 @@ using System.Collections.Generic; using System.Collections.Immutable; using System.Linq; +using System.Text.RegularExpressions; using System.Threading; using System.Threading.Tasks; using Microsoft.CodeAnalysis; @@ -61,7 +62,15 @@ public override void Initialize(AnalysisContext context) context.RegisterCompilationStartAction(compilationContext => { INamedTypeSymbol? taskSymbol = compilationContext.Compilation.GetTypeByMetadataName(Types.Task.FullName); - if (taskSymbol is object) + ImmutableArray configuredSyncBlockingMethods = CommonInterest.ReadMethods( + compilationContext.Options, + new Regex(@"^vs-threading\.SyncBlockingMethods(\..*)?.txt$", RegexOptions.IgnoreCase | RegexOptions.Singleline), + compilationContext.CancellationToken).ToImmutableArray(); + ImmutableArray methodsExcludedFromVSTHRD103 = CommonInterest.ReadMethods( + compilationContext.Options, + CommonInterest.FileNamePatternForSyncMethodsToExcludeFromVSTHRD103, + compilationContext.CancellationToken).ToImmutableArray(); + if (taskSymbol is object || !configuredSyncBlockingMethods.IsEmpty) { compilationContext.RegisterCodeBlockStartAction(codeBlockContext => { @@ -70,8 +79,13 @@ public override void Initialize(AnalysisContext context) if (propertySymbol is object || methodSymbol is object) { bool analyzeWholeCodeBlock = propertySymbol is object || !methodSymbol!.HasAsyncCompatibleReturnType(); - codeBlockContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(c => AnalyzeInvocation(c, taskSymbol, analyzeWholeCodeBlock)), SyntaxKind.InvocationExpression); - codeBlockContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(c => AnalyzeMemberAccess(c, taskSymbol, analyzeWholeCodeBlock)), SyntaxKind.SimpleMemberAccessExpression); + codeBlockContext.RegisterSyntaxNodeAction( + Utils.DebuggableWrapper(c => AnalyzeInvocation(c, configuredSyncBlockingMethods, methodsExcludedFromVSTHRD103, analyzeWholeCodeBlock, taskSymbol is object)), + SyntaxKind.InvocationExpression); + if (taskSymbol is object) + { + codeBlockContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(c => AnalyzeMemberAccess(c, analyzeWholeCodeBlock)), SyntaxKind.SimpleMemberAccessExpression); + } } }); } @@ -98,80 +112,138 @@ private static bool ShouldAnalyze(SyntaxNodeAnalysisContext context, bool analyz && !containingMethod.HasAsyncCompatibleReturnType(); } - private static ParameterSyntax? GetFirstParameter(AnonymousFunctionExpressionSyntax? anonymousFunctionSyntax) + private static void InspectMemberAccess( + SyntaxNodeAnalysisContext context, + MemberAccessExpressionSyntax? memberAccessSyntax, + IEnumerable problematicMethods) { - switch (anonymousFunctionSyntax) + if (memberAccessSyntax is null) { - case SimpleLambdaExpressionSyntax lambda: - return lambda.Parameter; - case ParenthesizedLambdaExpressionSyntax lambda: - return lambda.ParameterList.Parameters.FirstOrDefault(); - case AnonymousMethodExpressionSyntax anonymousMethod: - return anonymousMethod.ParameterList?.Parameters.FirstOrDefault(); + return; } - return null; + CSharpCommonInterest.InspectMemberAccess(context, memberAccessSyntax, Descriptor, problematicMethods); } - private static void InspectMemberAccess( + private static void AnalyzeInvocation( SyntaxNodeAnalysisContext context, - MemberAccessExpressionSyntax? memberAccessSyntax, - IEnumerable problematicMethods, - INamedTypeSymbol taskSymbol) + ImmutableArray configuredSyncBlockingMethods, + ImmutableArray methodsExcludedFromVSTHRD103, + bool analyzeWholeCodeBlock, + bool analyzeBuiltInBlockingMethods) { - if (memberAccessSyntax is null) + var invocationExpressionSyntax = (InvocationExpressionSyntax)context.Node; + if (analyzeBuiltInBlockingMethods && ShouldAnalyze(context, analyzeWholeCodeBlock)) + { + if (invocationExpressionSyntax.Expression is MemberAccessExpressionSyntax memberAccess) + { + InspectMemberAccess(context, memberAccess, CommonInterest.ProblematicSyncBlockingMethods); + } + else if (invocationExpressionSyntax.Expression is MemberBindingExpressionSyntax memberBinding + && invocationExpressionSyntax.FirstAncestorOrSelf() is { } conditionalAccess) + { + CSharpCommonInterest.InspectMemberBinding( + context, + memberBinding, + conditionalAccess.Expression, + conditionalAccess, + Descriptor, + CommonInterest.ProblematicSyncBlockingMethods); + } + } + + if (configuredSyncBlockingMethods.IsEmpty + || context.SemanticModel.GetSymbolInfo(invocationExpressionSyntax, context.CancellationToken).Symbol is not IMethodSymbol invokedMethod) + { + return; + } + + IMethodSymbol methodDefinition = invokedMethod.ReducedFrom ?? invokedMethod; + bool isConfiguredSyncBlockingMethod = configuredSyncBlockingMethods.Any( + method => method.IsMatch(invokedMethod) || method.IsMatch(methodDefinition)); + if (!isConfiguredSyncBlockingMethod) { return; } - // Are we in the context of an anonymous function that is passed directly in as an argument to another method? - AnonymousFunctionExpressionSyntax? anonymousFunctionSyntax = context.Node.FirstAncestorOrSelf(); - var anonFuncAsArgument = anonymousFunctionSyntax?.Parent as ArgumentSyntax; - var invocationPassingExpression = anonFuncAsArgument?.Parent?.Parent as InvocationExpressionSyntax; - var invokedMemberAccess = invocationPassingExpression?.Expression as MemberAccessExpressionSyntax; - if (invokedMemberAccess?.Name is object) + bool isBuiltInSyncBlockingMethod = CommonInterest.ProblematicSyncBlockingMethods.Any( + method => method.Method.IsMatch(invokedMethod) || method.Method.IsMatch(methodDefinition)); + bool coveredByVSTHRD103 = !methodsExcludedFromVSTHRD103.Contains(invokedMethod) + && !methodsExcludedFromVSTHRD103.Contains(methodDefinition) + && !invokedMethod.Name.EndsWith(VSTHRD200UseAsyncNamingConventionAnalyzer.MandatoryAsyncSuffix, StringComparison.CurrentCulture) + && !invokedMethod.HasAsyncCompatibleReturnType() + && IsInTaskReturningMethodOrDelegate(context) + && HasAsyncAlternative(context, invocationExpressionSyntax, invokedMethod); + if (!isBuiltInSyncBlockingMethod + && !coveredByVSTHRD103) { - // Does the anonymous function appear as the first argument to Task.ContinueWith? - var invokedMemberSymbol = context.SemanticModel.GetSymbolInfo(invokedMemberAccess.Name, context.CancellationToken).Symbol as IMethodSymbol; - if (invokedMemberSymbol?.Name == nameof(Task.ContinueWith) && - Utils.IsEqualToOrDerivedFrom(invokedMemberSymbol?.ContainingType, taskSymbol) && - invocationPassingExpression?.ArgumentList?.Arguments.FirstOrDefault() == anonFuncAsArgument) + SimpleNameSyntax? methodName = invocationExpressionSyntax.Expression switch { - // Does the member access being analyzed belong to the Task that just completed? - ParameterSyntax? firstParameter = GetFirstParameter(anonymousFunctionSyntax); - if (firstParameter is object) - { - // Are we accessing a member of the completed task? - ISymbol? invokedObjectSymbol = context.SemanticModel.GetSymbolInfo(memberAccessSyntax.Expression, context.CancellationToken).Symbol; - IParameterSymbol? completedTask = context.SemanticModel.GetDeclaredSymbol(firstParameter); - if (EqualityComparer.Default.Equals(invokedObjectSymbol, completedTask)) - { - // Skip analysis since Task.Result (et. al) of a completed Task is fair game. - return; - } - } + MemberAccessExpressionSyntax memberAccess => memberAccess.Name, + MemberBindingExpressionSyntax memberBinding => memberBinding.Name, + SimpleNameSyntax simpleName => simpleName, + _ => null, + }; + + if (methodName is object + && !CSharpCommonInterest.ShouldIgnoreContext(context) + && !CSharpUtils.IsWithinNameOf(invocationExpressionSyntax)) + { + ImmutableDictionary properties = ImmutableDictionary.Empty.Add("SuppressAwaitCodeFix", null); + context.ReportDiagnostic(Diagnostic.Create(Descriptor, methodName.GetLocation(), properties)); } } - - CSharpCommonInterest.InspectMemberAccess(context, memberAccessSyntax, Descriptor, problematicMethods); } - private static void AnalyzeInvocation(SyntaxNodeAnalysisContext context, INamedTypeSymbol taskSymbol, bool analyzeWholeCodeBlock) + private static bool HasAsyncAlternative( + SyntaxNodeAnalysisContext context, + InvocationExpressionSyntax invocation, + IMethodSymbol invokedMethod) { - if (!ShouldAnalyze(context, analyzeWholeCodeBlock)) + string asyncMethodName = invokedMethod.Name + VSTHRD200UseAsyncNamingConventionAnalyzer.MandatoryAsyncSuffix; + INamespaceOrTypeSymbol lookupContainer = invokedMethod.ContainingType; + if (invokedMethod.ReducedFrom is object) { - return; + ExpressionSyntax? receiver = invocation.Expression is MemberAccessExpressionSyntax memberAccess + ? memberAccess.Expression + : invocation.FirstAncestorOrSelf()?.Expression; + if (receiver is null + || context.SemanticModel.GetTypeInfo(receiver, context.CancellationToken).Type is not INamespaceOrTypeSymbol receiverType) + { + return false; + } + + lookupContainer = receiverType; } - var invocationExpressionSyntax = (InvocationExpressionSyntax)context.Node; - InspectMemberAccess( - context, - invocationExpressionSyntax.Expression as MemberAccessExpressionSyntax, - CommonInterest.ProblematicSyncBlockingMethods, - taskSymbol); + string? declaringMethodName = invocation.FirstAncestorOrSelf()?.Identifier.Text; + return context.SemanticModel.LookupSymbols( + invocation.Expression.SpanStart, + lookupContainer, + asyncMethodName, + includeReducedExtensionMethods: true) + .OfType() + .Any(candidate => !candidate.IsObsolete() + && candidate.Name != declaringMethodName + && candidate.HasAsyncCompatibleReturnType() + && CSharpCommonInterest.IsApplicableAsyncAlternative(context, invocation, candidate)); + } + + private static bool IsInTaskReturningMethodOrDelegate(SyntaxNodeAnalysisContext context) + { + SyntaxNode? containingFunction = context.Node.Ancestors().FirstOrDefault( + node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax or MethodDeclarationSyntax); + IMethodSymbol? containingMethod = containingFunction switch + { + AnonymousFunctionExpressionSyntax anonymousFunction => context.SemanticModel.GetSymbolInfo(anonymousFunction, context.CancellationToken).Symbol as IMethodSymbol, + LocalFunctionStatementSyntax localFunction => context.SemanticModel.GetDeclaredSymbol(localFunction, context.CancellationToken), + MethodDeclarationSyntax method => context.SemanticModel.GetDeclaredSymbol(method, context.CancellationToken), + _ => null, + }; + return containingMethod?.HasAsyncCompatibleReturnType() is true; } - private static void AnalyzeMemberAccess(SyntaxNodeAnalysisContext context, INamedTypeSymbol taskSymbol, bool analyzeWholeCodeBlock) + private static void AnalyzeMemberAccess(SyntaxNodeAnalysisContext context, bool analyzeWholeCodeBlock) { if (!ShouldAnalyze(context, analyzeWholeCodeBlock)) { @@ -182,7 +254,6 @@ private static void AnalyzeMemberAccess(SyntaxNodeAnalysisContext context, IName InspectMemberAccess( context, memberAccessSyntax, - CommonInterest.SyncBlockingProperties, - taskSymbol); + CommonInterest.SyncBlockingProperties); } } diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs index fa23bd4fe..3d76461be 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs @@ -110,7 +110,12 @@ internal void AnalyzeConditionalAccessExpression(SyntaxNodeAnalysisContext conte MemberBindingExpressionSyntax bindingExpr => bindingExpr.Name, _ => conditionalAccessSyntax.WhenNotNull, }; - this.InspectMemberAccess(context, rightSide, CommonInterest.SyncBlockingProperties); + this.InspectMemberAccess( + context, + rightSide, + CommonInterest.SyncBlockingProperties, + conditionalAccessSyntax.Expression, + conditionalAccessSyntax); } } @@ -119,8 +124,19 @@ internal void AnalyzeInvocation(SyntaxNodeAnalysisContext context) if (IsInTaskReturningMethodOrDelegate(context)) { var invocationExpressionSyntax = (InvocationExpressionSyntax)context.Node; - var memberAccessSyntax = invocationExpressionSyntax.Expression as MemberAccessExpressionSyntax; - if (memberAccessSyntax is not null && this.InspectMemberAccess(context, memberAccessSyntax.Name, CommonInterest.SyncBlockingMethods)) + bool handledBlockingMember = invocationExpressionSyntax.Expression switch + { + MemberAccessExpressionSyntax memberAccess => this.InspectMemberAccess(context, memberAccess.Name, CommonInterest.SyncBlockingMethods), + MemberBindingExpressionSyntax memberBinding when invocationExpressionSyntax.FirstAncestorOrSelf() is { } conditionalAccess => + this.InspectMemberAccess( + context, + memberBinding.Name, + CommonInterest.SyncBlockingMethods, + conditionalAccess.Expression, + conditionalAccess), + _ => false, + }; + if (handledBlockingMember) { // Don't return double-diagnostics. return; @@ -153,12 +169,12 @@ internal void AnalyzeInvocation(SyntaxNodeAnalysisContext context) foreach (IMethodSymbol m in symbols.OfType()) { if (!m.IsObsolete() - && HasSupersetOfParameterTypes(m, methodSymbol) + && CSharpCommonInterest.IsApplicableAsyncAlternative(context, invocationExpressionSyntax, m) && m.Name != invocationDeclaringMethod?.Identifier.Text && m.HasAsyncCompatibleReturnType()) { // Check if this method is excluded from VSTHRD103 diagnostics - if (this.excludedMethods.Contains(methodSymbol)) + if (this.IsExcluded(methodSymbol)) { return; } @@ -182,24 +198,6 @@ internal void AnalyzeInvocation(SyntaxNodeAnalysisContext context) } } - /// - /// Determines whether the given method has parameters to cover all the parameter types in another method. - /// - /// The candidate method. - /// The baseline method. - /// - /// if has a superset of parameter types found in ; otherwise . - /// - private static bool HasSupersetOfParameterTypes(IMethodSymbol candidateMethod, IMethodSymbol baselineMethod) - { - if (baselineMethod.Parameters.Length > candidateMethod.Parameters.Length) - { - return false; - } - - return baselineMethod.Parameters.All(baselineParameter => candidateMethod.Parameters.Any(candidateParameter => baselineParameter.Type?.Equals(candidateParameter.Type, SymbolEqualityComparer.Default) ?? false)); - } - private static bool IsInTaskReturningMethodOrDelegate(SyntaxNodeAnalysisContext context) { // We want to scan invocations that occur inside Task and Task-returning delegates or methods. @@ -231,7 +229,12 @@ private static bool IsInTaskReturningMethodOrDelegate(SyntaxNodeAnalysisContext return methodSymbol?.HasAsyncCompatibleReturnType() is true; } - private bool InspectMemberAccess(SyntaxNodeAnalysisContext context, ExpressionSyntax memberName, IEnumerable problematicMethods) + private bool InspectMemberAccess( + SyntaxNodeAnalysisContext context, + ExpressionSyntax memberName, + IEnumerable problematicMethods, + ExpressionSyntax? taskReceiver = null, + SyntaxNode? accessSyntax = null) { ISymbol? memberSymbol = context.SemanticModel.GetSymbolInfo(memberName, context.CancellationToken).Symbol; if (memberSymbol is object) @@ -240,16 +243,16 @@ private bool InspectMemberAccess(SyntaxNodeAnalysisContext context, ExpressionSy { if (item.Method.IsMatch(memberSymbol)) { - if (memberSymbol is IPropertySymbol { Name: nameof(Task.Result), ContainingType: { } containingType } - && Utils.IsTask(containingType) - && memberName.Parent is MemberAccessExpressionSyntax resultAccess - && TaskCompletionAnalysis.IsTaskKnownToBeCompleted(context, resultAccess)) + if ((memberName.Parent is MemberAccessExpressionSyntax memberAccess + && CSharpCommonInterest.HasTaskCompleted(context, memberAccess)) + || (taskReceiver is object + && CSharpCommonInterest.HasTaskCompleted(context, taskReceiver, accessSyntax ?? memberName))) { - return false; + return true; } // Check if this method is excluded from VSTHRD103 diagnostics - if (this.excludedMethods.Contains(memberSymbol)) + if (this.IsExcluded(memberSymbol)) { return false; } @@ -281,5 +284,10 @@ private bool InspectMemberAccess(SyntaxNodeAnalysisContext context, ExpressionSy return false; } + + private bool IsExcluded(ISymbol symbol) + => this.excludedMethods.Contains(symbol) + || (symbol is IMethodSymbol { ReducedFrom: { } reducedFrom } + && this.excludedMethods.Contains(reducedFrom)); } } diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index 78aff172a..d2f109ec2 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -6,6 +6,7 @@ using System.Collections.Immutable; using System.Diagnostics.CodeAnalysis; using System.Linq; +using System.Runtime.CompilerServices; using System.Text; using System.Threading; using System.Threading.Tasks; @@ -14,6 +15,7 @@ using Microsoft.CodeAnalysis.CodeFixes; using Microsoft.CodeAnalysis.CSharp; using Microsoft.CodeAnalysis.CSharp.Syntax; +using Microsoft.CodeAnalysis.FindSymbols; using Microsoft.CodeAnalysis.Simplification; using Microsoft.VisualStudio.Threading; @@ -22,6 +24,8 @@ namespace Microsoft.VisualStudio.Threading.Analyzers; [ExportCodeFixProvider(LanguageNames.CSharp)] public class VSTHRD002UseJtfRunCodeFixWithAwait : CodeFixProvider { + private const string SuppressAwaitCodeFixProperty = "SuppressAwaitCodeFix"; + private static readonly ImmutableArray ReusableFixableDiagnosticIds = ImmutableArray.Create( VSTHRD002UseJtfRunAnalyzer.Id); @@ -30,11 +34,27 @@ public class VSTHRD002UseJtfRunCodeFixWithAwait : CodeFixProvider public override async Task RegisterCodeFixesAsync(CodeFixContext context) { Diagnostic? diagnostic = context.Diagnostics.First(); + if (diagnostic.Properties.ContainsKey(SuppressAwaitCodeFixProperty)) + { + return; + } SyntaxNode root = await context.Document.GetSyntaxRootOrThrowAsync(context.CancellationToken).ConfigureAwait(false); - if (TryFindNodeAtSource(diagnostic, root, out _, out _)) + if (TryFindNodeAtSource(diagnostic, root, out ExpressionSyntax? target, out _)) { + SemanticModel? semanticModel = await context.Document.GetSemanticModelAsync(context.CancellationToken).ConfigureAwait(false); + MethodDeclarationSyntax? containingMethod = target.FirstAncestorOrSelf(); + if (semanticModel is null + || containingMethod is null + || IsAwaitForbiddenAt(target, containingMethod) + || !await CanConvertToAsyncAsync(context.Document, semanticModel, containingMethod, context.CancellationToken).ConfigureAwait(false) + || semanticModel.GetDiagnostics(target.FullSpan, context.CancellationToken).Any(d => d.Severity == DiagnosticSeverity.Error) + || !CanUseAwaitCodeFix(semanticModel, target, context.CancellationToken)) + { + return; + } + context.RegisterCodeFix( CodeAction.Create( Strings.VSTHRD002_CodeFix_Await_Title, @@ -46,7 +66,7 @@ public override async Task RegisterCodeFixesAsync(CodeFixContext context) (document, node, _) = await FixUtils.UpdateDocumentAsync( document, node, - n => SyntaxFactory.AwaitExpression(transform(n, ct)), + n => ParenthesizeAwaitIfRequired(SyntaxFactory.AwaitExpression(transform(n, ct)), n), ct).ConfigureAwait(false); MethodDeclarationSyntax? method = node.FirstAncestorOrSelf(); if (method is object) @@ -65,6 +85,257 @@ public override async Task RegisterCodeFixesAsync(CodeFixContext context) /// public override FixAllProvider GetFixAllProvider() => WellKnownFixAllProviders.BatchFixer; + private static ExpressionSyntax ParenthesizeAwaitIfRequired(AwaitExpressionSyntax awaitExpression, ExpressionSyntax replacedExpression) + => replacedExpression.Parent is MemberAccessExpressionSyntax + or ElementAccessExpressionSyntax + or InvocationExpressionSyntax + or ConditionalAccessExpressionSyntax + or PostfixUnaryExpressionSyntax + ? SyntaxFactory.ParenthesizedExpression(awaitExpression) + : awaitExpression; + + private static async Task CanConvertToAsyncAsync( + Document document, + SemanticModel semanticModel, + MethodDeclarationSyntax method, + CancellationToken cancellationToken) + { + if (!IsMethodLocallyConvertible(semanticModel, method, cancellationToken, out IMethodSymbol? methodSymbol)) + { + return false; + } + + if (method.Modifiers.Any(SyntaxKind.AsyncKeyword)) + { + return true; + } + + bool changesContract = !methodSymbol.HasAsyncCompatibleReturnType(); + if (!changesContract) + { + return true; + } + + if (!CanChangeMethodContract(method, methodSymbol) + || await HasMethodGroupReferenceAsync(document.Project.Solution, methodSymbol, cancellationToken).ConfigureAwait(false)) + { + return false; + } + + var visitedMethods = new HashSet(SymbolEqualityComparer.Default); + return await CanConvertCallerChainAsync( + document.Project.Solution, + methodSymbol, + visitedMethods, + cancellationToken).ConfigureAwait(false); + } + + private static bool IsAwaitForbiddenAt(SyntaxNode target, MethodDeclarationSyntax containingMethod) + => target.AncestorsAndSelf().TakeWhile(node => node != containingMethod) + .Any(node => node is LockStatementSyntax + or CatchFilterClauseSyntax + or UnsafeStatementSyntax + or FixedStatementSyntax) + || containingMethod.AncestorsAndSelf() + .OfType() + .Any(member => member.Modifiers.Any(SyntaxKind.UnsafeKeyword)); + + private static bool IsMethodLocallyConvertible( + SemanticModel semanticModel, + MethodDeclarationSyntax method, + CancellationToken cancellationToken, + [NotNullWhen(true)] out IMethodSymbol? methodSymbol) + { + methodSymbol = semanticModel.GetDeclaredSymbol(method, cancellationToken); + return methodSymbol is object + && !methodSymbol.Parameters.Any(parameter => parameter.RefKind != RefKind.None || parameter.Type.IsRefLikeType) + && !methodSymbol.ReturnsByRef + && !methodSymbol.ReturnsByRefReadonly + && !methodSymbol.ReturnType.IsRefLikeType + && !method.AncestorsAndSelf() + .OfType() + .Any(member => member.Modifiers.Any(SyntaxKind.UnsafeKeyword)) + && !method.DescendantNodes( + node => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax) + .OfType() + .Any() + && !method.DescendantNodes( + node => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax) + .Any(node => node switch + { + VariableDeclaratorSyntax variable => IsUnsupportedLocal(semanticModel.GetDeclaredSymbol(variable, cancellationToken)), + SingleVariableDesignationSyntax designation => IsUnsupportedLocal(semanticModel.GetDeclaredSymbol(designation, cancellationToken)), + _ => false, + }); + } + + private static bool IsUnsupportedLocal(ISymbol? symbol) + => symbol is ILocalSymbol local + && (local.RefKind != RefKind.None || local.Type.IsRefLikeType); + + private static bool CanChangeMethodContract(MethodDeclarationSyntax method, IMethodSymbol methodSymbol) + => !method.Modifiers.Any(SyntaxKind.PartialKeyword) + && methodSymbol.ContainingType.TypeKind != TypeKind.Interface + && !methodSymbol.IsVirtual + && !methodSymbol.IsOverride + && !methodSymbol.FindInterfacesImplemented().Any() + && !HasAsyncNameCollision(methodSymbol); + + private static bool HasAsyncNameCollision(IMethodSymbol method) + { + if (method.Name.EndsWith(VSTHRD200UseAsyncNamingConventionAnalyzer.MandatoryAsyncSuffix, StringComparison.Ordinal)) + { + return false; + } + + string asyncName = method.Name + VSTHRD200UseAsyncNamingConventionAnalyzer.MandatoryAsyncSuffix; + return method.ContainingType.GetMembers(asyncName) + .OfType() + .Any(candidate => candidate.Arity == method.Arity + && candidate.Parameters.Length >= method.Parameters.Length + && candidate.Parameters.Take(method.Parameters.Length).Zip(method.Parameters, ParametersHaveEquivalentSignatures).All(match => match) + && candidate.Parameters.Skip(method.Parameters.Length).All(parameter => parameter.IsOptional)); + } + + private static bool ParametersHaveEquivalentSignatures(IParameterSymbol left, IParameterSymbol right) + => left.RefKind == right.RefKind && SignatureTypesMatch(left.Type, right.Type); + + private static bool SignatureTypesMatch(ITypeSymbol left, ITypeSymbol right) + { + if (left is ITypeParameterSymbol { TypeParameterKind: TypeParameterKind.Method } leftTypeParameter + && right is ITypeParameterSymbol { TypeParameterKind: TypeParameterKind.Method } rightTypeParameter) + { + return leftTypeParameter.Ordinal == rightTypeParameter.Ordinal; + } + + if (left is IArrayTypeSymbol leftArray && right is IArrayTypeSymbol rightArray) + { + return leftArray.Rank == rightArray.Rank + && SignatureTypesMatch(leftArray.ElementType, rightArray.ElementType); + } + + if (left is IPointerTypeSymbol leftPointer && right is IPointerTypeSymbol rightPointer) + { + return SignatureTypesMatch(leftPointer.PointedAtType, rightPointer.PointedAtType); + } + + if (left is INamedTypeSymbol leftNamed && right is INamedTypeSymbol rightNamed) + { + return SymbolEqualityComparer.Default.Equals(leftNamed.OriginalDefinition, rightNamed.OriginalDefinition) + && leftNamed.TypeArguments.Length == rightNamed.TypeArguments.Length + && leftNamed.TypeArguments.Zip(rightNamed.TypeArguments, SignatureTypesMatch).All(match => match); + } + + return SymbolEqualityComparer.Default.Equals(left, right) + || (left.TypeKind == TypeKind.Dynamic && right.SpecialType == SpecialType.System_Object) + || (right.TypeKind == TypeKind.Dynamic && left.SpecialType == SpecialType.System_Object); + } + + private static async Task CanConvertCallerChainAsync( + Solution solution, + IMethodSymbol method, + HashSet visitedMethods, + CancellationToken cancellationToken) + { + if (!visitedMethods.Add(method.OriginalDefinition)) + { + return true; + } + + IEnumerable callers = await SymbolFinder.FindCallersAsync(method, solution, cancellationToken).ConfigureAwait(false); + foreach (SymbolCallerInfo caller in callers) + { + foreach (Location location in caller.Locations) + { + Document? document = location.SourceTree is object ? solution.GetDocument(location.SourceTree) : null; + SyntaxNode? root = document is object ? await document.GetSyntaxRootAsync(cancellationToken).ConfigureAwait(false) : null; + InvocationExpressionSyntax? invocation = root? + .FindNode(location.SourceSpan, getInnermostNodeForTie: true) + .FirstAncestorOrSelf(); + MethodDeclarationSyntax? callingMethod = invocation?.FirstAncestorOrSelf(); + if (document is null + || invocation is null + || callingMethod is null + || IsAwaitForbiddenAt(invocation, callingMethod) + || invocation.Ancestors().TakeWhile(node => node != callingMethod) + .Any(node => node is AnonymousFunctionExpressionSyntax + or LocalFunctionStatementSyntax + or ConditionalAccessExpressionSyntax)) + { + return false; + } + + SemanticModel? semanticModel = await document.GetSemanticModelAsync(cancellationToken).ConfigureAwait(false); + if (semanticModel is null + || !IsMethodLocallyConvertible(semanticModel, callingMethod, cancellationToken, out IMethodSymbol? callingMethodSymbol)) + { + return false; + } + + if (callingMethodSymbol.ReturnType is INamedTypeSymbol callerReturnType + && callerReturnType.IsAsyncCompatibleReturnType() + && ((invocation.FirstAncestorOrSelf() is { Expression: { } returnExpression } + && returnExpression.FullSpan.Contains(invocation.Span)) + || (callingMethod.ExpressionBody?.Expression.FullSpan.Contains(invocation.Span) is true))) + { + return false; + } + + if (!callingMethodSymbol.HasAsyncCompatibleReturnType() + && (!CanChangeMethodContract(callingMethod, callingMethodSymbol) + || await HasMethodGroupReferenceAsync(solution, callingMethodSymbol, cancellationToken).ConfigureAwait(false) + || !await CanConvertCallerChainAsync(solution, callingMethodSymbol, visitedMethods, cancellationToken).ConfigureAwait(false))) + { + return false; + } + } + } + + return true; + } + + private static async Task HasMethodGroupReferenceAsync( + Solution solution, + IMethodSymbol method, + CancellationToken cancellationToken) + { + IEnumerable references = await SymbolFinder.FindReferencesAsync(method, solution, cancellationToken).ConfigureAwait(false); + foreach (ReferenceLocation reference in references.SelectMany(result => result.Locations)) + { + SyntaxNode? root = await reference.Document.GetSyntaxRootAsync(cancellationToken).ConfigureAwait(false); + if (root is null) + { + continue; + } + + SyntaxNode referenceNode = root.FindNode(reference.Location.SourceSpan, getInnermostNodeForTie: true); + SimpleNameSyntax? methodName = referenceNode.FirstAncestorOrSelf(); + if (methodName is null) + { + return true; + } + + if (CSharpUtils.IsWithinNameOf(methodName)) + { + continue; + } + + ExpressionSyntax invokedExpression = methodName.Parent switch + { + MemberAccessExpressionSyntax memberAccess when memberAccess.Name == methodName => memberAccess, + MemberBindingExpressionSyntax memberBinding when memberBinding.Name == methodName => memberBinding, + _ => methodName, + }; + if (invokedExpression.Parent is not InvocationExpressionSyntax invocation + || invocation.Expression != invokedExpression) + { + return true; + } + } + + return false; + } + private static bool TryFindNodeAtSource(Diagnostic diagnostic, SyntaxNode root, [NotNullWhen(true)] out ExpressionSyntax? target, [NotNullWhen(true)] out Func? transform) { transform = null; @@ -102,19 +373,43 @@ private static bool TryFindNodeAtSource(Diagnostic diagnostic, SyntaxNode root, return from.ReplaceToken(name.Identifier, SyntaxFactory.Identifier(newIdentifier)).WithoutAnnotations(FixUtils.BookmarkAnnotationName); } - ExpressionSyntax? FindTwoLevelDeepIdentifierInvocation(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) => - ((((from as InvocationExpressionSyntax)?.Expression as MemberAccessExpressionSyntax)?.Expression as InvocationExpressionSyntax)?.Expression as MemberAccessExpressionSyntax)?.Expression; - ExpressionSyntax? FindOneLevelDeepIdentifierInvocation(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) => - ((from as InvocationExpressionSyntax)?.Expression as MemberAccessExpressionSyntax)?.Expression; + ExpressionSyntax? FindGetAwaiterReceiver(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) + { + var getResultAccess = (from as InvocationExpressionSyntax)?.Expression as MemberAccessExpressionSyntax; + if (getResultAccess?.Name.Identifier.ValueText != nameof(TaskAwaiter.GetResult)) + { + return null; + } + + ExpressionSyntax getAwaiterInvocationExpression = getResultAccess.Expression; + while (getAwaiterInvocationExpression is ParenthesizedExpressionSyntax parenthesized) + { + getAwaiterInvocationExpression = parenthesized.Expression; + } + + var getAwaiterInvocation = getAwaiterInvocationExpression as InvocationExpressionSyntax; + var getAwaiterAccess = getAwaiterInvocation?.Expression as MemberAccessExpressionSyntax; + return getAwaiterAccess?.Name.Identifier.ValueText == "GetAwaiter" + && getAwaiterInvocation!.ArgumentList.Arguments.Count == 0 + ? getAwaiterAccess.Expression + : null; + } + + ExpressionSyntax? FindInstanceWaitReceiver(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) + { + var waitAccess = (from as InvocationExpressionSyntax)?.Expression as MemberAccessExpressionSyntax; + return waitAccess?.Name.Identifier.ValueText == nameof(Task.Wait) ? waitAccess.Expression : null; + } + ExpressionSyntax? FindParentMemberAccess(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) => (from as MemberAccessExpressionSyntax)?.Expression; InvocationExpressionSyntax? parentInvocation = syntaxNode.FirstAncestorOrSelf(); MemberAccessExpressionSyntax? parentMemberAccess = syntaxNode.FirstAncestorOrSelf(); - if (FindTwoLevelDeepIdentifierInvocation(parentInvocation) is object) + if (FindGetAwaiterReceiver(parentInvocation) is object) { // This method will not return null for the provided 'target' argument - transform = NullableHelpers.AsNonNullReturnUnchecked(FindTwoLevelDeepIdentifierInvocation); + transform = NullableHelpers.AsNonNullReturnUnchecked(FindGetAwaiterReceiver); target = parentInvocation!; return true; } @@ -125,14 +420,15 @@ private static bool TryFindNodeAtSource(Diagnostic diagnostic, SyntaxNode root, target = parentInvocation!; return true; } - else if (FindOneLevelDeepIdentifierInvocation(parentInvocation) is object) + else if (FindInstanceWaitReceiver(parentInvocation) is object) { // This method will not return null for the provided 'target' argument - transform = NullableHelpers.AsNonNullReturnUnchecked(FindOneLevelDeepIdentifierInvocation); + transform = NullableHelpers.AsNonNullReturnUnchecked(FindInstanceWaitReceiver); target = parentInvocation!; return true; } - else if (FindParentMemberAccess(parentMemberAccess) is object) + else if (parentMemberAccess?.Name.Identifier.ValueText == nameof(Task.Result) + && FindParentMemberAccess(parentMemberAccess) is object) { // This method will not return null for the provided 'target' argument transform = NullableHelpers.AsNonNullReturnUnchecked(FindParentMemberAccess); @@ -144,4 +440,64 @@ private static bool TryFindNodeAtSource(Diagnostic diagnostic, SyntaxNode root, return false; } } + + private static bool CanUseAwaitCodeFix(SemanticModel semanticModel, ExpressionSyntax target, CancellationToken cancellationToken) + { + if (target is not InvocationExpressionSyntax invocation + || semanticModel.GetSymbolInfo(invocation, cancellationToken).Symbol is not IMethodSymbol method) + { + return true; + } + + if (method.Name == nameof(Task.Wait)) + { + return method.ReducedFrom is null + && !method.IsStatic + && Utils.IsTask(method.ContainingType) + && method.Parameters.IsEmpty; + } + + if (method.Name == nameof(TaskAwaiter.GetResult)) + { + ExpressionSyntax? getAwaiterExpression = (invocation.Expression as MemberAccessExpressionSyntax)?.Expression; + while (getAwaiterExpression is ParenthesizedExpressionSyntax parenthesized) + { + getAwaiterExpression = parenthesized.Expression; + } + + if (getAwaiterExpression is not InvocationExpressionSyntax { ArgumentList.Arguments.Count: 0 } getAwaiterInvocation + || semanticModel.GetSymbolInfo(getAwaiterInvocation, cancellationToken).Symbol is not IMethodSymbol getAwaiterMethod + || getAwaiterMethod.ReducedFrom is object + || getAwaiterMethod.IsStatic + || !getAwaiterMethod.Parameters.IsEmpty) + { + return false; + } + + ExpressionSyntax? receiver = (getAwaiterInvocation.Expression as MemberAccessExpressionSyntax)?.Expression; + if (receiver is not { } awaitableReceiver + || semanticModel.GetSymbolInfo(awaitableReceiver, cancellationToken).Symbol is INamedTypeSymbol) + { + return false; + } + + ITypeSymbol? receiverType = semanticModel.GetTypeInfo(awaitableReceiver, cancellationToken).Type; + return receiverType.IsAwaitable(semanticModel, awaitableReceiver.SpanStart); + } + + if (method.Name is nameof(Task.WaitAll) or nameof(Task.WaitAny) + && Utils.IsTask(method.ContainingType)) + { + if (method.Name == nameof(Task.WaitAny) && invocation.Parent is not ExpressionStatementSyntax) + { + return false; + } + + return method.Parameters.All( + parameter => Utils.IsTask(parameter.Type) + || (parameter.Type is IArrayTypeSymbol arrayType && Utils.IsTask(arrayType.ElementType))); + } + + return true; + } } diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs index 9d4a61437..5f2a2e609 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs @@ -32,7 +32,7 @@ Task FooAsync() { static void SetTaskSourceIfCompleted(Task task, TaskCompletionSource tcs) { if (task.IsCompleted) { - tcs.SetResult(task.Result); + tcs.SetResult(task.Result); // No VSTHRD002 because task is known to be completed. } } }"; @@ -41,7 +41,6 @@ static void SetTaskSourceIfCompleted(Task task, TaskCompletionSource tc { CSVerify.Diagnostic(VSTHRD103UseAsyncOptionAnalyzer.DescriptorNoAlternativeMethod).WithSpan(10, 24, 10, 33).WithArguments("GetResult"), CSVerify.Diagnostic(VSTHRD103UseAsyncOptionAnalyzer.Descriptor).WithSpan(11, 13, 11, 16).WithArguments("Run", "RunAsync"), - CSVerify.Diagnostic(VSTHRD002UseJtfRunAnalyzer.Descriptor).WithSpan(19, 32, 19, 38), }; // All expected diagnostics should include a location @@ -63,6 +62,139 @@ static void SetTaskSourceIfCompleted(Task task, TaskCompletionSource tc await verifyTest.RunAsync(); } + [Fact] + public async Task ConfiguredBlockingMethodWithAsyncAlternativeProducesOneDiagnostic() + { + var test = @" +using System.Threading.Tasks; + +class CustomWaiter { + public void Join() { } + public Task JoinAsync() => Task.CompletedTask; +} + +class Test { + Task FAsync(CustomWaiter waiter) { + waiter.Join(); + return Task.CompletedTask; + } +} +"; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + TestState = { MarkupHandling = MarkupMode.None }, + }; + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncBlockingMethods.txt", "[CustomWaiter]::Join")); + verifyTest.ExpectedDiagnostics.Add( + CSVerify.Diagnostic(VSTHRD103UseAsyncOptionAnalyzer.Descriptor) + .WithSpan(11, 16, 11, 20) + .WithArguments("Join", "JoinAsync")); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task VSTHRD103ExclusionLeavesConfiguredVSTHRD002Diagnostic() + { + var test = @" +using System.Threading.Tasks; + +class CustomWaiter { + public void Join() { } + public Task JoinAsync() => Task.CompletedTask; +} + +class Test { + Task FAsync(CustomWaiter waiter) { + waiter.Join(); + return Task.CompletedTask; + } +} +"; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + TestState = { MarkupHandling = MarkupMode.None }, + }; + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncBlockingMethods.txt", "[CustomWaiter]::Join")); + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncMethodsToExcludeFromVSTHRD103.txt", "[CustomWaiter]::Join")); + verifyTest.ExpectedDiagnostics.Add( + CSVerify.Diagnostic(VSTHRD002UseJtfRunAnalyzer.Descriptor) + .WithSpan(11, 16, 11, 20)); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task VSTHRD103ExclusionMatchesReducedExtensionDefinition() + { + var test = @" +using System.Threading.Tasks; + +class CustomWaiter { } + +static class Extensions { + public static void Join(this CustomWaiter waiter) { } + public static Task JoinAsync(this CustomWaiter waiter) => Task.CompletedTask; +} + +class Test { + Task FAsync(CustomWaiter waiter) { + waiter.Join(); + return Task.CompletedTask; + } +} +"; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + TestState = { MarkupHandling = MarkupMode.None }, + }; + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncBlockingMethods.txt", "[Extensions]::Join")); + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncMethodsToExcludeFromVSTHRD103.txt", "[Extensions]::Join")); + verifyTest.ExpectedDiagnostics.Add( + CSVerify.Diagnostic(VSTHRD002UseJtfRunAnalyzer.Descriptor) + .WithSpan(13, 16, 13, 20)); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task ConfiguredBlockerInAsyncLambdaProducesOneDiagnostic() + { + var test = @" +using System; +using System.Threading.Tasks; + +class CustomWaiter { + public void Join() { } + public Task JoinAsync() => Task.CompletedTask; +} + +class Test { + void F(CustomWaiter waiter) { + Func action = async () => { + waiter.Join(); + await Task.Yield(); + }; + } +} +"; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + TestState = { MarkupHandling = MarkupMode.None }, + }; + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncBlockingMethods.txt", "[CustomWaiter]::Join")); + verifyTest.ExpectedDiagnostics.Add( + CSVerify.Diagnostic(VSTHRD103UseAsyncOptionAnalyzer.Descriptor) + .WithSpan(13, 20, 13, 24) + .WithArguments("Join", "JoinAsync")); + await verifyTest.RunAsync(); + } + /// /// Verifies that no analyzer throws due to a missing interface member. /// diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 6ac1ac67d..3f43f8548 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -71,6 +71,22 @@ async Task FAsync() { await CSVerify.VerifyCodeFixAsync(test, expected, withFix); } + [Fact] + public async Task TaskWaitAnyDoesNotOfferCodeFixWhenResultIsConsumed() + { + var test = @" +using System.Threading.Tasks; + +class Test { + void F(Task task1, Task task2) { + int index = Task.[|WaitAny|](task1, task2); + } +} +"; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + [Fact] public async Task TaskWhenAll_CompareWithAndWithout() { @@ -204,6 +220,91 @@ void Foo() { await CSVerify.VerifyAnalyzerAsync(test, expected); } + [Fact] + public async Task TaskWhenAll_LocalArrayCompletesContainedTask() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async void GetResultAsync(Task task) + { + Task[] tasks = { task }; + await Task.WhenAll(tasks); + _ = task.Result; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task TaskWhenAll_LocalArrayReassignedByClosureDoesNotCompleteContainedTask() + { + string test = """ + using System; + using System.Threading.Tasks; + + class Test + { + async void GetResultAsync(Task task, Task replacement) + { + Task[] tasks = { task }; + Action replace = () => tasks = new Task[] { replacement }; + replace(); + await Task.WhenAll(tasks); + _ = task.[|Result|]; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task TaskWhenAll_LocalArrayDoesNotCompleteReassignedTask() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async void GetResultAsync(Task task, Task replacement) + { + Task[] tasks = { task }; + task = replacement; + await Task.WhenAll(tasks); + _ = task.[|Result|]; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task TaskWhenAll_LocalArrayElementWriteInvalidatesProof() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async void GetResultAsync(Task task, Task replacement) + { + Task[] tasks = { task }; + tasks[0] = replacement; + await Task.WhenAll(tasks); + _ = task.[|Result|]; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + [Fact] public async Task TaskWhenAll_TaskPassedByValue_NoWarning() { @@ -531,6 +632,64 @@ async Task FAsync() { await CSVerify.VerifyCodeFixAsync(test, expected, withFix); } + [Fact] + public async Task AwaitedValueTaskResultReportsWarningButGuardedResultDoesNot() + { + var test = @" +using System.Threading.Tasks; + +class Test { + async void Awaited(ValueTask task) { + await task; + _ = task.[|Result|]; + } + + void Guarded(ValueTask task) { + if (task.IsCompleted) { + _ = task.Result; + } + } +} +"; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task SynchronouslyConsumedValueTaskInvalidatesCompletionGuard() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + int GetResult(ValueTask task) + { + _ = task.[|Result|]; + if (task.IsCompletedSuccessfully) + { + return task.[|Result|]; + } + + return 0; + } + + int GetAwaiterResult(ValueTask task) + { + _ = task.GetAwaiter().[|GetResult|](); + if (task.IsCompletedSuccessfully) + { + return task.[|Result|]; + } + + return 0; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + [Fact] public async Task TaskResultShouldReportWarning_WithinAnonymousDelegate() { @@ -568,6 +727,16 @@ void F() { task.ContinueWith(t => t.Wait()); ((Task)task).ContinueWith(t => t.Wait()); task.ContinueWith((t, s) => t.Result, new object()); + task.ContinueWith(t => (t.GetAwaiter()).GetResult()); + task.ContinueWith(t => ((t.ConfigureAwait(false)).GetAwaiter()).GetResult()); + task.ContinueWith(t => { + ref Task alias = ref t; + return alias.Result; + }); + task.ContinueWith(t => { + Console.WriteLine(t.Result); + Action replaceLater = () => t = Task.Run(() => 6); + }); } void ContinueWith(Func, int> del) { } @@ -584,7 +753,29 @@ void ContinueWith(Func, int> del) { } } [Fact] - public async Task Task_GetAwaiter_GetResult_ShouldReportWarning() + public async Task TaskResultShouldNotReportWarning_WithinParenthesizedContinuationDelegate() + { + string test = """ + using System; + using System.Threading; + using System.Threading.Tasks; + + class Test + { + void F(Task task) + { + task.ContinueWith((Func, int>)(t => t.Result)); + task.ContinueWith(((t) => t.Result)); + task.ContinueWith(cancellationToken: CancellationToken.None, continuationFunction: t => t.Result); + } + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task TaskResultShouldNotReportWarning_WithinNestedDelegateInItsOwnContinuation() { var test = @" using System; @@ -592,28 +783,20 @@ public async Task Task_GetAwaiter_GetResult_ShouldReportWarning() class Test { void F() { - var task = Task.Run(() => 1); - task.GetAwaiter().GetResult(); + var task = Task.Run(() => 5); + task.ContinueWith(t => { + Action useResultLater = () => Console.WriteLine(t.Result); + useResultLater(); + }); } } "; - var withFix = @" -using System; -using System.Threading.Tasks; -class Test { - async Task FAsync() { - var task = Task.Run(() => 1); - await task; - } -} -"; - DiagnosticResult expected = CSVerify.Diagnostic().WithSpan(8, 27, 8, 36); - await CSVerify.VerifyCodeFixAsync(test, expected, withFix); + await CSVerify.VerifyAnalyzerAsync(test); } [Fact] - public async Task ConfiguredTask_GetAwaiter_GetResult_ShouldReportWarning() + public async Task TaskResultReportsWarning_WhenContinuationParameterIsReassigned() { var test = @" using System; @@ -621,192 +804,1608 @@ public async Task ConfiguredTask_GetAwaiter_GetResult_ShouldReportWarning() class Test { void F() { - var task = Task.Run(() => 1); - task.ConfigureAwait(false).GetAwaiter().[|GetResult|](); + var task = Task.Run(() => 5); + task.ContinueWith(t => { + Action useResultLater = () => Console.WriteLine(t.[|Result|]); + t = Task.Run(() => 6); + useResultLater(); + }); + task.ContinueWith(t => { + Task other = Task.Run(() => 7); + ref Task alias = ref t; + alias = ref other; + return alias.[|Result|]; + }); + task.ContinueWith(t => { + ref Task alias = ref t; + alias = Task.Run(() => 8); + Func useResultLater = () => t.[|Result|]; + return useResultLater(); + }); + task.ContinueWith(t => { + Func useResultLater = () => { + ref Task alias = ref t; + alias = Task.Run(() => 9); + return t.[|Result|]; + }; + return useResultLater(); + }); } } "; - var withFix = @" -using System; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task RefReturningInvocationCreatesPotentialAlias() + { + var test = @" using System.Threading.Tasks; class Test { - async Task FAsync() { - var task = Task.Run(() => 1); - await task.ConfigureAwait(false); + void F(Task task) { + ref Task alias = ref GetTaskRef(ref task); + if (task.IsCompleted) { + alias = Task.Run(() => 1); + _ = task.[|Result|]; + } } + + static ref Task GetTaskRef(ref Task task) => ref task; } "; - await CSVerify.VerifyCodeFixAsync(test, withFix); + + await CSVerify.VerifyCodeFixAsync(test, test); } [Fact] - public async Task ValueTask_GetAwaiter_GetResult_ShouldReportWarning() + public async Task ContinuationParameterDeconstructionReportsWarning() { var test = @" -using System; using System.Threading.Tasks; class Test { void F() { - ValueTask task = default; - task.GetAwaiter().GetResult(); - } -} -"; - var withFix = @" -using System; -using System.Threading.Tasks; + var task = Task.Run(() => 5); + task.ContinueWith(t => { + (t, _) = (Task.Run(() => 6), 0); + return t.[|Result|]; + }); + task.ContinueWith(t => { + Task other = Task.Run(() => 7); + ref Task alias = ref other; + alias = ref t; + alias = Task.Run(() => 7); + return t.[|Result|]; + }); + task.ContinueWith(t => { + Replace(); + return t.[|Result|]; -class Test { - async Task FAsync() { - ValueTask task = default; - await task; + void Replace() => t = Task.Run(() => 8); + }); } } "; - DiagnosticResult expected = CSVerify.Diagnostic().WithSpan(8, 27, 8, 36); - await CSVerify.VerifyCodeFixAsync(test, expected, withFix); + + await CSVerify.VerifyCodeFixAsync(test, test); } [Fact] - public async Task ConfiguredValueTask_GetAwaiter_GetResult_ShouldReportWarning() + public async Task TaskWhenAllResultReportsWarningWithoutAnalyzerFailure() { var test = @" -using System; using System.Threading.Tasks; class Test { void F() { - ValueTask task = default; - task.ConfigureAwait(false).GetAwaiter().[|GetResult|](); + var task = Task.Run(() => 1); + _ = Task.WhenAll(task).[|Result|]; } } "; + var withFix = @" -using System; using System.Threading.Tasks; class Test { async Task FAsync() { - ValueTask task = default; - await task.ConfigureAwait(false); + var task = Task.Run(() => 1); + _ = await Task.WhenAll(task); } } "; + await CSVerify.VerifyCodeFixAsync(test, withFix); } [Fact] - public async Task TaskResult_FixUpdatesCallers() + public async Task CompletedTaskResultDoesNotReportWarning() { - var test = new SourceFileList("Test", "cs") - { - @" + var test = @" using System; using System.Threading.Tasks; class Test { - internal static int GetNumber(int a) { - var task = Task.Run(() => a); - return task.Result; + async void Awaited() { + var task = Task.Run(() => 1); + await task; + _ = task.Result; } - int Add(int a, int b) { - return GetNumber(a) + b; + async void AwaitedAsArgument() { + var task = Task.Run(() => 1); + Consume(await task); + _ = task.Result; } - int Subtract(int a, int b) { - return GetNumber(a) - b; + async void AwaitedEarlierInSameStatement() { + var task = Task.Run(() => 1); + Consume(await task, task.Result); } - static int Main(string[] args) - { - return new Test().Add(1, 2); - } -} -", - @" -class TestClient { - int Multiply(int a, int b) { - return Test.GetNumber(a) * b; + async void AwaitedBeforeOtherTaskInSameStatement() { + var task = Task.Run(() => 1); + var otherTask = Task.Run(() => 2); + Consume(await task, await otherTask); + _ = task.Result; } -} -", - }; - var withFix = new SourceFileList("Test", "cs") - { - @" -using System; -using System.Threading.Tasks; -class Test { - internal static async Task GetNumberAsync(int a) { - var task = Task.Run(() => a); - return await task; + async void AwaitedInWhenAllArray() { + var task = Task.Run(() => 1); + await Task.WhenAll(new[] { task }); + _ = task.Result; } - async Task AddAsync(int a, int b) { - return await GetNumberAsync(a) + b; + async void AwaitedInConditionalCondition() { + var task = Task.Run(() => true); + _ = (await task) ? 1 : 0; + _ = task.Result; } - async Task SubtractAsync(int a, int b) { - return await GetNumberAsync(a) - b; - } + async void AwaitedInLeftShortCircuitOperand(bool condition) { + var task = Task.Run(() => true); + if (await task && condition) { + } - static async Task Main(string[] args) - { - return await new Test().AddAsync(1, 2); - } -} -", - @" -class TestClient { - async System.Threading.Tasks.Task MultiplyAsync(int a, int b) { - return await Test.GetNumberAsync(a) * b; + _ = task.Result; } -} -", - }; - var verifyTest = new CSVerify.Test + async void AwaitedInNestedBlock() { + var task = Task.Run(() => 1); { - TestState = - { - OutputKind = OutputKind.ConsoleApplication, - }, - ExpectedDiagnostics = - { - CSVerify.Diagnostic().WithSpan("Test0.cs", 8, 21, 8, 27), - }, - }; + await task; + } - verifyTest.TestState.Sources.AddRange(test); - verifyTest.FixedState.Sources.AddRange(withFix); - await verifyTest.RunAsync(); + _ = task.Result; } - [Fact] - public async Task DoNotReportWarningInTaskReturningMethods() - { - var test = @" -using System.Threading.Tasks; + async void Guarded() { + var task = Task.Run(() => 1); + if (!task.IsCompleted) { + await task.ConfigureAwait(false); + } -class Test { - Task F() { + _ = task.Result; + } + + async void GuardedWithElse() { var task = Task.Run(() => 1); - task.GetAwaiter().GetResult(); - return Task.CompletedTask; + if (task.IsCompleted) { + } else { + await task.ConfigureAwait(false); + } + + _ = task.Result; } -} -"; - await CSVerify.VerifyAnalyzerAsync(test); + + async void GuardedWithNestedConditional(bool condition) { + var task = Task.Run(() => 1); + if (!task.IsCompleted) { + if (condition) { + await task; + } else { + await task; + } + } + + _ = task.Result; } - [Fact] - public async Task DoNotReportWarningOnCodeGeneratedByXaml2CS() - { - var test = @" + async void GuardedInsideNestedBlock() { + var task = Task.Run(() => 1); + { + if (!task.IsCompleted) { + await task; + } + } + + _ = task.Result; + } + + async void AwaitedBeforeNestedBlock(bool condition) { + var task = Task.Run(() => 1); + await task; + if (condition) { + _ = task.Result; + } + } + + void CompletionProperties(Task task) { + if (task.IsCompleted) { + _ = task.Result; + } + + if (task.IsCanceled) { + _ = task.Result; + } + + if (task.IsFaulted) { + _ = task.Result; + } + + if (task.IsCompletedSuccessfully) { + _ = task.Result; + } + + if (task.IsCompleted || task.IsCanceled) { + _ = task.Result; + } + + if (task.Status == TaskStatus.RanToCompletion) { + _ = task.Result; + } + + if (task.IsCompleted && task.Result == 1) { + } + } + + async void ConditionalAwait(bool condition) { + var task = Task.Run(() => 1); + if (condition) { + await task; + } + + _ = task.[|Result|]; + } + + async void ReassignedAfterAwait() { + var task = Task.Run(() => 1); + await task; + task = Task.Run(() => 2); + _ = task.[|Result|]; + } + + async void ReassignedLaterInAwaitStatement() { + var task = Task.Run(() => 1); + Consume(await task, task = Task.Run(() => 2)); + _ = task.[|Result|]; + } + + async void ReassignedInWhenAllArgument() { + var task = Task.Run(() => 1); + await Task.WhenAll(task, Replace(ref task)); + _ = task.[|Result|]; + } + + async void ReassignedInConfigureAwaitArgument() { + var task = Task.Run(() => 1); + await task.ConfigureAwait(ReplaceFlag(ref task)); + _ = task.[|Result|]; + } + + async void ReassignedInGuardCondition() { + var task = Task.Run(() => 1); + if (!task.IsCompleted || (task = Replace(ref task)) == null) { + await task; + } + + _ = task.[|Result|]; + } + + async void ReassignedByDeconstruction() { + var task = Task.Run(() => 1); + await task; + (task, _) = (Task.Run(() => 2), 0); + _ = task.[|Result|]; + } + + async void ReassignedByInvokedClosure() { + var task = Task.Run(() => 1); + await task; + Action replace = () => task = Task.Run(() => 2); + replace(); + _ = task.[|Result|]; + } + + async void ReassignedByClosureDeclaredBeforeAwait() { + var task = Task.Run(() => 1); + Action replace = () => task = Task.Run(() => 2); + await task; + replace(); + _ = task.[|Result|]; + } + + void ReboundRefAliasDoesNotConnectTargets(Task task, Task other) { + ref Task alias = ref task; + alias = ref other; + if (other.IsCompleted) { + _ = task.[|Result|]; + } + } + + void ConditionallyReboundRefAliasStillInvalidatesOriginal(Task task, Task other, bool condition) { + ref Task alias = ref task; + if (condition) { + alias = ref other; + } + + if (task.IsCompleted) { + alias = Task.Run(() => 3); + _ = task.[|Result|]; + } + } + + int ReassignedByClosureInGetter { + get { + var task = Task.Run(() => 1); + Action replace = () => task = Task.Run(() => 2); + if (task.IsCompleted) { + replace(); + return task.[|Result|]; + } + + return 0; + } + } + + void LaterLambdaDoesNotInvalidateEarlierGuard(Task task) { + if (task.IsCompleted) { + _ = task.Result; + } + + Action replace = () => task = Task.Run(() => 4); + } + + void LaterLocalFunctionStillInvalidatesEarlierGuard(Task task) { + if (task.IsCompleted) { + Replace(); + _ = task.[|Result|]; + } + + void Replace() { + Action nestedReplace = () => task = Task.Run(() => 5); + nestedReplace(); + } + } + + async void AwaitInDoWhileConditionIsNotDefinite() { + var task = Task.Run(() => true); + do { + break; + } while (await task); + + _ = task.[|Result|]; + } + + async void AwaitInSwitchGuardIsNotDefinite(int value) { + var task = Task.Run(() => true); + switch (value) { + case 0 when await task: + break; + default: + break; + } + + _ = task.[|Result|]; + } + + async void ReassignedEarlierInResultStatement() { + var task = Task.Run(() => 1); + await task; + Consume(task = Task.Run(() => 2), task.[|Result|]); + } + + void ReassignedInsideGuard(Task task) { + if (task.IsCompleted) { + task = Task.Run(() => 2); + _ = task.[|Result|]; + } + } + + void ReassignedThroughRefAliasInsideGuard(Task task) { + ref Task alias = ref task; + if (task.IsCompleted) { + alias = Task.Run(() => 2); + _ = task.[|Result|]; + } + } + + void ReassignedThroughReboundRefAliasInsideGuard(ref Task task, ref Task other) { + ref Task alias = ref other; + alias = ref task; + if (task.IsCompleted) { + alias = Task.Run(() => 2); + _ = task.[|Result|]; + } + } + + void RefAliasCompletionGuard(Task task) { + ref Task alias = ref task; + if (alias.IsCompleted) { + _ = task.Result; + } + } + + void RefAliasReceiverReassignedThroughOriginal(Task task) { + ref Task alias = ref task; + if (alias.IsCompleted) { + task = Task.Run(() => 2); + _ = alias.[|Result|]; + } + } + + void CapturedTaskIsNotProvenComplete() { + var task = Task.Run(() => 1); + Local(); + task = Task.Run(() => 2); + + async void Local() { + await task; + _ = task.[|Result|]; + } + } + + void Consume(int value) { } + void Consume(int first, int second) { } + void Consume(int value, Task task) { } + void Consume(Task task, int value) { } + Task Replace(ref Task task) => task = Task.Run(() => 2); + bool ReplaceFlag(ref Task task) { + task = Task.Run(() => 2); + return false; + } +} +"; + + await new CSVerify.Test + { + TestCode = test, + ReferenceAssemblies = Microsoft.CodeAnalysis.Testing.ReferenceAssemblies.Net.Net80, + }.RunAsync(); + } + + [Fact] + public async Task CapturedTaskReassignmentInTopLevelStatementsReportsWarning() + { + var test = @" +using System; +using System.Threading.Tasks; + +var task = Task.Run(() => 1); +Action replace = () => task = Task.Run(() => 2); +if (task.IsCompleted) { + replace(); + _ = task.[|Result|]; +} +"; + + await new CSVerify.Test + { + TestCode = test, + TestState = + { + OutputKind = OutputKind.ConsoleApplication, + }, + }.RunAsync(); + } + + [Fact] + public async Task ConfiguredSyncBlockingMethodsReportWithoutCodeFix() + { + var test = @" +using System.Threading.Tasks; +using Contoso.Threading; + +namespace Contoso.Threading { + static class TaskExtensions { + internal static T WaitSynchronously(this Task task) => default; + } + + class CustomWaiter { + internal void Join() { } + } +} + +class Test { + void F(Task task, Contoso.Threading.CustomWaiter waiter) { + _ = task.[|WaitSynchronously|](); + waiter.[|Join|](); + waiter?.[|Join|](); + } + + Task FAsync(Task task) { + _ = task.[|WaitSynchronously|](); + return task; + } +} +"; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + FixedCode = test, + }; + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncBlockingMethods.txt", @" +[Contoso.Threading.TaskExtensions]::WaitSynchronously +[Contoso.Threading.CustomWaiter]::Join +")); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task ConditionalTaskWaitInSynchronousMethodReports() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + void F(Task task) + { + task?.[|Wait|](); + } + } + """; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task ConfiguredSyncBlockingMethodReportsWithoutTaskType() + { + string test = """ + class CustomWaiter + { + internal void Join() { } + } + + class Test + { + void F(CustomWaiter waiter) + { + waiter.[|Join|](); + } + } + """; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + FixedCode = test, + ReferenceAssemblies = Microsoft.CodeAnalysis.Testing.ReferenceAssemblies.NetFramework.Net20.Default, + }; + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncBlockingMethods.txt", "[CustomWaiter]::Join")); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task ConfiguredSyncBlockingMethodRequiresApplicableAsyncAlternative() + { + string test = """ + using System.Threading.Tasks; + + class CustomWaiter + { + internal void Join(int value) { } + + internal Task JoinAsync(string required, int value) => Task.CompletedTask; + } + + class ReorderedWaiter + { + internal void Join(int first, string second) { } + + internal Task JoinAsync(string second, int first) => Task.CompletedTask; + } + + class OptionalWaiter + { + internal void Join(int value) { } + + internal Task JoinAsync(int value, string optional = null) => Task.CompletedTask; + } + + class Test + { + Task FAsync(CustomWaiter waiter, ReorderedWaiter reordered, OptionalWaiter optional) + { + waiter.[|Join|](1); + reordered.[|Join|](1, ""); + optional.Join(1); + return Task.CompletedTask; + } + } + """; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + FixedCode = test, + }; + verifyTest.TestState.AdditionalFiles.Add( + ("vs-threading.SyncBlockingMethods.txt", """ + [CustomWaiter]::Join + [ReorderedWaiter]::Join + [OptionalWaiter]::Join + """)); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task ConfiguredGenericExtensionReceiverDoesNotThrow() + { + var test = @" +using System.Threading.Tasks; + +class Test { + Task FAsync(object value) { + value.[|WaitSynchronously|](); + return Task.CompletedTask; + } +} + +static class Extensions { + public static void WaitSynchronously(this T value) { } +} +"; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + FixedCode = test, + }; + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncBlockingMethods.txt", "[Extensions]::WaitSynchronously")); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task ConfiguredAsyncSuffixedMethodIsNotTreatedAsCoveredByVSTHRD103() + { + var test = @" +using System.Threading.Tasks; + +class CustomWaiter { + internal void JoinAsync() { } + internal Task JoinAsyncAsync() => Task.CompletedTask; +} + +class Test { + async Task FAsync(CustomWaiter waiter) { + waiter.[|JoinAsync|](); + await Task.Yield(); + } +} +"; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + FixedCode = test, + }; + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncBlockingMethods.txt", "[CustomWaiter]::JoinAsync")); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task ConfiguredSyncBlockingMethodInsideNameOfDoesNotReport() + { + var test = @" +namespace Contoso.Threading { + class CustomWaiter { + internal void Join() { } + } +} + +class Test { + void F(Contoso.Threading.CustomWaiter waiter) { + _ = nameof({|CS8081:waiter.Join()|}); + } +} +"; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + }; + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncBlockingMethods.txt", "[Contoso.Threading.CustomWaiter]::Join")); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task KnownAwaiterFromCustomMethodDoesNotOfferCodeFix() + { + var test = @" +using System.Runtime.CompilerServices; +using System.Threading.Tasks; + +class Test { + void F(Task task) { + GetCustomAwaiter(task).[|GetResult|](); + } + + TaskAwaiter GetCustomAwaiter(Task task) => task.GetAwaiter(); +} +"; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task ParameterizedGetAwaiterDoesNotOfferCodeFix() + { + var test = @" +using System.Runtime.CompilerServices; +using System.Threading.Tasks; + +class Test { + void F(CustomAwaitable value) { + value.GetAwaiter(1).[|GetResult|](); + } +} + +class CustomAwaitable { + public TaskAwaiter GetAwaiter(int value) => Task.CompletedTask.GetAwaiter(); +} +"; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task ParameterizedTaskGetAwaiterDoesNotUseCompletionProofOrOfferCodeFix() + { + var test = @" +using System.Runtime.CompilerServices; +using System.Threading.Tasks; + +class Test { + void F(Task task) { + if (task.IsCompleted) { + _ = task.GetAwaiter(1).[|GetResult|](); + } + + task.ContinueWith(t => t.GetAwaiter(1).[|GetResult|]()); + } +} + +static class TaskExtensions { + public static TaskAwaiter GetAwaiter(this Task task, int mode) + => Task.Run(() => mode).GetAwaiter(); +} +"; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task ConfiguredAwaitExtensionDoesNotProveOriginalTaskCompleted() + { + var test = @" +using System.Threading.Tasks; + +class Test { + async void F(Task task) { + await task.ConfigureAwait(""custom""); + _ = task.[|Result|]; + } +} + +static class TaskExtensions { + public static Task ConfigureAwait(this Task task, string mode) + => Task.Run(() => mode.Length); +} +"; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task TaskWaitExtensionDoesNotOfferCodeFix() + { + var test = @" +using System.Threading.Tasks; + +class Test { + void F(Task task) { + task.[|Wait|](""custom""); + } +} + +static class TaskExtensions { + public static void Wait(this Task task, string mode) { } +} +"; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + FixedCode = test, + }; + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncBlockingMethods.txt", "[TaskExtensions]::Wait")); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task CodeFixIsNotOfferedOutsideMethodDeclarations() + { + var test = @" +using System.Threading.Tasks; + +class Test { + Test() { + Task.Delay(1).[|Wait|](); + } + + int Value { + get { + return Task.FromResult(1).[|Result|]; + } + } +} +"; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task CodeFixIsNotOfferedForMethodsThatCannotBeAsync() + { + var test = @" +using System; +using System.Collections.Generic; +using System.Threading.Tasks; + +ref struct RefLike { + internal int Value; +} + +partial class Test { + int RefParameter(Task task, ref int value) { + return task.[|Result|]; + } + + int OutParameter(Task task, out int value) { + value = 0; + return task.[|Result|]; + } + + IEnumerable Iterator(Task task) { + yield return task.[|Result|]; + } + + int RefLikeParameter(Task task, RefLike value) { + return task.[|Result|]; + } + + RefLike RefLikeReturn(Task task) { + return new RefLike { Value = task.[|Result|] }; + } + + int RefLocal(Task task) { + int value = 0; + ref int alias = ref value; + return task.[|Result|]; + } + + int RefLikeLocal(Task task) { + RefLike value = default; + return task.[|Result|] + value.Value; + } + + int OutDeclaredRefLikeLocal(Task task) { + Create(out RefLike value); + return task.[|Result|] + value.Value; + } + + private partial int PartialMethod(Task task); + + private partial int PartialMethod(Task task) { + return task.[|Result|]; + } + + void MethodGroup(Task task) { + task.[|Wait|](); + } + + void UseMethodGroup() { + Action action = MethodGroup; + } + + static void Create(out RefLike value) { + value = default; + } +} +"; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task CodeFixIsNotOfferedInAwaitForbiddenContexts() + { + var test = @" +using System; +using System.Threading.Tasks; + +class Test { + void F(object gate) { + lock (gate) { + _ = Task.FromResult(1).[|Result|]; + } + + try { + } catch (Exception) when (Task.FromResult(false).[|Result|]) { + } + } +} +"; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task CodeFixIsNotOfferedInUnsafeOrFixedContexts() + { + string test = """ + using System.Threading.Tasks; + + unsafe class Test + { + private int[] values = new int[1]; + + int UnsafeType(Task task) + { + return task.[|Result|]; + } + + int UnsafeBlock(Task task) + { + unsafe + { + return task.[|Result|]; + } + } + + int FixedBlock(Task task) + { + fixed (int* pointer = values) + { + return task.[|Result|]; + } + } + } + """; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + FixedCode = test, + }; + await verifyTest.RunAsync(); + } + + [Fact] + public async Task CodeFixIsNotOfferedWhenCallerCannotBecomeAsync() + { + string test = """ + using System.Threading.Tasks; + + ref struct RefLike + { + internal int Value; + } + + class Test + { + Test(Task task) + { + _ = CalledByConstructor(task); + } + + static int CalledByConstructor(Task task) + { + return task.[|Result|]; + } + + static int CalledByRefLikeLocal(Task task) + { + return task.[|Result|]; + } + + static int CallerWithRefLikeLocal(Task task) + { + RefLike value = default; + return CalledByRefLikeLocal(task) + value.Value; + } + } + """; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task CodeFixIsNotOfferedWhenTaskCallerWouldLoseReturnedInvocation() + { + string test = """ + using System; + using System.Threading.Tasks; + + class DerivedTask : Task + { + internal DerivedTask() + : base(() => { }) + { + } + } + + class Test + { + static DerivedTask GetTask(Task task) + { + task.[|Wait|](); + return new DerivedTask(); + } + + static Task Caller(Task task) + { + return GetTask(task); + } + } + """; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task CodeFixParenthesizesAwaitUsedAsReceiver() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + int GetLength(Task task) + { + return task.[|Result|].Length; + } + } + """; + string withFix = """ + using System.Threading.Tasks; + + class Test + { + async Task GetLengthAsync(Task task) + { + return (await task).Length; + } + } + """; + + await CSVerify.VerifyCodeFixAsync(test, withFix); + } + + [Fact] + public async Task CodeFixIsNotOfferedWhenGenericTaskLikeCallerWouldLoseReturnedInvocation() + { + string test = """ + using System.Threading.Tasks; + + class DerivedTask : Task + { + internal DerivedTask() + : base(() => default) + { + } + } + + class Test + { + static DerivedTask GetValue(Task task) + { + _ = task.[|Result|]; + return new DerivedTask(); + } + + static Task Caller(Task task) => GetValue(task); + } + """; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task CodeFixIsNotOfferedForDirectReturnFromAsyncCompatibleCaller() + { + string test = """ + using System; + using System.Runtime.CompilerServices; + using System.Threading.Tasks; + + [AsyncMethodBuilder(typeof(CustomTaskMethodBuilder<>))] + class CustomTask + { + public TaskAwaiter GetAwaiter() => Task.FromResult(default(T)).GetAwaiter(); + public static implicit operator CustomTask(T value) => new(); + } + + struct CustomTaskMethodBuilder + { + public static CustomTaskMethodBuilder Create() => default; + public CustomTask Task => new(); + public void SetResult(T result) { } + public void SetException(Exception exception) { } + public void SetStateMachine(IAsyncStateMachine stateMachine) { } + public void Start(ref TStateMachine stateMachine) + where TStateMachine : IAsyncStateMachine => stateMachine.MoveNext(); + public void AwaitOnCompleted(ref TAwaiter awaiter, ref TStateMachine stateMachine) + where TAwaiter : INotifyCompletion + where TStateMachine : IAsyncStateMachine { } + public void AwaitUnsafeOnCompleted(ref TAwaiter awaiter, ref TStateMachine stateMachine) + where TAwaiter : ICriticalNotifyCompletion + where TStateMachine : IAsyncStateMachine { } + } + + class Test + { + static int GetValue(Task task) + { + return task.[|Result|]; + } + + static CustomTask Caller(Task task) + { + return GetValue(task); + } + } + """; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task CodeFixIsNotOfferedWhenCallerUsesConditionalAccess() + { + string test = """ + using System.Threading.Tasks; + + class Receiver + { + internal int GetValue(Task task) + { + task.[|Wait|](); + return 1; + } + } + + class Test + { + static int? Caller(Receiver receiver, Task task) => receiver?.GetValue(task); + } + """; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task CodeFixIsNotOfferedWhenAsyncNameHasSameSignature() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + int GetValue(Task task) + { + return task.[|Result|]; + } + + Task GetValueAsync(Task task) + { + return task; + } + } + """; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task CodeFixIsNotOfferedWhenAsyncNameHasApplicableOptionalOverload() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + int GetValue(Task task) + { + return task.[|Result|]; + } + + Task GetValueAsync(Task task, bool optional = false) + { + return task; + } + } + """; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task CodeFixIsNotOfferedWhenChangingMethodContract() + { + var test = @" +using System.Threading.Tasks; + +interface ITest { + int InterfaceMethod(Task task); +} + +abstract class Base { + public abstract int OverrideMethod(Task task); +} + +class Test : Base, ITest { + public virtual int VirtualMethod(Task task) { + return task.[|Result|]; + } + + public override int OverrideMethod(Task task) { + return task.[|Result|]; + } + + public int InterfaceMethod(Task task) { + return task.[|Result|]; + } +} +"; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task CodeFixIsNotOfferedForDefaultInterfaceMethod() + { + string test = """ + using System.Threading.Tasks; + + interface ITest + { + int DefaultInterfaceMethod(Task task) + { + return task.[|Result|]; + } + } + """; + + await new CSVerify.Test + { + TestCode = test, + FixedCode = test, + ReferenceAssemblies = Microsoft.CodeAnalysis.Testing.ReferenceAssemblies.Net.Net80, + }.RunAsync(); + } + + [Fact] + public async Task StaticGetAwaiterFactoryDoesNotOfferCodeFix() + { + var test = @" +using System.Runtime.CompilerServices; +using System.Threading.Tasks; + +class Test { + void F() { + AwaiterFactory.GetAwaiter().[|GetResult|](); + } +} + +static class AwaiterFactory { + public static TaskAwaiter GetAwaiter() => Task.CompletedTask.GetAwaiter(); +} +"; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task CodeFixIsNotOfferedWhenBlockingExpressionHasCompileErrors() + { + var test = @" +using System.Threading.Tasks; + +class Test { + void F() { + Task.Delay(1, CancellationToken.None).GetAwaiter().[|GetResult|](); + } +} +"; + + DiagnosticResult compilerError = DiagnosticResult.CompilerError("CS0103").WithSpan(6, 23, 6, 40).WithArguments("CancellationToken"); + await CSVerify.VerifyCodeFixAsync(test, new[] { compilerError }, test); + } + + [Fact] + public async Task CodeFixIsNotOfferedForWaitWithCancellation() + { + var test = @" +using System.Threading; +using System.Threading.Tasks; + +class Test { + void F(CancellationToken cancellationToken) { + Task.Delay(2, cancellationToken).[|Wait|](cancellationToken); + } +} +"; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task Task_GetAwaiter_GetResult_ShouldReportWarning() + { + var test = @" +using System; +using System.Threading.Tasks; + +class Test { + void F() { + var task = Task.Run(() => 1); + task.GetAwaiter().GetResult(); + } +} +"; + var withFix = @" +using System; +using System.Threading.Tasks; + +class Test { + async Task FAsync() { + var task = Task.Run(() => 1); + await task; + } +} +"; + DiagnosticResult expected = CSVerify.Diagnostic().WithSpan(8, 27, 8, 36); + await CSVerify.VerifyCodeFixAsync(test, expected, withFix); + } + + [Fact] + public async Task ParenthesizedTask_GetAwaiter_GetResult_ShouldReportWarning() + { + var test = @" +using System.Threading.Tasks; + +class Test { + void F() { + var task = Task.Run(() => 1); + (task.GetAwaiter()).[|GetResult|](); + } +} +"; + var withFix = @" +using System.Threading.Tasks; + +class Test { + async Task FAsync() { + var task = Task.Run(() => 1); + await task; + } +} +"; + + await CSVerify.VerifyCodeFixAsync(test, withFix); + } + + [Fact] + public async Task ConfiguredTask_GetAwaiter_GetResult_ShouldReportWarning() + { + var test = @" +using System; +using System.Threading.Tasks; + +class Test { + void F() { + var task = Task.Run(() => 1); + task.ConfigureAwait(false).GetAwaiter().[|GetResult|](); + } +} +"; + var withFix = @" +using System; +using System.Threading.Tasks; + +class Test { + async Task FAsync() { + var task = Task.Run(() => 1); + await task.ConfigureAwait(false); + } +} +"; + await CSVerify.VerifyCodeFixAsync(test, withFix); + } + + [Fact] + public async Task ValueTask_GetAwaiter_GetResult_ShouldReportWarning() + { + var test = @" +using System; +using System.Threading.Tasks; + +class Test { + void F() { + ValueTask task = default; + task.GetAwaiter().GetResult(); + } +} +"; + var withFix = @" +using System; +using System.Threading.Tasks; + +class Test { + async Task FAsync() { + ValueTask task = default; + await task; + } +} +"; + DiagnosticResult expected = CSVerify.Diagnostic().WithSpan(8, 27, 8, 36); + await CSVerify.VerifyCodeFixAsync(test, expected, withFix); + } + + [Fact] + public async Task ConfiguredValueTask_GetAwaiter_GetResult_ShouldReportWarning() + { + var test = @" +using System; +using System.Threading.Tasks; + +class Test { + void F() { + ValueTask task = default; + task.ConfigureAwait(false).GetAwaiter().[|GetResult|](); + } +} +"; + var withFix = @" +using System; +using System.Threading.Tasks; + +class Test { + async Task FAsync() { + ValueTask task = default; + await task.ConfigureAwait(false); + } +} +"; + await CSVerify.VerifyCodeFixAsync(test, withFix); + } + + [Fact] + public async Task TaskResult_FixUpdatesCallers() + { + var test = new SourceFileList("Test", "cs") + { + @" +using System; +using System.Threading.Tasks; + +class Test { + internal static int GetNumber(int a) { + var task = Task.Run(() => a); + return task.Result; + } + + int Add(int a, int b) { + return GetNumber(a) + b; + } + + int Subtract(int a, int b) { + return GetNumber(a) - b; + } + + static int Main(string[] args) + { + return new Test().Add(1, 2); + } +} +", + @" +class TestClient { + int Multiply(int a, int b) { + return Test.GetNumber(a) * b; + } +} +", + }; + var withFix = new SourceFileList("Test", "cs") + { + @" +using System; +using System.Threading.Tasks; + +class Test { + internal static async Task GetNumberAsync(int a) { + var task = Task.Run(() => a); + return await task; + } + + async Task AddAsync(int a, int b) { + return await GetNumberAsync(a) + b; + } + + async Task SubtractAsync(int a, int b) { + return await GetNumberAsync(a) - b; + } + + static async Task Main(string[] args) + { + return await new Test().AddAsync(1, 2); + } +} +", + @" +class TestClient { + async System.Threading.Tasks.Task MultiplyAsync(int a, int b) { + return await Test.GetNumberAsync(a) * b; + } +} +", + }; + + var verifyTest = new CSVerify.Test + { + TestState = + { + OutputKind = OutputKind.ConsoleApplication, + }, + ExpectedDiagnostics = + { + CSVerify.Diagnostic().WithSpan("Test0.cs", 8, 21, 8, 27), + }, + }; + + verifyTest.TestState.Sources.AddRange(test); + verifyTest.FixedState.Sources.AddRange(withFix); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task DoNotReportWarningInTaskReturningMethods() + { + var test = @" +using System.Threading.Tasks; + +class Test { + Task F() { + var task = Task.Run(() => 1); + task.GetAwaiter().GetResult(); + return Task.CompletedTask; + } +} +"; + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task DoNotReportWarningOnCodeGeneratedByXaml2CS() + { + var test = @" //------------------------------------------------------------------------------ // // This code was generated by a tool. @@ -833,6 +2432,35 @@ void F() { await CSVerify.VerifyAnalyzerAsync(test); } + [Fact] + public async Task DoNotReportConfiguredWarningOnCodeGeneratedByXaml2CS() + { + var test = @" +//------------------------------------------------------------------------------ +// +//------------------------------------------------------------------------------ + +namespace Contoso.Threading { + class CustomWaiter { + internal void Join() { } + } + + class Test { + void F(CustomWaiter waiter) { + waiter.Join(); + } + } +} +"; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + }; + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncBlockingMethods.txt", "[Contoso.Threading.CustomWaiter]::Join")); + await verifyTest.RunAsync(); + } + [Fact] public async Task DoNotReportWarningOnJTFRun() { diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD102AvoidJtfRunInNonPublicMembersAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD102AvoidJtfRunInNonPublicMembersAnalyzerTests.cs index cbe31a6c0..3fdb7fe2d 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD102AvoidJtfRunInNonPublicMembersAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD102AvoidJtfRunInNonPublicMembersAnalyzerTests.cs @@ -5,6 +5,24 @@ public class VSTHRD102AvoidJtfRunInNonPublicMembersAnalyzerTests { + [Fact] + public async Task JoinableTaskCompletionGuardDoesNotSuppressDiagnostic() + { + var test = @" +using Microsoft.VisualStudio.Threading; + +class Test { + void F(JoinableTask joinableTask) { + if (joinableTask.IsCompleted) { + joinableTask.[|Join|](); + } + } +} +"; + + await CSVerify.VerifyAnalyzerAsync(test); + } + [Fact] public async Task JtfRunInPublicMethodsOfInternalType_ProducesDiagnostic() { diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs index 83ed2c74e..f34335c4f 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs @@ -124,6 +124,143 @@ void Run() { } await CSVerify.VerifyCodeFixAsync(test, expected, withFix); } + [Fact] + public async Task AwaitedTaskResultDoesNotGenerateWarning() + { + var test = @" +using System.Threading.Tasks; + +class Test { + async Task T(Task task) { + await task; + _ = task.Result; + } +} +"; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task AwaitedTaskRemainsCompletedAcrossStatementWrappers() + { + var test = @" +using System.IO; +using System.Threading.Tasks; + +class Test { + async Task T(Task task, int value) { + await task; + try { + _ = task.Result; + } catch { + } + + using (new MemoryStream()) { + _ = task.Result; + } + + switch (value) { + case 0: + _ = task.Result; + break; + } + } +} +"; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task AwaitedTaskReassignedAfterAccessInLoopGeneratesWarning() + { + var test = @" +using System.Threading.Tasks; + +class Test { + async Task T(Task task, bool condition) { + await task; + while (condition) { + _ = task.{|#0:Result|}; + task = Task.FromResult(1); + } + } +} +"; + + DiagnosticResult expected = CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result"); + await CSVerify.VerifyAnalyzerAsync(test, expected); + } + + [Fact] + public async Task CompletedConditionalTaskResultDoesNotGenerateWarning() + { + var test = @" +using System.Threading.Tasks; + +class Test { + async Task Awaited(Task task) { + await task; + _ = task?.Result; + } + + Task Guarded(Task task) { + if (task.IsCompleted) { + _ = task?.Result; + } + + return Task.CompletedTask; + } +} +"; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task AwaitedValueTaskCompletionGuardStillGeneratesWarning() + { + var test = @" +using System.Threading.Tasks; + +class Test { + async Task T(ValueTask task) { + await task; + if (task.IsCompleted) { + _ = task.{|#0:Result|}; + } + } +} +"; + + DiagnosticResult expected = CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result"); + await CSVerify.VerifyAnalyzerAsync(test, expected); + } + + [Fact] + public async Task AwaitedTaskWaitDoesNotFallBackToAsyncAlternativeWarning() + { + var test = @" +using System.Threading; +using System.Threading.Tasks; + +class Test { + async Task T(Task task) { + await task; + task.Wait(); + } +} + +static class TaskExtensions { + internal static Task WaitAsync(this Task task, CancellationToken cancellationToken = default) + => task; +} +"; + + await CSVerify.VerifyAnalyzerAsync(test); + } + [Fact] public async Task JTFRunOfTInTaskReturningMethodGeneratesWarning() { @@ -790,6 +927,26 @@ Task T() { await CSVerify.VerifyAnalyzerAsync(test); } + [Fact] + public async Task TaskOfTResultInAsyncContinuationGeneratesNoWarning() + { + var test = @" +using System.Threading.Tasks; + +class Test { + Task T(Task task) { + task.ContinueWith(async t => { + _ = t.Result; + await Task.Yield(); + }); + return Task.CompletedTask; + } +} +"; + + await CSVerify.VerifyAnalyzerAsync(test); + } + [Fact] public async Task TaskResultGuardedByIsCompletedSuccessfully_GeneratesNoWarning() { @@ -971,7 +1128,7 @@ Task GetResultAsync(Task task, Task replacement) } [Fact] - public async Task ValueTaskResultGuardedByIsCompletedSuccessfully_GeneratesWarning() + public async Task ValueTaskResultGuardedByIsCompletedSuccessfully_GeneratesNoWarning() { string test = """ using System.Threading.Tasks; @@ -982,7 +1139,7 @@ Task GetResultAsync(ValueTask task) { if (task.IsCompletedSuccessfully) { - return Task.FromResult(task.{|#0:Result|}); + return Task.FromResult(task.Result); } return Task.FromResult(0); @@ -995,7 +1152,6 @@ Task GetResultAsync(ValueTask task) TestCode = test, ReferenceAssemblies = Microsoft.CodeAnalysis.Testing.ReferenceAssemblies.Net.Net80, }; - analyzerTest.ExpectedDiagnostics.Add(CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result")); await analyzerTest.RunAsync(); } @@ -1241,6 +1397,591 @@ Task GetResultAsync(ref Task task) }.RunAsync(); } + [Fact] + public async Task AliasedRefParameterResultGuardedByCompletionGeneratesWarning() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + Task GetResultAsync(ref Task task, ref Task possibleAlias) + { + if (task.IsCompleted) + { + possibleAlias = Task.FromResult(1); + return Task.FromResult(task.{|#0:Result|}); + } + + return task; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync( + test, + CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result")); + } + + [Fact] + public async Task SeparateRefParameterDoesNotInvalidateOrdinaryTaskCompletionGuard() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + Task GetResultAsync(Task task, ref Task other) + { + if (task.IsCompleted) + { + other = Task.FromResult(1); + return Task.FromResult(task.Result); + } + + return task; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task GotoCanBypassAwait_GeneratesWarning() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async Task GetResultAsync(Task task, bool skip) + { + if (skip) + { + goto AccessResult; + } + + await task; + + AccessResult: + return task.{|#0:Result|}; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync( + test, + CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result")); + } + + [Fact] + public async Task RefReturningAssignmentAfterAwait_GeneratesWarning() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async Task GetResultAsync(Task task, Task replacement) + { + await task; + GetReference(ref task) = replacement; + return task.{|#0:Result|}; + } + + static ref Task GetReference(ref Task task) => ref task; + } + """; + + await CSVerify.VerifyAnalyzerAsync( + test, + CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result")); + } + + [Fact] + public async Task AwaitInConditionalAccessReceiverCompletesTask() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async Task GetResultAsync(Task task) + { + _ = (await task)?.Length; + return task.Result; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task AwaitInSwitchGoverningExpressionCompletesTask() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async Task GetResultAsync(Task task) + { + _ = (await task) switch { _ => 0 }; + return task.Result; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task MutationAfterNestedCompletionProof_GeneratesWarning() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + Task GetResultAsync(Task task, Task replacement) + { + if ((task.IsCompleted && (task = replacement) != null) && task.{|#0:Result|} > 0) + { + return task; + } + + return replacement; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync( + test, + CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result")); + } + + [Fact] + public async Task AwaitedValueTaskCopyInvalidatesCompletionGuard() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async Task GetResultAsync(ValueTask task) + { + ValueTask copy = task; + await copy; + if (task.IsCompletedSuccessfully) + { + return task.{|#0:Result|}; + } + + return 0; + } + } + """; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + ReferenceAssemblies = Microsoft.CodeAnalysis.Testing.ReferenceAssemblies.Net.Net80, + }; + verifyTest.ExpectedDiagnostics.Add(CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result")); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task OverwrittenValueTaskCopyDoesNotInvalidateCompletionGuard() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async Task GetResultAsync(ValueTask task, ValueTask other) + { + ValueTask copy = task; + { + copy = other; + } + + await copy; + if (task.IsCompletedSuccessfully) + { + return task.Result; + } + + return 0; + } + } + """; + + await new CSVerify.Test + { + TestCode = test, + ReferenceAssemblies = Microsoft.CodeAnalysis.Testing.ReferenceAssemblies.Net.Net80, + }.RunAsync(); + } + + [Fact] + public async Task AwaitedConfiguredValueTaskCopyInvalidatesCompletionGuard() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async Task GetResultAsync(ValueTask task) + { + var configured = task.ConfigureAwait(false); + await configured; + if (task.IsCompletedSuccessfully) + { + return task.{|#0:Result|}; + } + + return 0; + } + } + """; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + ReferenceAssemblies = Microsoft.CodeAnalysis.Testing.ReferenceAssemblies.Net.Net80, + }; + verifyTest.ExpectedDiagnostics.Add(CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result")); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task AwaitedValueTaskParameterCopyInvalidatesCompletionGuard() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async Task GetResultAsync(ValueTask task, ValueTask copy) + { + copy = task; + await copy; + if (task.IsCompletedSuccessfully) + { + return task.{|#0:Result|}; + } + + return 0; + } + } + """; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + ReferenceAssemblies = Microsoft.CodeAnalysis.Testing.ReferenceAssemblies.Net.Net80, + }; + verifyTest.ExpectedDiagnostics.Add(CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result")); + await verifyTest.RunAsync(); + } + + [Fact] + public async Task CustomAwaitableAliasInvalidatesValueTaskCompletionGuard() + { + string test = """ + using System.Threading.Tasks; + + static class Extensions + { + internal static ValueTask Preserve(this ValueTask task) => task; + } + + class Test + { + async Task GetResultAsync(ValueTask task) + { + await Extensions.Preserve(task); + if (task.IsCompletedSuccessfully) + { + return task.{|#0:Result|}; + } + + return 0; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync( + test, + CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result")); + } + + [Fact] + public async Task UnevaluatedValueTaskMembersDoNotInvalidateCompletionGuard() + { + string test = """ + using System; + using System.Threading.Tasks; + + class Test + { + int GetResult(ValueTask task) + { + _ = nameof(task.Result); + Func getResult = task.GetAwaiter().GetResult; + if (task.IsCompletedSuccessfully) + { + return task.Result; + } + + return 0; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task InvokedNestedFunctionInvalidatesValueTaskCompletionGuard() + { + string test = """ + using System; + using System.Threading.Tasks; + + class Test + { + async Task GetResultAsync(ValueTask task) + { + ValueTask copy = task; + async Task ConsumeAsync() + { + await copy; + } + + Func consume = async () => + { + _ = task.{|#0:Result|}; + await Task.Yield(); + }; + await ConsumeAsync(); + await consume(); + if (task.IsCompletedSuccessfully) + { + return task.{|#1:Result|}; + } + + return 0; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync( + test, + CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result"), + CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(1).WithArguments("Result")); + } + + [Fact] + public async Task RefReturningExtensionReceiverInvalidatesCompletionGuard() + { + string test = """ + using System.Threading.Tasks; + + static class Extensions + { + internal static ref Task AsRef(this int ignored, ref Task task) => ref task; + } + + class Test + { + Task GetResultAsync(Task task) + { + if (task.IsCompleted) + { + ref Task alias = ref 0.AsRef(ref task); + alias = Task.FromResult(1); + return Task.FromResult(task.{|#0:Result|}); + } + + return task; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync( + test, + CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result")); + } + + [Fact] + public async Task CoalesceAssignmentAwaitDoesNotProveCompletion() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async Task GetResultAsync(Task task, int? other) + { + other ??= await task; + return task.{|#0:Result|}; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync( + test, + CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result")); + } + + [Fact] + public async Task RefConditionalAliasMutationAfterAwait_GeneratesWarning() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + Task GetResultAsync(Task task, Task other, Task replacement, bool chooseTask) + { + ref Task alias = ref (chooseTask ? ref task : ref other); + if (task.IsCompleted) + { + alias = replacement; + return Task.FromResult(task.{|#0:Result|}); + } + + return task; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync( + test, + CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result")); + } + + [Fact] + public async Task AwaitBeforeResultInExpressionBodyCompletesTask() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async Task GetResultAsync(Task task) => (await task) + task.Result; + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task CompletedConditionalWaitDoesNotFallBackToAsyncAlternativeAnalysis() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + async Task WaitAsync(Task task) + { + await task; + task?.Wait(); + } + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task ConditionalWaitOnIncompleteTaskGeneratesWarning() + { + string test = """ + using System.Threading.Tasks; + + class Test + { + Task WaitAsync(Task task) + { + task?.{|#0:Wait|}(); + return Task.CompletedTask; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync( + test, + CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Wait")); + } + + [Fact] + public async Task AsyncAlternativeMustBeApplicableToInvocation() + { + string test = """ + using System.Threading.Tasks; + + class CustomWaiter + { + internal void Join(int value) { } + + internal Task JoinAsync(string required, int value) => Task.CompletedTask; + } + + class ReorderedWaiter + { + internal void Join(int first, string second) { } + + internal Task JoinAsync(string second, int first) => Task.CompletedTask; + } + + class Test + { + Task FAsync(CustomWaiter waiter, ReorderedWaiter reordered) + { + waiter.Join(1); + reordered.Join(1, ""); + return Task.CompletedTask; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task AsyncAlternativeMayHaveOptionalAdditionalParameter() + { + string test = """ + using System.Threading.Tasks; + + class CustomWaiter + { + internal void Join(int value) { } + + internal Task JoinAsync(int value, string optional = null) => Task.CompletedTask; + } + + class Test + { + Task FAsync(CustomWaiter waiter) + { + waiter.{|#0:Join|}(1); + return Task.CompletedTask; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync( + test, + CSVerify.Diagnostic(Descriptor).WithLocation(0).WithArguments("Join", "JoinAsync")); + } + [Fact] public async Task TaskGetAwaiterGetResultInTaskReturningMethodGeneratesWarning() {