From b5d41badada5987dbeded365f4c3464bb9351678 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 10:48:30 -0600 Subject: [PATCH 01/26] Fix VSTHRD002 completion analysis and extensibility Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- docfx/analyzers/VSTHRD002.md | 8 + docfx/analyzers/configuration.md | 13 + .../CSharpCommonInterest.cs | 339 ++++++++++++------ .../VSTHRD002UseJtfRunAnalyzer.cs | 146 ++++++-- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 76 +++- .../MultiAnalyzerTests.cs | 3 +- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 258 +++++++++++++ 7 files changed, 699 insertions(+), 144 deletions(-) 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 8cf653fd8..cf816802e 100644 --- a/docfx/analyzers/configuration.md +++ b/docfx/analyzers/configuration.md @@ -104,3 +104,16 @@ excluded from VSTHRD103 analysis by specifying them in a configuration file. **Sample:** `[System.Data.SqlClient.SqlDataReader]::Read` **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` diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index 497469f95..92692fcf8 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -4,8 +4,10 @@ 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; @@ -113,178 +115,307 @@ internal static void InspectMemberAccess( } } - private static SyntaxNode? GetEnclosingBlock(SyntaxNode? node) + private static ExpressionSyntax UnwrapParentheses(ExpressionSyntax expression) { - while (node is not null) + while (expression is ParenthesizedExpressionSyntax parenthesized) { - if (node.IsKind(SyntaxKind.Block)) - { - return node; - } + expression = parenthesized.Expression; + } + + return expression; + } - node = node.Parent; + private static ExpressionSyntax GetTaskReceiver(MemberAccessExpressionSyntax memberAccessSyntax) + { + ExpressionSyntax receiver = memberAccessSyntax.Expression; + if (receiver is InvocationExpressionSyntax getAwaiterInvocation + && getAwaiterInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: "GetAwaiter" } getAwaiterAccess) + { + receiver = getAwaiterAccess.Expression; } - return null; + if (receiver is InvocationExpressionSyntax configureAwaitInvocation + && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess) + { + receiver = configureAwaitAccess.Expression; + } + + return UnwrapParentheses(receiver); } - private static bool IsVariablePassedToInvocation(InvocationExpressionSyntax invocationExpr, string variableName, bool byRef) + private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) { - ArgumentListSyntax? argList = invocationExpr.ChildNodes().OfType().FirstOrDefault(); - if (argList is null) + ISymbol? taskSymbol = context.SemanticModel.GetSymbolInfo(GetTaskReceiver(memberAccessSyntax), context.CancellationToken).Symbol; + if (taskSymbol is not ILocalSymbol and not IParameterSymbol) { return false; } - foreach (ArgumentSyntax arg in argList.ChildNodes().OfType()) + if (IsWithinCompletedTaskBranch(context, memberAccessSyntax, taskSymbol)) { - // `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)) + return true; + } + + StatementSyntax? containingStatement = memberAccessSyntax.FirstAncestorOrSelf(); + if (containingStatement is null) + { + return false; + } + + while (containingStatement.Parent is BlockSyntax block) + { + if (MayReassignTask(context, containingStatement, taskSymbol, containingStatement.SpanStart - 1, memberAccessSyntax.SpanStart)) { - continue; + return false; + } + + int statementIndex = block.Statements.IndexOf(containingStatement); + for (int i = statementIndex - 1; i >= 0; i--) + { + StatementSyntax statement = block.Statements[i]; + if (StatementCompletesTask(context, statement, taskSymbol)) + { + return true; + } + + if (MayReassignTask(context, statement, taskSymbol)) + { + return false; + } } - IdentifierNameSyntax identiferName = arg.ChildNodes().OfType().FirstOrDefault(); - if (identiferName is null) + StatementSyntax? outerStatement = containingStatement.Ancestors().OfType().FirstOrDefault(statement => statement.Parent is BlockSyntax); + if (outerStatement is not IfStatementSyntax and not BlockSyntax) { return false; } - if (identiferName.Identifier.ValueText == variableName) + if (MayReassignTask(context, outerStatement, taskSymbol, outerStatement.SpanStart - 1, containingStatement.SpanStart)) { - return true; + return false; } + + containingStatement = outerStatement; } return false; } - private static bool IsTaskCompletedWithWhenAll(SyntaxNodeAnalysisContext context, InvocationExpressionSyntax invocationExpr, string taskVariableName) + private static bool IsWithinCompletedTaskBranch(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax, ISymbol taskSymbol) { - // We only care about awaited invocations, because an un-awaited Task.WhenAll will be an error. - if (invocationExpr.Parent is not AwaitExpressionSyntax) + foreach (IfStatementSyntax ifStatement in memberAccessSyntax.Ancestors().OfType()) { - return false; + IEnumerable nodesBetweenAccessAndCondition = memberAccessSyntax.Ancestors().TakeWhile(node => node != ifStatement); + if (nodesBetweenAccessAndCondition.Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax)) + { + continue; + } + + if (ifStatement.Statement.FullSpan.Contains(memberAccessSyntax.Span) + && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: true) + && !MayReassignTask(context, ifStatement, taskSymbol, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) + { + return true; + } + + if (ifStatement.Else?.Statement.FullSpan.Contains(memberAccessSyntax.Span) is true + && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: false) + && !MayReassignTask(context, ifStatement, taskSymbol, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) + { + return true; + } } - IEnumerable? memberAccessList = invocationExpr.ChildNodes().OfType(); - if (memberAccessList.Count() != 1) + return false; + } + + private static bool ConditionProvesCompletion(SyntaxNodeAnalysisContext context, ExpressionSyntax condition, ISymbol taskSymbol, bool conditionValue) + { + condition = UnwrapParentheses(condition); + if (condition is PrefixUnaryExpressionSyntax { RawKind: (int)SyntaxKind.LogicalNotExpression } logicalNot) { - return false; + return ConditionProvesCompletion(context, logicalNot.Operand, taskSymbol, !conditionValue); } - MemberAccessExpressionSyntax? memberAccess = memberAccessList.First(); + if (condition is BinaryExpressionSyntax binary) + { + if (conditionValue && binary.IsKind(SyntaxKind.LogicalAndExpression)) + { + return ConditionProvesCompletion(context, binary.Left, taskSymbol, conditionValue: true) + || ConditionProvesCompletion(context, binary.Right, taskSymbol, conditionValue: 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 (!conditionValue && binary.IsKind(SyntaxKind.LogicalOrExpression)) + { + return ConditionProvesCompletion(context, binary.Left, taskSymbol, conditionValue: false) + || ConditionProvesCompletion(context, binary.Right, taskSymbol, conditionValue: false); + } - if (!correctSyntax) + 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, taskSymbol)) + { + return equalityHolds; + } + } + + if (condition is MemberAccessExpressionSyntax completedProperty + && IsSameTask(context, completedProperty.Expression, taskSymbol) + && completedProperty.Name.Identifier.ValueText is nameof(Task.IsCompleted) + or nameof(Task.IsCanceled) + or nameof(Task.IsFaulted) + or "IsCompletedSuccessfully") { - return false; + return conditionValue; } - // 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 false; + } + + private static bool IsRanToCompletionComparison(SyntaxNodeAnalysisContext context, ExpressionSyntax left, ExpressionSyntax right, ISymbol taskSymbol) + { + return (IsTaskStatus(context, left, taskSymbol) && IsRanToCompletion(context, right)) + || (IsTaskStatus(context, right, taskSymbol) && IsRanToCompletion(context, left)); + } + + private static bool IsTaskStatus(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, ISymbol taskSymbol) + => expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.Status) } statusAccess + && IsSameTask(context, statusAccess.Expression, taskSymbol); + + 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, ISymbol taskSymbol) + { + if (TryGetAwaitExpression(statement, out AwaitExpressionSyntax? awaitExpression) + && AwaitCompletesTask(context, awaitExpression, taskSymbol) + && !MayReassignTask(context, statement, taskSymbol, awaitExpression.Span.End, statement.Span.End + 1)) { - return false; + return true; } - // Is the task variable passed as an argument to `Task.WhenAll`? - return IsVariablePassedToInvocation(invocationExpr, taskVariableName, byRef: false); + return statement is IfStatementSyntax { Else: null } ifStatement + && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: false) + && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbol); } - private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) + private static bool StatementDefinitelyAwaitsTask(SyntaxNodeAnalysisContext context, StatementSyntax statement, ISymbol taskSymbol) { - SyntaxNode? enclosingBlock = GetEnclosingBlock(memberAccessSyntax); - if (enclosingBlock is null) + if (TryGetAwaitExpression(statement, out AwaitExpressionSyntax? awaitExpression)) { - return false; + return AwaitCompletesTask(context, awaitExpression, taskSymbol) + && !MayReassignTask(context, statement, taskSymbol, awaitExpression.Span.End, statement.Span.End + 1); } - // 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) + if (statement is BlockSyntax block) { - if (parentExpr is IdentifierNameSyntax identifierExpr) + for (int i = block.Statements.Count - 1; i >= 0; i--) { - taskVariableName = identifierExpr.Identifier.ValueText; - break; - } - else if (parentExpr is MemberAccessExpressionSyntax memberAccessExpr) - { - parentExpr = memberAccessExpr.Expression; - } - else if (parentExpr is InvocationExpressionSyntax invocExpr) - { - parentExpr = invocExpr.Expression; + if (TryGetAwaitExpression(block.Statements[i], out awaitExpression) + && AwaitCompletesTask(context, awaitExpression, taskSymbol) + && !MayReassignTask(context, block.Statements[i], taskSymbol, awaitExpression.Span.End, block.Statements[i].Span.End + 1)) + { + return true; + } + + if (MayReassignTask(context, block.Statements[i], taskSymbol)) + { + return false; + } } - else + } + + return false; + } + + private static bool TryGetAwaitExpression(StatementSyntax statement, [NotNullWhen(true)] out AwaitExpressionSyntax? awaitExpression) + { + static bool DescendIntoChildren(SyntaxNode node) => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; + + foreach (AwaitExpressionSyntax candidate in statement.DescendantNodes(DescendIntoChildren).OfType().Reverse()) + { + IEnumerable ancestorsWithinStatement = candidate.Ancestors().TakeWhile(node => node != statement); + bool isConditionallyExecuted = ancestorsWithinStatement.Any( + node => node is StatementSyntax or ConditionalExpressionSyntax or SwitchExpressionSyntax or ConditionalAccessExpressionSyntax + || (node is BinaryExpressionSyntax binary + && (binary.IsKind(SyntaxKind.LogicalAndExpression) + || binary.IsKind(SyntaxKind.LogicalOrExpression) + || binary.IsKind(SyntaxKind.CoalesceExpression)))); + if (!isConditionallyExecuted) { - break; + awaitExpression = candidate; + return true; } } - if (taskVariableName is null) + awaitExpression = null; + return false; + } + + private static bool AwaitCompletesTask(SyntaxNodeAnalysisContext context, AwaitExpressionSyntax awaitExpression, ISymbol taskSymbol) + { + ExpressionSyntax awaitedExpression = UnwrapParentheses(awaitExpression.Expression); + if (awaitedExpression is InvocationExpressionSyntax configureAwaitInvocation + && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess) { - return false; + awaitedExpression = UnwrapParentheses(configureAwaitAccess.Expression); } - // 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; + if (IsSameTask(context, awaitedExpression, taskSymbol)) + { + return true; + } - if (!taskWhenAllInvocationList.Any()) + 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 false; + return whenAllInvocation.ArgumentList.Arguments.Any(argument => IsSameTask(context, argument.Expression, taskSymbol)); } - // 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) + return false; + } + + private static bool MayReassignTask(SyntaxNodeAnalysisContext context, SyntaxNode node, ISymbol taskSymbol) + => MayReassignTask(context, node, taskSymbol, node.SpanStart - 1, node.Span.End + 1); + + private static bool MayReassignTask(SyntaxNodeAnalysisContext context, SyntaxNode node, ISymbol taskSymbol, int afterPosition, int beforePosition) + { + static bool DescendIntoChildren(SyntaxNode node) => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; + + foreach (AssignmentExpressionSyntax assignment in node.DescendantNodes(DescendIntoChildren).OfType()) { - // 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; - - if (assignmentList.Any()) + if (assignment.SpanStart > afterPosition + && assignment.SpanStart < beforePosition + && SymbolEqualityComparer.Default.Equals(context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol, taskSymbol)) { - return false; + return true; } + } - // 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 !invocationList.Any(); + foreach (ArgumentSyntax argument in node.DescendantNodes(DescendIntoChildren).OfType()) + { + if (argument.SpanStart > afterPosition + && argument.SpanStart < beforePosition + && (argument.RefKindKeyword.IsKind(SyntaxKind.RefKeyword) || argument.RefKindKeyword.IsKind(SyntaxKind.OutKeyword)) + && IsSameTask(context, argument.Expression, taskSymbol)) + { + return true; + } } return false; } + + private static bool IsSameTask(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, ISymbol taskSymbol) + => SymbolEqualityComparer.Default.Equals( + context.SemanticModel.GetSymbolInfo(UnwrapParentheses(expression), context.CancellationToken).Symbol, + taskSymbol); } diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index 3d2ba5c1e..6fe7f3919 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,6 +62,10 @@ public override void Initialize(AnalysisContext context) context.RegisterCompilationStartAction(compilationContext => { INamedTypeSymbol? taskSymbol = compilationContext.Compilation.GetTypeByMetadataName(Types.Task.FullName); + ImmutableArray configuredSyncBlockingMethods = CommonInterest.ReadMethods( + compilationContext.Options, + new Regex(@"^vs-threading\.SyncBlockingMethods(\..*)?.txt$", RegexOptions.IgnoreCase | RegexOptions.Singleline), + compilationContext.CancellationToken).ToImmutableArray(); if (taskSymbol is object) { compilationContext.RegisterCodeBlockStartAction(codeBlockContext => @@ -70,7 +75,7 @@ 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 => AnalyzeInvocation(c, taskSymbol, configuredSyncBlockingMethods, analyzeWholeCodeBlock)), SyntaxKind.InvocationExpression); codeBlockContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(c => AnalyzeMemberAccess(c, taskSymbol, analyzeWholeCodeBlock)), SyntaxKind.SimpleMemberAccessExpression); } }); @@ -124,51 +129,77 @@ private static void InspectMemberAccess( 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) + // A continuation's antecedent is complete throughout its delegate, including nested delegates that capture it. + foreach (AnonymousFunctionExpressionSyntax anonymousFunctionSyntax in context.Node.Ancestors().OfType()) { - // 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) + var anonymousFunctionArgument = anonymousFunctionSyntax.Parent as ArgumentSyntax; + var continuationInvocation = anonymousFunctionArgument?.Parent?.Parent as InvocationExpressionSyntax; + if (continuationInvocation is null || continuationInvocation.ArgumentList.Arguments.FirstOrDefault() != anonymousFunctionArgument) { - // 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; - } - } + continue; + } + + var invokedMemberSymbol = context.SemanticModel.GetSymbolInfo(continuationInvocation, context.CancellationToken).Symbol as IMethodSymbol; + if (invokedMemberSymbol?.Name != nameof(Task.ContinueWith) + || !Utils.IsEqualToOrDerivedFrom(invokedMemberSymbol.ContainingType, taskSymbol)) + { + continue; + } + + ParameterSyntax? firstParameter = GetFirstParameter(anonymousFunctionSyntax); + if (firstParameter is object + && context.SemanticModel.GetDeclaredSymbol(firstParameter, context.CancellationToken) is IParameterSymbol completedTask + && SymbolEqualityComparer.Default.Equals(GetTaskReceiverSymbol(context, memberAccessSyntax), completedTask) + && !IsTaskReassignedInContinuation(context, anonymousFunctionSyntax, memberAccessSyntax, completedTask)) + { + return; } } CSharpCommonInterest.InspectMemberAccess(context, memberAccessSyntax, Descriptor, problematicMethods); } - private static void AnalyzeInvocation(SyntaxNodeAnalysisContext context, INamedTypeSymbol taskSymbol, bool analyzeWholeCodeBlock) + private static void AnalyzeInvocation( + SyntaxNodeAnalysisContext context, + INamedTypeSymbol taskSymbol, + ImmutableArray configuredSyncBlockingMethods, + bool analyzeWholeCodeBlock) { - if (!ShouldAnalyze(context, analyzeWholeCodeBlock)) + var invocationExpressionSyntax = (InvocationExpressionSyntax)context.Node; + if (ShouldAnalyze(context, analyzeWholeCodeBlock)) + { + InspectMemberAccess( + context, + invocationExpressionSyntax.Expression as MemberAccessExpressionSyntax, + CommonInterest.ProblematicSyncBlockingMethods, + taskSymbol); + } + + if (configuredSyncBlockingMethods.IsEmpty + || context.SemanticModel.GetSymbolInfo(invocationExpressionSyntax, context.CancellationToken).Symbol is not IMethodSymbol invokedMethod) { return; } - var invocationExpressionSyntax = (InvocationExpressionSyntax)context.Node; - InspectMemberAccess( - context, - invocationExpressionSyntax.Expression as MemberAccessExpressionSyntax, - CommonInterest.ProblematicSyncBlockingMethods, - taskSymbol); + IMethodSymbol methodDefinition = invokedMethod.ReducedFrom ?? invokedMethod; + bool isBuiltInSyncBlockingMethod = CommonInterest.ProblematicSyncBlockingMethods.Any( + method => method.Method.IsMatch(invokedMethod) || method.Method.IsMatch(methodDefinition)); + if (!isBuiltInSyncBlockingMethod + && configuredSyncBlockingMethods.Any(method => method.IsMatch(invokedMethod) || method.IsMatch(methodDefinition))) + { + SimpleNameSyntax? methodName = invocationExpressionSyntax.Expression switch + { + MemberAccessExpressionSyntax memberAccess => memberAccess.Name, + SimpleNameSyntax simpleName => simpleName, + _ => null, + }; + + if (methodName is object) + { + ImmutableDictionary properties = ImmutableDictionary.Empty.Add("SuppressAwaitCodeFix", null); + context.ReportDiagnostic(Diagnostic.Create(Descriptor, methodName.GetLocation(), properties)); + } + } } private static void AnalyzeMemberAccess(SyntaxNodeAnalysisContext context, INamedTypeSymbol taskSymbol, bool analyzeWholeCodeBlock) @@ -185,4 +216,53 @@ private static void AnalyzeMemberAccess(SyntaxNodeAnalysisContext context, IName CommonInterest.SyncBlockingProperties, taskSymbol); } + + private static ISymbol? GetTaskReceiverSymbol(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) + { + ExpressionSyntax receiver = memberAccessSyntax.Expression; + if (receiver is InvocationExpressionSyntax getAwaiterInvocation + && getAwaiterInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: "GetAwaiter" } getAwaiterAccess) + { + receiver = getAwaiterAccess.Expression; + } + + if (receiver is InvocationExpressionSyntax configureAwaitInvocation + && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess) + { + receiver = configureAwaitAccess.Expression; + } + + return context.SemanticModel.GetSymbolInfo(receiver, context.CancellationToken).Symbol; + } + + private static bool IsTaskReassignedInContinuation( + SyntaxNodeAnalysisContext context, + AnonymousFunctionExpressionSyntax continuation, + MemberAccessExpressionSyntax memberAccess, + IParameterSymbol taskParameter) + { + bool accessIsNested = memberAccess.Ancestors().TakeWhile(node => node != continuation).Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax); + int beforePosition = accessIsNested ? continuation.Span.End + 1 : memberAccess.SpanStart; + + foreach (AssignmentExpressionSyntax assignment in continuation.DescendantNodes().OfType()) + { + if (assignment.SpanStart < beforePosition + && SymbolEqualityComparer.Default.Equals(context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol, taskParameter)) + { + return true; + } + } + + foreach (ArgumentSyntax argument in continuation.DescendantNodes().OfType()) + { + if (argument.SpanStart < beforePosition + && (argument.RefKindKeyword.IsKind(SyntaxKind.RefKeyword) || argument.RefKindKeyword.IsKind(SyntaxKind.OutKeyword)) + && SymbolEqualityComparer.Default.Equals(context.SemanticModel.GetSymbolInfo(argument.Expression, context.CancellationToken).Symbol, taskParameter)) + { + return true; + } + } + + return false; + } } diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index 78aff172a..6337d3eb4 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; @@ -22,6 +23,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 +33,23 @@ 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); + if (semanticModel is null + || 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, @@ -102,8 +117,30 @@ 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? FindGetAwaiterReceiver(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) + { + if (from is InvocationExpressionSyntax + { + Expression: MemberAccessExpressionSyntax + { + Name.Identifier.ValueText: nameof(TaskAwaiter.GetResult), + Expression: InvocationExpressionSyntax + { + Expression: MemberAccessExpressionSyntax + { + Name.Identifier.ValueText: "GetAwaiter", + Expression: ExpressionSyntax receiver, + }, + }, + }, + }) + { + return receiver; + } + + return null; + } + ExpressionSyntax? FindOneLevelDeepIdentifierInvocation(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) => ((from as InvocationExpressionSyntax)?.Expression as MemberAccessExpressionSyntax)?.Expression; ExpressionSyntax? FindParentMemberAccess(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) => @@ -111,10 +148,10 @@ private static bool TryFindNodeAtSource(Diagnostic diagnostic, SyntaxNode root, 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; } @@ -144,4 +181,33 @@ 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) && Utils.IsTask(method.ContainingType)) + { + return method.Parameters.IsEmpty; + } + + 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..61d2f3cda 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 diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 6ac1ac67d..c0f88797d 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() { @@ -583,6 +599,248 @@ void ContinueWith(Func, int> del) { } await CSVerify.VerifyCodeFixAsync(test, expected, test); } + [Fact] + public async Task TaskResultShouldNotReportWarning_WithinNestedDelegateInItsOwnContinuation() + { + var test = @" +using System; +using System.Threading.Tasks; + +class Test { + void F() { + var task = Task.Run(() => 5); + task.ContinueWith(t => { + Action useResultLater = () => Console.WriteLine(t.Result); + useResultLater(); + }); + } +} +"; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task TaskResultReportsWarning_WhenContinuationParameterIsReassigned() + { + var test = @" +using System; +using System.Threading.Tasks; + +class Test { + void F() { + var task = Task.Run(() => 5); + task.ContinueWith(t => { + Action useResultLater = () => Console.WriteLine(t.[|Result|]); + t = Task.Run(() => 6); + useResultLater(); + }); + } +} +"; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + + [Fact] + public async Task TaskWhenAllResultReportsWarningWithoutAnalyzerFailure() + { + var test = @" +using System.Threading.Tasks; + +class Test { + void F() { + var task = Task.Run(() => 1); + _ = Task.WhenAll(task).[|Result|]; + } +} +"; + + await CSVerify.VerifyAnalyzerAsync(test); + } + + [Fact] + public async Task CompletedTaskResultDoesNotReportWarning() + { + var test = @" +using System.Threading.Tasks; + +class Test { + async void Awaited() { + var task = Task.Run(() => 1); + await task; + _ = task.Result; + } + + async void AwaitedAsArgument() { + var task = Task.Run(() => 1); + Consume(await task); + _ = task.Result; + } + + async void Guarded() { + var task = Task.Run(() => 1); + if (!task.IsCompleted) { + await task.ConfigureAwait(false); + } + + _ = 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.Status == TaskStatus.RanToCompletion) { + _ = task.Result; + } + } + + 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 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 Consume(int value) { } + void Consume(int value, Task task) { } + void Consume(Task task, int value) { } +} +"; + + await new CSVerify.Test + { + TestCode = test, + ReferenceAssemblies = Microsoft.CodeAnalysis.Testing.ReferenceAssemblies.Net.Net80, + }.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|](); + } + + 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 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() { From 05a9fc4cab621b827d006351fffdcf86540c1664 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 11:52:42 -0600 Subject: [PATCH 02/26] Address VSTHRD002 review feedback Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 20 ++++++-- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 48 ++++++++++++------ .../VSTHRD002UseJtfRunAnalyzerTests.cs | 50 ++++++++++++++++++- 3 files changed, 98 insertions(+), 20 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index 92692fcf8..adf36a197 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -299,9 +299,23 @@ private static bool StatementCompletesTask(SyntaxNodeAnalysisContext context, St return true; } - return statement is IfStatementSyntax { Else: null } ifStatement - && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: false) - && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbol); + if (statement is not IfStatementSyntax ifStatement) + { + return false; + } + + if (ifStatement.Else is null) + { + return ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: false) + && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbol); + } + + return (ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: true) + && !MayReassignTask(context, ifStatement.Statement, taskSymbol) + && StatementDefinitelyAwaitsTask(context, ifStatement.Else.Statement, taskSymbol)) + || (ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: false) + && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbol) + && !MayReassignTask(context, ifStatement.Else.Statement, taskSymbol)); } private static bool StatementDefinitelyAwaitsTask(SyntaxNodeAnalysisContext context, StatementSyntax statement, ISymbol taskSymbol) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index 6337d3eb4..5fcbff2e9 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -119,30 +119,46 @@ private static bool TryFindNodeAtSource(Diagnostic diagnostic, SyntaxNode root, ExpressionSyntax? FindGetAwaiterReceiver(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) { - if (from is InvocationExpressionSyntax + if (from is not InvocationExpressionSyntax { Expression: MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(TaskAwaiter.GetResult), - Expression: InvocationExpressionSyntax - { - Expression: MemberAccessExpressionSyntax - { - Name.Identifier.ValueText: "GetAwaiter", - Expression: ExpressionSyntax receiver, - }, - }, - }, + } getResultAccess }) { - return receiver; + return null; + } + + ExpressionSyntax getAwaiterInvocationExpression = getResultAccess.Expression; + while (getAwaiterInvocationExpression is ParenthesizedExpressionSyntax parenthesized) + { + getAwaiterInvocationExpression = parenthesized.Expression; } - return null; + return getAwaiterInvocationExpression is InvocationExpressionSyntax + { + Expression: MemberAccessExpressionSyntax + { + Name.Identifier.ValueText: "GetAwaiter", + Expression: ExpressionSyntax receiver, + } + } + ? receiver + : null; } - ExpressionSyntax? FindOneLevelDeepIdentifierInvocation(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) => - ((from as InvocationExpressionSyntax)?.Expression as MemberAccessExpressionSyntax)?.Expression; + ExpressionSyntax? FindInstanceWaitReceiver(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) => + from is InvocationExpressionSyntax + { + Expression: MemberAccessExpressionSyntax + { + Name.Identifier.ValueText: nameof(Task.Wait), + Expression: ExpressionSyntax receiver, + } + } + ? receiver + : null; ExpressionSyntax? FindParentMemberAccess(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) => (from as MemberAccessExpressionSyntax)?.Expression; @@ -162,10 +178,10 @@ 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; } diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index c0f88797d..9236df1a2 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -656,7 +656,18 @@ void F() { } "; - await CSVerify.VerifyAnalyzerAsync(test); + var withFix = @" +using System.Threading.Tasks; + +class Test { + async Task FAsync() { + var task = Task.Run(() => 1); + _ = await Task.WhenAll(task); + } +} +"; + + await CSVerify.VerifyCodeFixAsync(test, withFix); } [Fact] @@ -687,6 +698,16 @@ async void Guarded() { _ = task.Result; } + async void GuardedWithElse() { + var task = Task.Run(() => 1); + if (task.IsCompleted) { + } else { + await task.ConfigureAwait(false); + } + + _ = task.Result; + } + async void AwaitedBeforeNestedBlock(bool condition) { var task = Task.Run(() => 1); await task; @@ -870,6 +891,33 @@ async Task FAsync() { 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() { From 8d0cc24b1563ad17900281c23fb4d7496b89cf82 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 12:04:39 -0600 Subject: [PATCH 03/26] Fix code fix analyzer style violations Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 38 +++++-------------- 1 file changed, 10 insertions(+), 28 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index 5fcbff2e9..d7bf03d79 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -119,13 +119,8 @@ private static bool TryFindNodeAtSource(Diagnostic diagnostic, SyntaxNode root, ExpressionSyntax? FindGetAwaiterReceiver(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) { - if (from is not InvocationExpressionSyntax - { - Expression: MemberAccessExpressionSyntax - { - Name.Identifier.ValueText: nameof(TaskAwaiter.GetResult), - } getResultAccess - }) + var getResultAccess = (from as InvocationExpressionSyntax)?.Expression as MemberAccessExpressionSyntax; + if (getResultAccess?.Name.Identifier.ValueText != nameof(TaskAwaiter.GetResult)) { return null; } @@ -136,29 +131,16 @@ private static bool TryFindNodeAtSource(Diagnostic diagnostic, SyntaxNode root, getAwaiterInvocationExpression = parenthesized.Expression; } - return getAwaiterInvocationExpression is InvocationExpressionSyntax - { - Expression: MemberAccessExpressionSyntax - { - Name.Identifier.ValueText: "GetAwaiter", - Expression: ExpressionSyntax receiver, - } - } - ? receiver - : null; + var getAwaiterAccess = (getAwaiterInvocationExpression as InvocationExpressionSyntax)?.Expression as MemberAccessExpressionSyntax; + return getAwaiterAccess?.Name.Identifier.ValueText == "GetAwaiter" ? 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? FindInstanceWaitReceiver(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) => - from is InvocationExpressionSyntax - { - Expression: MemberAccessExpressionSyntax - { - Name.Identifier.ValueText: nameof(Task.Wait), - Expression: ExpressionSyntax receiver, - } - } - ? receiver - : null; ExpressionSyntax? FindParentMemberAccess(ExpressionSyntax? from, CancellationToken cancellationToken = default(CancellationToken)) => (from as MemberAccessExpressionSyntax)?.Expression; From a6c02eaf2648f02492f81f7b5ce70f72336fe007 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 12:07:38 -0600 Subject: [PATCH 04/26] Handle parenthesized awaiters in continuations Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 8 ++++---- .../VSTHRD002UseJtfRunAnalyzer.cs | 16 +++++++++++++--- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 2 ++ 3 files changed, 19 insertions(+), 7 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index adf36a197..1193b4ed9 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -127,20 +127,20 @@ private static ExpressionSyntax UnwrapParentheses(ExpressionSyntax expression) private static ExpressionSyntax GetTaskReceiver(MemberAccessExpressionSyntax memberAccessSyntax) { - ExpressionSyntax receiver = memberAccessSyntax.Expression; + ExpressionSyntax receiver = UnwrapParentheses(memberAccessSyntax.Expression); if (receiver is InvocationExpressionSyntax getAwaiterInvocation && getAwaiterInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: "GetAwaiter" } getAwaiterAccess) { - receiver = getAwaiterAccess.Expression; + receiver = UnwrapParentheses(getAwaiterAccess.Expression); } if (receiver is InvocationExpressionSyntax configureAwaitInvocation && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess) { - receiver = configureAwaitAccess.Expression; + receiver = UnwrapParentheses(configureAwaitAccess.Expression); } - return UnwrapParentheses(receiver); + return receiver; } private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index 6fe7f3919..e624e01a1 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -219,22 +219,32 @@ private static void AnalyzeMemberAccess(SyntaxNodeAnalysisContext context, IName private static ISymbol? GetTaskReceiverSymbol(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) { - ExpressionSyntax receiver = memberAccessSyntax.Expression; + ExpressionSyntax receiver = UnwrapParentheses(memberAccessSyntax.Expression); if (receiver is InvocationExpressionSyntax getAwaiterInvocation && getAwaiterInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: "GetAwaiter" } getAwaiterAccess) { - receiver = getAwaiterAccess.Expression; + receiver = UnwrapParentheses(getAwaiterAccess.Expression); } if (receiver is InvocationExpressionSyntax configureAwaitInvocation && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess) { - receiver = configureAwaitAccess.Expression; + receiver = UnwrapParentheses(configureAwaitAccess.Expression); } return context.SemanticModel.GetSymbolInfo(receiver, context.CancellationToken).Symbol; } + private static ExpressionSyntax UnwrapParentheses(ExpressionSyntax expression) + { + while (expression is ParenthesizedExpressionSyntax parenthesized) + { + expression = parenthesized.Expression; + } + + return expression; + } + private static bool IsTaskReassignedInContinuation( SyntaxNodeAnalysisContext context, AnonymousFunctionExpressionSyntax continuation, diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 9236df1a2..cfa3f4bfd 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -584,6 +584,8 @@ 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()); } void ContinueWith(Func, int> del) { } From 1d594a30f360d373da3382296d79606e50c3d10e Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 12:31:06 -0600 Subject: [PATCH 05/26] Harden VSTHRD002 completion proofs Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 70 ++++++++- .../VSTHRD002UseJtfRunAnalyzer.cs | 16 +- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 3 +- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 143 ++++++++++++++++++ 4 files changed, 225 insertions(+), 7 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index 1193b4ed9..cf0328bae 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -151,6 +151,11 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc return false; } + if (NestedFunctionMayReassignTask(context, memberAccessSyntax, taskSymbol)) + { + return false; + } + if (IsWithinCompletedTaskBranch(context, memberAccessSyntax, taskSymbol)) { return true; @@ -162,6 +167,13 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc return false; } + if (TryGetAwaitExpression(containingStatement, memberAccessSyntax.SpanStart, out AwaitExpressionSyntax? precedingAwait) + && AwaitCompletesTask(context, precedingAwait, taskSymbol) + && !MayReassignTask(context, containingStatement, taskSymbol, precedingAwait.Span.End, memberAccessSyntax.SpanStart)) + { + return true; + } + while (containingStatement.Parent is BlockSyntax block) { if (MayReassignTask(context, containingStatement, taskSymbol, containingStatement.SpanStart - 1, memberAccessSyntax.SpanStart)) @@ -299,11 +311,21 @@ private static bool StatementCompletesTask(SyntaxNodeAnalysisContext context, St return true; } + if (statement is BlockSyntax block) + { + return StatementDefinitelyAwaitsTask(context, block, taskSymbol); + } + if (statement is not IfStatementSyntax ifStatement) { return false; } + if (MayReassignTask(context, ifStatement.Condition, taskSymbol)) + { + return false; + } + if (ifStatement.Else is null) { return ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: false) @@ -348,11 +370,25 @@ private static bool StatementDefinitelyAwaitsTask(SyntaxNodeAnalysisContext cont } private static bool TryGetAwaitExpression(StatementSyntax statement, [NotNullWhen(true)] out AwaitExpressionSyntax? awaitExpression) + => TryGetAwaitExpression(statement, statement.Span.End + 1, out awaitExpression); + + private static bool TryGetAwaitExpression(StatementSyntax statement, int beforePosition, [NotNullWhen(true)] out AwaitExpressionSyntax? awaitExpression) { static bool DescendIntoChildren(SyntaxNode node) => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; + if (statement is WhileStatementSyntax or DoStatementSyntax or ForStatementSyntax or ForEachStatementSyntax or ForEachVariableStatementSyntax) + { + awaitExpression = null; + return false; + } + foreach (AwaitExpressionSyntax candidate in statement.DescendantNodes(DescendIntoChildren).OfType().Reverse()) { + if (candidate.Span.End >= beforePosition) + { + continue; + } + IEnumerable ancestorsWithinStatement = candidate.Ancestors().TakeWhile(node => node != statement); bool isConditionallyExecuted = ancestorsWithinStatement.Any( node => node is StatementSyntax or ConditionalExpressionSyntax or SwitchExpressionSyntax or ConditionalAccessExpressionSyntax @@ -373,6 +409,11 @@ private static bool TryGetAwaitExpression(StatementSyntax statement, [NotNullWhe private static bool AwaitCompletesTask(SyntaxNodeAnalysisContext context, AwaitExpressionSyntax awaitExpression, ISymbol taskSymbol) { + if (MayReassignTask(context, awaitExpression.Expression, taskSymbol)) + { + return false; + } + ExpressionSyntax awaitedExpression = UnwrapParentheses(awaitExpression.Expression); if (awaitedExpression is InvocationExpressionSyntax configureAwaitInvocation && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess) @@ -402,19 +443,20 @@ private static bool MayReassignTask(SyntaxNodeAnalysisContext context, SyntaxNod private static bool MayReassignTask(SyntaxNodeAnalysisContext context, SyntaxNode node, ISymbol taskSymbol, int afterPosition, int beforePosition) { - static bool DescendIntoChildren(SyntaxNode node) => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; + bool DescendIntoChildren(SyntaxNode child) => + child == node || child is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; - foreach (AssignmentExpressionSyntax assignment in node.DescendantNodes(DescendIntoChildren).OfType()) + foreach (AssignmentExpressionSyntax assignment in node.DescendantNodesAndSelf(DescendIntoChildren).OfType()) { if (assignment.SpanStart > afterPosition && assignment.SpanStart < beforePosition - && SymbolEqualityComparer.Default.Equals(context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol, taskSymbol)) + && IsAssignmentToTask(context, assignment.Left, taskSymbol)) { return true; } } - foreach (ArgumentSyntax argument in node.DescendantNodes(DescendIntoChildren).OfType()) + foreach (ArgumentSyntax argument in node.DescendantNodesAndSelf(DescendIntoChildren).OfType()) { if (argument.SpanStart > afterPosition && argument.SpanStart < beforePosition @@ -428,6 +470,26 @@ private static bool MayReassignTask(SyntaxNodeAnalysisContext context, SyntaxNod return false; } + private static bool NestedFunctionMayReassignTask(SyntaxNodeAnalysisContext context, SyntaxNode node, ISymbol taskSymbol) + { + SyntaxNode? containingFunction = node.Ancestors().FirstOrDefault( + ancestor => ancestor is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax or BaseMethodDeclarationSyntax); + return containingFunction?.DescendantNodes() + .Where(descendant => descendant is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax) + .Any(nestedFunction => MayReassignTask(context, nestedFunction, taskSymbol)) is true; + } + + private static bool IsAssignmentToTask(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, ISymbol taskSymbol) + { + expression = UnwrapParentheses(expression); + if (expression is TupleExpressionSyntax tuple) + { + return tuple.Arguments.Any(argument => IsAssignmentToTask(context, argument.Expression, taskSymbol)); + } + + return IsSameTask(context, expression, taskSymbol); + } + private static bool IsSameTask(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, ISymbol taskSymbol) => SymbolEqualityComparer.Default.Equals( context.SemanticModel.GetSymbolInfo(UnwrapParentheses(expression), context.CancellationToken).Symbol, diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index e624e01a1..ada10c898 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -190,11 +190,12 @@ private static void AnalyzeInvocation( SimpleNameSyntax? methodName = invocationExpressionSyntax.Expression switch { MemberAccessExpressionSyntax memberAccess => memberAccess.Name, + MemberBindingExpressionSyntax memberBinding => memberBinding.Name, SimpleNameSyntax simpleName => simpleName, _ => null, }; - if (methodName is object) + if (methodName is object && !CSharpCommonInterest.ShouldIgnoreContext(context)) { ImmutableDictionary properties = ImmutableDictionary.Empty.Add("SuppressAwaitCodeFix", null); context.ReportDiagnostic(Diagnostic.Create(Descriptor, methodName.GetLocation(), properties)); @@ -257,7 +258,7 @@ private static bool IsTaskReassignedInContinuation( foreach (AssignmentExpressionSyntax assignment in continuation.DescendantNodes().OfType()) { if (assignment.SpanStart < beforePosition - && SymbolEqualityComparer.Default.Equals(context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol, taskParameter)) + && IsAssignmentToParameter(context, assignment.Left, taskParameter)) { return true; } @@ -275,4 +276,15 @@ private static bool IsTaskReassignedInContinuation( return false; } + + private static bool IsAssignmentToParameter(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, IParameterSymbol taskParameter) + { + expression = UnwrapParentheses(expression); + if (expression is TupleExpressionSyntax tuple) + { + return tuple.Arguments.Any(argument => IsAssignmentToParameter(context, argument.Expression, taskParameter)); + } + + return SymbolEqualityComparer.Default.Equals(context.SemanticModel.GetSymbolInfo(expression, context.CancellationToken).Symbol, taskParameter); + } } diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index d7bf03d79..4711e1915 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -167,7 +167,8 @@ private static bool TryFindNodeAtSource(Diagnostic diagnostic, SyntaxNode root, 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); diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index cfa3f4bfd..b888aeadf 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -644,6 +644,26 @@ void F() { await CSVerify.VerifyCodeFixAsync(test, test); } + [Fact] + public async Task ContinuationParameterDeconstructionReportsWarning() + { + var test = @" +using System.Threading.Tasks; + +class Test { + void F() { + var task = Task.Run(() => 5); + task.ContinueWith(t => { + (t, _) = (Task.Run(() => 6), 0); + return t.[|Result|]; + }); + } +} +"; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + [Fact] public async Task TaskWhenAllResultReportsWarningWithoutAnalyzerFailure() { @@ -676,6 +696,7 @@ async Task FAsync() { public async Task CompletedTaskResultDoesNotReportWarning() { var test = @" +using System; using System.Threading.Tasks; class Test { @@ -691,6 +712,20 @@ async void AwaitedAsArgument() { _ = task.Result; } + async void AwaitedEarlierInSameStatement() { + var task = Task.Run(() => 1); + Consume(await task, task.Result); + } + + async void AwaitedInNestedBlock() { + var task = Task.Run(() => 1); + { + await task; + } + + _ = task.Result; + } + async void Guarded() { var task = Task.Run(() => 1); if (!task.IsCompleted) { @@ -762,6 +797,59 @@ async void ReassignedLaterInAwaitStatement() { _ = 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|]; + } + + async void AwaitInDoWhileConditionIsNotDefinite() { + var task = Task.Run(() => true); + do { + break; + } while (await task); + + _ = task.[|Result|]; + } + async void ReassignedEarlierInResultStatement() { var task = Task.Run(() => 1); await task; @@ -776,8 +864,14 @@ void ReassignedInsideGuard(Task task) { } 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; + } } "; @@ -809,6 +903,7 @@ class Test { void F(Task task, Contoso.Threading.CustomWaiter waiter) { _ = task.[|WaitSynchronously|](); waiter.[|Join|](); + waiter?.[|Join|](); } Task FAsync(Task task) { @@ -830,6 +925,25 @@ Task FAsync(Task task) { 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 CodeFixIsNotOfferedWhenBlockingExpressionHasCompileErrors() { @@ -1141,6 +1255,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() { From 28db7ee5895aadbcc57d4947d49e0c4228747097 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 12:43:05 -0600 Subject: [PATCH 06/26] Track VSTHRD002 ref aliases Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 200 ++++++++++++++---- .../VSTHRD002UseJtfRunAnalyzer.cs | 14 +- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 35 +++ 3 files changed, 197 insertions(+), 52 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index cf0328bae..fc2768a23 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -29,6 +29,62 @@ internal static class CSharpCommonInterest SyntaxKind.AddAccessorDeclaration, SyntaxKind.RemoveAccessorDeclaration); + /// + /// Gets a symbol and any ref locals that alias it within the containing function. + /// + internal static ImmutableHashSet GetSymbolAndRefAliases(SyntaxNodeAnalysisContext context, SyntaxNode node, ISymbol symbol) + { + ImmutableHashSet.Builder symbols = ImmutableHashSet.CreateBuilder(SymbolEqualityComparer.Default); + symbols.Add(symbol); + + SyntaxNode searchRoot = node.AncestorsAndSelf().FirstOrDefault( + ancestor => ancestor is AnonymousFunctionExpressionSyntax + or LocalFunctionStatementSyntax + or BaseMethodDeclarationSyntax + or AccessorDeclarationSyntax) ?? node; + + bool addedAlias; + do + { + addedAlias = false; + foreach (VariableDeclaratorSyntax variable in searchRoot.DescendantNodes().OfType()) + { + if (variable.Initializer is null + || context.SemanticModel.GetDeclaredSymbol(variable, context.CancellationToken) is not ILocalSymbol { RefKind: not RefKind.None } local) + { + continue; + } + + ExpressionSyntax initializer = variable.Initializer.Value is RefExpressionSyntax refExpression + ? refExpression.Expression + : variable.Initializer.Value; + ISymbol? initializedFrom = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(initializer), context.CancellationToken).Symbol; + if (initializedFrom is object && symbols.Contains(initializedFrom) && symbols.Add(local)) + { + addedAlias = true; + } + } + + foreach (AssignmentExpressionSyntax assignment in searchRoot.DescendantNodes().OfType()) + { + if (assignment.Right is not RefExpressionSyntax refExpression + || context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol is not ILocalSymbol { RefKind: not RefKind.None } local) + { + continue; + } + + ISymbol? assignedFrom = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(refExpression.Expression), context.CancellationToken).Symbol; + if (assignedFrom is object && symbols.Contains(assignedFrom) && symbols.Add(local)) + { + addedAlias = true; + } + } + } + while (addedAlias); + + return symbols.ToImmutable(); + } + /// /// This is an explicit rule to ignore the code that was generated by Xaml2CS. /// @@ -151,12 +207,13 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc return false; } - if (NestedFunctionMayReassignTask(context, memberAccessSyntax, taskSymbol)) + ImmutableHashSet taskSymbols = GetSymbolAndRefAliases(context, memberAccessSyntax, taskSymbol); + if (NestedFunctionMayReassignTask(context, memberAccessSyntax, taskSymbols)) { return false; } - if (IsWithinCompletedTaskBranch(context, memberAccessSyntax, taskSymbol)) + if (IsWithinCompletedTaskBranch(context, memberAccessSyntax, taskSymbol, taskSymbols)) { return true; } @@ -167,16 +224,15 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc return false; } - if (TryGetAwaitExpression(containingStatement, memberAccessSyntax.SpanStart, out AwaitExpressionSyntax? precedingAwait) - && AwaitCompletesTask(context, precedingAwait, taskSymbol) - && !MayReassignTask(context, containingStatement, taskSymbol, precedingAwait.Span.End, memberAccessSyntax.SpanStart)) + if (TryGetAwaitExpression(context, containingStatement, taskSymbol, taskSymbols, memberAccessSyntax.SpanStart, out AwaitExpressionSyntax? precedingAwait) + && !MayReassignTask(context, containingStatement, taskSymbols, precedingAwait.Span.End, memberAccessSyntax.SpanStart)) { return true; } while (containingStatement.Parent is BlockSyntax block) { - if (MayReassignTask(context, containingStatement, taskSymbol, containingStatement.SpanStart - 1, memberAccessSyntax.SpanStart)) + if (MayReassignTask(context, containingStatement, taskSymbols, containingStatement.SpanStart - 1, memberAccessSyntax.SpanStart)) { return false; } @@ -185,12 +241,12 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc for (int i = statementIndex - 1; i >= 0; i--) { StatementSyntax statement = block.Statements[i]; - if (StatementCompletesTask(context, statement, taskSymbol)) + if (StatementCompletesTask(context, statement, taskSymbol, taskSymbols)) { return true; } - if (MayReassignTask(context, statement, taskSymbol)) + if (MayReassignTask(context, statement, taskSymbols)) { return false; } @@ -202,7 +258,7 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc return false; } - if (MayReassignTask(context, outerStatement, taskSymbol, outerStatement.SpanStart - 1, containingStatement.SpanStart)) + if (MayReassignTask(context, outerStatement, taskSymbols, outerStatement.SpanStart - 1, containingStatement.SpanStart)) { return false; } @@ -213,7 +269,11 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc return false; } - private static bool IsWithinCompletedTaskBranch(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax, ISymbol taskSymbol) + private static bool IsWithinCompletedTaskBranch( + SyntaxNodeAnalysisContext context, + MemberAccessExpressionSyntax memberAccessSyntax, + ISymbol taskSymbol, + IImmutableSet taskSymbols) { foreach (IfStatementSyntax ifStatement in memberAccessSyntax.Ancestors().OfType()) { @@ -225,14 +285,14 @@ private static bool IsWithinCompletedTaskBranch(SyntaxNodeAnalysisContext contex if (ifStatement.Statement.FullSpan.Contains(memberAccessSyntax.Span) && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: true) - && !MayReassignTask(context, ifStatement, taskSymbol, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) + && !MayReassignTask(context, ifStatement, taskSymbols, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) { return true; } if (ifStatement.Else?.Statement.FullSpan.Contains(memberAccessSyntax.Span) is true && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: false) - && !MayReassignTask(context, ifStatement, taskSymbol, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) + && !MayReassignTask(context, ifStatement, taskSymbols, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) { return true; } @@ -263,6 +323,18 @@ private static bool ConditionProvesCompletion(SyntaxNodeAnalysisContext context, || ConditionProvesCompletion(context, binary.Right, taskSymbol, conditionValue: false); } + if (conditionValue && binary.IsKind(SyntaxKind.LogicalOrExpression)) + { + return ConditionProvesCompletion(context, binary.Left, taskSymbol, conditionValue: true) + && ConditionProvesCompletion(context, binary.Right, taskSymbol, conditionValue: true); + } + + if (!conditionValue && binary.IsKind(SyntaxKind.LogicalAndExpression)) + { + return ConditionProvesCompletion(context, binary.Left, taskSymbol, conditionValue: false) + && ConditionProvesCompletion(context, binary.Right, taskSymbol, conditionValue: false); + } + bool equalityHolds = binary.IsKind(SyntaxKind.EqualsExpression) ? conditionValue : binary.IsKind(SyntaxKind.NotEqualsExpression) ? !conditionValue : false; @@ -302,18 +374,21 @@ private static bool IsRanToCompletion(SyntaxNodeAnalysisContext context, Express && status.ContainingType.Name == nameof(TaskStatus) && status.ContainingType.BelongsToNamespace(Namespaces.SystemThreadingTasks); - private static bool StatementCompletesTask(SyntaxNodeAnalysisContext context, StatementSyntax statement, ISymbol taskSymbol) + private static bool StatementCompletesTask( + SyntaxNodeAnalysisContext context, + StatementSyntax statement, + ISymbol taskSymbol, + IImmutableSet taskSymbols) { - if (TryGetAwaitExpression(statement, out AwaitExpressionSyntax? awaitExpression) - && AwaitCompletesTask(context, awaitExpression, taskSymbol) - && !MayReassignTask(context, statement, taskSymbol, awaitExpression.Span.End, statement.Span.End + 1)) + if (TryGetAwaitExpression(context, statement, taskSymbol, taskSymbols, out AwaitExpressionSyntax? awaitExpression) + && !MayReassignTask(context, statement, taskSymbols, awaitExpression.Span.End, statement.Span.End + 1)) { return true; } if (statement is BlockSyntax block) { - return StatementDefinitelyAwaitsTask(context, block, taskSymbol); + return StatementDefinitelyAwaitsTask(context, block, taskSymbol, taskSymbols); } if (statement is not IfStatementSyntax ifStatement) @@ -321,7 +396,7 @@ private static bool StatementCompletesTask(SyntaxNodeAnalysisContext context, St return false; } - if (MayReassignTask(context, ifStatement.Condition, taskSymbol)) + if (MayReassignTask(context, ifStatement.Condition, taskSymbols)) { return false; } @@ -329,37 +404,39 @@ private static bool StatementCompletesTask(SyntaxNodeAnalysisContext context, St if (ifStatement.Else is null) { return ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: false) - && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbol); + && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbol, taskSymbols); } return (ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: true) - && !MayReassignTask(context, ifStatement.Statement, taskSymbol) - && StatementDefinitelyAwaitsTask(context, ifStatement.Else.Statement, taskSymbol)) + && !MayReassignTask(context, ifStatement.Statement, taskSymbols) + && StatementDefinitelyAwaitsTask(context, ifStatement.Else.Statement, taskSymbol, taskSymbols)) || (ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: false) - && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbol) - && !MayReassignTask(context, ifStatement.Else.Statement, taskSymbol)); + && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbol, taskSymbols) + && !MayReassignTask(context, ifStatement.Else.Statement, taskSymbols)); } - private static bool StatementDefinitelyAwaitsTask(SyntaxNodeAnalysisContext context, StatementSyntax statement, ISymbol taskSymbol) + private static bool StatementDefinitelyAwaitsTask( + SyntaxNodeAnalysisContext context, + StatementSyntax statement, + ISymbol taskSymbol, + IImmutableSet taskSymbols) { - if (TryGetAwaitExpression(statement, out AwaitExpressionSyntax? awaitExpression)) + if (TryGetAwaitExpression(context, statement, taskSymbol, taskSymbols, out AwaitExpressionSyntax? awaitExpression)) { - return AwaitCompletesTask(context, awaitExpression, taskSymbol) - && !MayReassignTask(context, statement, taskSymbol, awaitExpression.Span.End, statement.Span.End + 1); + return !MayReassignTask(context, statement, taskSymbols, awaitExpression.Span.End, statement.Span.End + 1); } if (statement is BlockSyntax block) { for (int i = block.Statements.Count - 1; i >= 0; i--) { - if (TryGetAwaitExpression(block.Statements[i], out awaitExpression) - && AwaitCompletesTask(context, awaitExpression, taskSymbol) - && !MayReassignTask(context, block.Statements[i], taskSymbol, awaitExpression.Span.End, block.Statements[i].Span.End + 1)) + if (TryGetAwaitExpression(context, block.Statements[i], taskSymbol, taskSymbols, out awaitExpression) + && !MayReassignTask(context, block.Statements[i], taskSymbols, awaitExpression.Span.End, block.Statements[i].Span.End + 1)) { return true; } - if (MayReassignTask(context, block.Statements[i], taskSymbol)) + if (MayReassignTask(context, block.Statements[i], taskSymbols)) { return false; } @@ -369,10 +446,21 @@ private static bool StatementDefinitelyAwaitsTask(SyntaxNodeAnalysisContext cont return false; } - private static bool TryGetAwaitExpression(StatementSyntax statement, [NotNullWhen(true)] out AwaitExpressionSyntax? awaitExpression) - => TryGetAwaitExpression(statement, statement.Span.End + 1, out awaitExpression); + private static bool TryGetAwaitExpression( + SyntaxNodeAnalysisContext context, + StatementSyntax statement, + ISymbol taskSymbol, + IImmutableSet taskSymbols, + [NotNullWhen(true)] out AwaitExpressionSyntax? awaitExpression) + => TryGetAwaitExpression(context, statement, taskSymbol, taskSymbols, statement.Span.End + 1, out awaitExpression); - private static bool TryGetAwaitExpression(StatementSyntax statement, int beforePosition, [NotNullWhen(true)] out AwaitExpressionSyntax? awaitExpression) + private static bool TryGetAwaitExpression( + SyntaxNodeAnalysisContext context, + StatementSyntax statement, + ISymbol taskSymbol, + IImmutableSet taskSymbols, + int beforePosition, + [NotNullWhen(true)] out AwaitExpressionSyntax? awaitExpression) { static bool DescendIntoChildren(SyntaxNode node) => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; @@ -396,7 +484,7 @@ private static bool TryGetAwaitExpression(StatementSyntax statement, int beforeP && (binary.IsKind(SyntaxKind.LogicalAndExpression) || binary.IsKind(SyntaxKind.LogicalOrExpression) || binary.IsKind(SyntaxKind.CoalesceExpression)))); - if (!isConditionallyExecuted) + if (!isConditionallyExecuted && AwaitCompletesTask(context, candidate, taskSymbol, taskSymbols)) { awaitExpression = candidate; return true; @@ -407,9 +495,13 @@ private static bool TryGetAwaitExpression(StatementSyntax statement, int beforeP return false; } - private static bool AwaitCompletesTask(SyntaxNodeAnalysisContext context, AwaitExpressionSyntax awaitExpression, ISymbol taskSymbol) + private static bool AwaitCompletesTask( + SyntaxNodeAnalysisContext context, + AwaitExpressionSyntax awaitExpression, + ISymbol taskSymbol, + IImmutableSet taskSymbols) { - if (MayReassignTask(context, awaitExpression.Expression, taskSymbol)) + if (MayReassignTask(context, awaitExpression.Expression, taskSymbols)) { return false; } @@ -438,10 +530,15 @@ private static bool AwaitCompletesTask(SyntaxNodeAnalysisContext context, AwaitE return false; } - private static bool MayReassignTask(SyntaxNodeAnalysisContext context, SyntaxNode node, ISymbol taskSymbol) - => MayReassignTask(context, node, taskSymbol, node.SpanStart - 1, node.Span.End + 1); + 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, ISymbol taskSymbol, int afterPosition, int beforePosition) + 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; @@ -450,7 +547,7 @@ bool DescendIntoChildren(SyntaxNode child) => { if (assignment.SpanStart > afterPosition && assignment.SpanStart < beforePosition - && IsAssignmentToTask(context, assignment.Left, taskSymbol)) + && IsAssignmentToTask(context, assignment.Left, taskSymbols)) { return true; } @@ -461,7 +558,7 @@ bool DescendIntoChildren(SyntaxNode child) => if (argument.SpanStart > afterPosition && argument.SpanStart < beforePosition && (argument.RefKindKeyword.IsKind(SyntaxKind.RefKeyword) || argument.RefKindKeyword.IsKind(SyntaxKind.OutKeyword)) - && IsSameTask(context, argument.Expression, taskSymbol)) + && IsOneOfSymbols(context, argument.Expression, taskSymbols)) { return true; } @@ -470,28 +567,37 @@ bool DescendIntoChildren(SyntaxNode child) => return false; } - private static bool NestedFunctionMayReassignTask(SyntaxNodeAnalysisContext context, SyntaxNode node, ISymbol taskSymbol) + private static bool NestedFunctionMayReassignTask( + SyntaxNodeAnalysisContext context, + SyntaxNode node, + IImmutableSet taskSymbols) { SyntaxNode? containingFunction = node.Ancestors().FirstOrDefault( ancestor => ancestor is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax or BaseMethodDeclarationSyntax); return containingFunction?.DescendantNodes() .Where(descendant => descendant is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax) - .Any(nestedFunction => MayReassignTask(context, nestedFunction, taskSymbol)) is true; + .Any(nestedFunction => MayReassignTask(context, nestedFunction, taskSymbols)) is true; } - private static bool IsAssignmentToTask(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, ISymbol taskSymbol) + 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, taskSymbol)); + return tuple.Arguments.Any(argument => IsAssignmentToTask(context, argument.Expression, taskSymbols)); } - return IsSameTask(context, expression, taskSymbol); + return IsOneOfSymbols(context, expression, taskSymbols); } private static bool IsSameTask(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, ISymbol taskSymbol) => SymbolEqualityComparer.Default.Equals( context.SemanticModel.GetSymbolInfo(UnwrapParentheses(expression), context.CancellationToken).Symbol, taskSymbol); + + 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/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index ada10c898..29ea1050e 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -254,11 +254,12 @@ private static bool IsTaskReassignedInContinuation( { bool accessIsNested = memberAccess.Ancestors().TakeWhile(node => node != continuation).Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax); int beforePosition = accessIsNested ? continuation.Span.End + 1 : memberAccess.SpanStart; + ImmutableHashSet taskSymbols = CSharpCommonInterest.GetSymbolAndRefAliases(context, continuation, taskParameter); foreach (AssignmentExpressionSyntax assignment in continuation.DescendantNodes().OfType()) { if (assignment.SpanStart < beforePosition - && IsAssignmentToParameter(context, assignment.Left, taskParameter)) + && IsAssignmentToParameter(context, assignment.Left, taskSymbols)) { return true; } @@ -266,9 +267,11 @@ private static bool IsTaskReassignedInContinuation( foreach (ArgumentSyntax argument in continuation.DescendantNodes().OfType()) { + ISymbol? argumentSymbol = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(argument.Expression), context.CancellationToken).Symbol; if (argument.SpanStart < beforePosition && (argument.RefKindKeyword.IsKind(SyntaxKind.RefKeyword) || argument.RefKindKeyword.IsKind(SyntaxKind.OutKeyword)) - && SymbolEqualityComparer.Default.Equals(context.SemanticModel.GetSymbolInfo(argument.Expression, context.CancellationToken).Symbol, taskParameter)) + && argumentSymbol is object + && taskSymbols.Contains(argumentSymbol)) { return true; } @@ -277,14 +280,15 @@ private static bool IsTaskReassignedInContinuation( return false; } - private static bool IsAssignmentToParameter(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, IParameterSymbol taskParameter) + private static bool IsAssignmentToParameter(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, IImmutableSet taskSymbols) { expression = UnwrapParentheses(expression); if (expression is TupleExpressionSyntax tuple) { - return tuple.Arguments.Any(argument => IsAssignmentToParameter(context, argument.Expression, taskParameter)); + return tuple.Arguments.Any(argument => IsAssignmentToParameter(context, argument.Expression, taskSymbols)); } - return SymbolEqualityComparer.Default.Equals(context.SemanticModel.GetSymbolInfo(expression, context.CancellationToken).Symbol, taskParameter); + ISymbol? symbol = context.SemanticModel.GetSymbolInfo(expression, context.CancellationToken).Symbol; + return symbol is object && taskSymbols.Contains(symbol); } } diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index b888aeadf..03e7d8dcc 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -657,6 +657,13 @@ void F() { (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|]; + }); } } "; @@ -717,6 +724,13 @@ async void AwaitedEarlierInSameStatement() { Consume(await task, task.Result); } + async void AwaitedBeforeOtherTaskInSameStatement() { + var task = Task.Run(() => 1); + var otherTask = Task.Run(() => 2); + Consume(await task, await otherTask); + _ = task.Result; + } + async void AwaitedInNestedBlock() { var task = Task.Run(() => 1); { @@ -770,6 +784,10 @@ void CompletionProperties(Task task) { _ = task.Result; } + if (task.IsCompleted || task.IsCanceled) { + _ = task.Result; + } + if (task.Status == TaskStatus.RanToCompletion) { _ = task.Result; } @@ -863,6 +881,23 @@ void ReassignedInsideGuard(Task task) { } } + 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 Consume(int value) { } void Consume(int first, int second) { } void Consume(int value, Task task) { } From 60af3a241fd06ae1c16c14d221021a4e4580a99a Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 12:47:42 -0600 Subject: [PATCH 07/26] Refine VSTHRD002 control flow proofs Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 15 ++++- .../CSharpUtils.cs | 7 +-- .../VSTHRD002UseJtfRunAnalyzer.cs | 4 +- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 56 +++++++++++++++++++ 4 files changed, 75 insertions(+), 7 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index fc2768a23..f51ebda79 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -207,6 +207,12 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc return false; } + if (context.SemanticModel.GetEnclosingSymbol(memberAccessSyntax.SpanStart, context.CancellationToken) is not IMethodSymbol enclosingMethod + || !SymbolEqualityComparer.Default.Equals(taskSymbol.ContainingSymbol, enclosingMethod)) + { + return false; + } + ImmutableHashSet taskSymbols = GetSymbolAndRefAliases(context, memberAccessSyntax, taskSymbol); if (NestedFunctionMayReassignTask(context, memberAccessSyntax, taskSymbols)) { @@ -479,11 +485,16 @@ private static bool TryGetAwaitExpression( IEnumerable ancestorsWithinStatement = candidate.Ancestors().TakeWhile(node => node != statement); bool isConditionallyExecuted = ancestorsWithinStatement.Any( - node => node is StatementSyntax or ConditionalExpressionSyntax or SwitchExpressionSyntax or ConditionalAccessExpressionSyntax + node => node is StatementSyntax + or ConditionalExpressionSyntax + or SwitchExpressionSyntax + or ConditionalAccessExpressionSyntax + or WhenClauseSyntax || (node is BinaryExpressionSyntax binary && (binary.IsKind(SyntaxKind.LogicalAndExpression) || binary.IsKind(SyntaxKind.LogicalOrExpression) - || binary.IsKind(SyntaxKind.CoalesceExpression)))); + || binary.IsKind(SyntaxKind.CoalesceExpression)) + && binary.Right.FullSpan.Contains(candidate.Span))); if (!isConditionallyExecuted && AwaitCompletesTask(context, candidate, taskSymbol, taskSymbols)) { awaitExpression = candidate; 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 29ea1050e..afdbe4554 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -195,7 +195,9 @@ private static void AnalyzeInvocation( _ => null, }; - if (methodName is object && !CSharpCommonInterest.ShouldIgnoreContext(context)) + 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)); diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 03e7d8dcc..2e5f09821 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -731,6 +731,14 @@ async void AwaitedBeforeOtherTaskInSameStatement() { _ = task.Result; } + async void AwaitedInLeftShortCircuitOperand(bool condition) { + var task = Task.Run(() => true); + if (await task && condition) { + } + + _ = task.Result; + } + async void AwaitedInNestedBlock() { var task = Task.Run(() => 1); { @@ -868,6 +876,18 @@ async void AwaitInDoWhileConditionIsNotDefinite() { _ = 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; @@ -898,6 +918,17 @@ void ReassignedThroughReboundRefAliasInsideGuard(ref Task task, ref Task 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) { } @@ -960,6 +991,31 @@ Task FAsync(Task task) { 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() { From a0b2eb46799108c9a40649e42858dff7e513e55c Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 12:59:40 -0600 Subject: [PATCH 08/26] Complete VSTHRD002 alias analysis Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 109 ++++++++++-------- .../VSTHRD002UseJtfRunAnalyzer.cs | 24 ++-- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 8 +- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 46 ++++++++ ...idJtfRunInNonPublicMembersAnalyzerTests.cs | 18 +++ 5 files changed, 145 insertions(+), 60 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index f51ebda79..086ddf37d 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -59,7 +59,9 @@ or BaseMethodDeclarationSyntax ? refExpression.Expression : variable.Initializer.Value; ISymbol? initializedFrom = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(initializer), context.CancellationToken).Symbol; - if (initializedFrom is object && symbols.Contains(initializedFrom) && symbols.Add(local)) + if (initializedFrom is object + && (symbols.Contains(initializedFrom) || symbols.Contains(local)) + && (symbols.Add(initializedFrom) | symbols.Add(local))) { addedAlias = true; } @@ -74,7 +76,9 @@ or BaseMethodDeclarationSyntax } ISymbol? assignedFrom = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(refExpression.Expression), context.CancellationToken).Symbol; - if (assignedFrom is object && symbols.Contains(assignedFrom) && symbols.Add(local)) + if (assignedFrom is object + && (symbols.Contains(assignedFrom) || symbols.Contains(local)) + && (symbols.Add(assignedFrom) | symbols.Add(local))) { addedAlias = true; } @@ -201,7 +205,15 @@ private static ExpressionSyntax GetTaskReceiver(MemberAccessExpressionSyntax mem private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) { - ISymbol? taskSymbol = context.SemanticModel.GetSymbolInfo(GetTaskReceiver(memberAccessSyntax), context.CancellationToken).Symbol; + ExpressionSyntax taskReceiver = GetTaskReceiver(memberAccessSyntax); + ITypeSymbol? taskType = context.SemanticModel.GetTypeInfo(taskReceiver, context.CancellationToken).Type; + if (!Utils.IsTask(taskType) + && !(taskType?.Name == nameof(ValueTask) && taskType.BelongsToNamespace(Namespaces.SystemThreadingTasks))) + { + return false; + } + + ISymbol? taskSymbol = context.SemanticModel.GetSymbolInfo(taskReceiver, context.CancellationToken).Symbol; if (taskSymbol is not ILocalSymbol and not IParameterSymbol) { return false; @@ -219,7 +231,7 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc return false; } - if (IsWithinCompletedTaskBranch(context, memberAccessSyntax, taskSymbol, taskSymbols)) + if (IsWithinCompletedTaskBranch(context, memberAccessSyntax, taskSymbols)) { return true; } @@ -230,7 +242,7 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc return false; } - if (TryGetAwaitExpression(context, containingStatement, taskSymbol, taskSymbols, memberAccessSyntax.SpanStart, out AwaitExpressionSyntax? precedingAwait) + if (TryGetAwaitExpression(context, containingStatement, taskSymbols, memberAccessSyntax.SpanStart, out AwaitExpressionSyntax? precedingAwait) && !MayReassignTask(context, containingStatement, taskSymbols, precedingAwait.Span.End, memberAccessSyntax.SpanStart)) { return true; @@ -247,7 +259,7 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc for (int i = statementIndex - 1; i >= 0; i--) { StatementSyntax statement = block.Statements[i]; - if (StatementCompletesTask(context, statement, taskSymbol, taskSymbols)) + if (StatementCompletesTask(context, statement, taskSymbols)) { return true; } @@ -278,7 +290,6 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc private static bool IsWithinCompletedTaskBranch( SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax, - ISymbol taskSymbol, IImmutableSet taskSymbols) { foreach (IfStatementSyntax ifStatement in memberAccessSyntax.Ancestors().OfType()) @@ -290,14 +301,14 @@ private static bool IsWithinCompletedTaskBranch( } if (ifStatement.Statement.FullSpan.Contains(memberAccessSyntax.Span) - && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: true) + && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: true) && !MayReassignTask(context, ifStatement, taskSymbols, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) { return true; } if (ifStatement.Else?.Statement.FullSpan.Contains(memberAccessSyntax.Span) is true - && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: false) + && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: false) && !MayReassignTask(context, ifStatement, taskSymbols, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) { return true; @@ -307,52 +318,56 @@ private static bool IsWithinCompletedTaskBranch( return false; } - private static bool ConditionProvesCompletion(SyntaxNodeAnalysisContext context, ExpressionSyntax condition, ISymbol taskSymbol, bool conditionValue) + 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, taskSymbol, !conditionValue); + return ConditionProvesCompletion(context, logicalNot.Operand, taskSymbols, !conditionValue); } if (condition is BinaryExpressionSyntax binary) { if (conditionValue && binary.IsKind(SyntaxKind.LogicalAndExpression)) { - return ConditionProvesCompletion(context, binary.Left, taskSymbol, conditionValue: true) - || ConditionProvesCompletion(context, binary.Right, taskSymbol, conditionValue: true); + 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, taskSymbol, conditionValue: false) - || ConditionProvesCompletion(context, binary.Right, taskSymbol, conditionValue: false); + 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, taskSymbol, conditionValue: true) - && ConditionProvesCompletion(context, binary.Right, taskSymbol, conditionValue: true); + 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, taskSymbol, conditionValue: false) - && ConditionProvesCompletion(context, binary.Right, taskSymbol, conditionValue: false); + 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, taskSymbol)) + && IsRanToCompletionComparison(context, binary.Left, binary.Right, taskSymbols)) { return equalityHolds; } } if (condition is MemberAccessExpressionSyntax completedProperty - && IsSameTask(context, completedProperty.Expression, taskSymbol) + && IsOneOfSymbols(context, completedProperty.Expression, taskSymbols) && completedProperty.Name.Identifier.ValueText is nameof(Task.IsCompleted) or nameof(Task.IsCanceled) or nameof(Task.IsFaulted) @@ -364,15 +379,19 @@ or nameof(Task.IsFaulted) return false; } - private static bool IsRanToCompletionComparison(SyntaxNodeAnalysisContext context, ExpressionSyntax left, ExpressionSyntax right, ISymbol taskSymbol) + private static bool IsRanToCompletionComparison( + SyntaxNodeAnalysisContext context, + ExpressionSyntax left, + ExpressionSyntax right, + IImmutableSet taskSymbols) { - return (IsTaskStatus(context, left, taskSymbol) && IsRanToCompletion(context, right)) - || (IsTaskStatus(context, right, taskSymbol) && IsRanToCompletion(context, left)); + return (IsTaskStatus(context, left, taskSymbols) && IsRanToCompletion(context, right)) + || (IsTaskStatus(context, right, taskSymbols) && IsRanToCompletion(context, left)); } - private static bool IsTaskStatus(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, ISymbol taskSymbol) + private static bool IsTaskStatus(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, IImmutableSet taskSymbols) => expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.Status) } statusAccess - && IsSameTask(context, statusAccess.Expression, taskSymbol); + && IsOneOfSymbols(context, statusAccess.Expression, taskSymbols); private static bool IsRanToCompletion(SyntaxNodeAnalysisContext context, ExpressionSyntax expression) => context.SemanticModel.GetSymbolInfo(expression, context.CancellationToken).Symbol is IFieldSymbol status @@ -383,10 +402,9 @@ private static bool IsRanToCompletion(SyntaxNodeAnalysisContext context, Express private static bool StatementCompletesTask( SyntaxNodeAnalysisContext context, StatementSyntax statement, - ISymbol taskSymbol, IImmutableSet taskSymbols) { - if (TryGetAwaitExpression(context, statement, taskSymbol, taskSymbols, out AwaitExpressionSyntax? awaitExpression) + if (TryGetAwaitExpression(context, statement, taskSymbols, out AwaitExpressionSyntax? awaitExpression) && !MayReassignTask(context, statement, taskSymbols, awaitExpression.Span.End, statement.Span.End + 1)) { return true; @@ -394,7 +412,7 @@ private static bool StatementCompletesTask( if (statement is BlockSyntax block) { - return StatementDefinitelyAwaitsTask(context, block, taskSymbol, taskSymbols); + return StatementDefinitelyAwaitsTask(context, block, taskSymbols); } if (statement is not IfStatementSyntax ifStatement) @@ -409,25 +427,24 @@ private static bool StatementCompletesTask( if (ifStatement.Else is null) { - return ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: false) - && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbol, taskSymbols); + return ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: false) + && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbols); } - return (ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: true) + return (ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: true) && !MayReassignTask(context, ifStatement.Statement, taskSymbols) - && StatementDefinitelyAwaitsTask(context, ifStatement.Else.Statement, taskSymbol, taskSymbols)) - || (ConditionProvesCompletion(context, ifStatement.Condition, taskSymbol, conditionValue: false) - && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbol, taskSymbols) + && StatementDefinitelyAwaitsTask(context, ifStatement.Else.Statement, taskSymbols)) + || (ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: false) + && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbols) && !MayReassignTask(context, ifStatement.Else.Statement, taskSymbols)); } private static bool StatementDefinitelyAwaitsTask( SyntaxNodeAnalysisContext context, StatementSyntax statement, - ISymbol taskSymbol, IImmutableSet taskSymbols) { - if (TryGetAwaitExpression(context, statement, taskSymbol, taskSymbols, out AwaitExpressionSyntax? awaitExpression)) + if (TryGetAwaitExpression(context, statement, taskSymbols, out AwaitExpressionSyntax? awaitExpression)) { return !MayReassignTask(context, statement, taskSymbols, awaitExpression.Span.End, statement.Span.End + 1); } @@ -436,7 +453,7 @@ private static bool StatementDefinitelyAwaitsTask( { for (int i = block.Statements.Count - 1; i >= 0; i--) { - if (TryGetAwaitExpression(context, block.Statements[i], taskSymbol, taskSymbols, out awaitExpression) + if (TryGetAwaitExpression(context, block.Statements[i], taskSymbols, out awaitExpression) && !MayReassignTask(context, block.Statements[i], taskSymbols, awaitExpression.Span.End, block.Statements[i].Span.End + 1)) { return true; @@ -455,15 +472,13 @@ private static bool StatementDefinitelyAwaitsTask( private static bool TryGetAwaitExpression( SyntaxNodeAnalysisContext context, StatementSyntax statement, - ISymbol taskSymbol, IImmutableSet taskSymbols, [NotNullWhen(true)] out AwaitExpressionSyntax? awaitExpression) - => TryGetAwaitExpression(context, statement, taskSymbol, taskSymbols, statement.Span.End + 1, out awaitExpression); + => TryGetAwaitExpression(context, statement, taskSymbols, statement.Span.End + 1, out awaitExpression); private static bool TryGetAwaitExpression( SyntaxNodeAnalysisContext context, StatementSyntax statement, - ISymbol taskSymbol, IImmutableSet taskSymbols, int beforePosition, [NotNullWhen(true)] out AwaitExpressionSyntax? awaitExpression) @@ -495,7 +510,7 @@ or WhenClauseSyntax || binary.IsKind(SyntaxKind.LogicalOrExpression) || binary.IsKind(SyntaxKind.CoalesceExpression)) && binary.Right.FullSpan.Contains(candidate.Span))); - if (!isConditionallyExecuted && AwaitCompletesTask(context, candidate, taskSymbol, taskSymbols)) + if (!isConditionallyExecuted && AwaitCompletesTask(context, candidate, taskSymbols)) { awaitExpression = candidate; return true; @@ -509,7 +524,6 @@ or WhenClauseSyntax private static bool AwaitCompletesTask( SyntaxNodeAnalysisContext context, AwaitExpressionSyntax awaitExpression, - ISymbol taskSymbol, IImmutableSet taskSymbols) { if (MayReassignTask(context, awaitExpression.Expression, taskSymbols)) @@ -524,7 +538,7 @@ private static bool AwaitCompletesTask( awaitedExpression = UnwrapParentheses(configureAwaitAccess.Expression); } - if (IsSameTask(context, awaitedExpression, taskSymbol)) + if (IsOneOfSymbols(context, awaitedExpression, taskSymbols)) { return true; } @@ -535,7 +549,7 @@ private static bool AwaitCompletesTask( && whenAllMethod.ContainingType.Name == nameof(Task) && whenAllMethod.ContainingType.BelongsToNamespace(Namespaces.SystemThreadingTasks)) { - return whenAllInvocation.ArgumentList.Arguments.Any(argument => IsSameTask(context, argument.Expression, taskSymbol)); + return whenAllInvocation.ArgumentList.Arguments.Any(argument => IsOneOfSymbols(context, argument.Expression, taskSymbols)); } return false; @@ -601,11 +615,6 @@ private static bool IsAssignmentToTask(SyntaxNodeAnalysisContext context, Expres return IsOneOfSymbols(context, expression, taskSymbols); } - private static bool IsSameTask(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, ISymbol taskSymbol) - => SymbolEqualityComparer.Default.Equals( - context.SemanticModel.GetSymbolInfo(UnwrapParentheses(expression), context.CancellationToken).Symbol, - taskSymbol); - private static bool IsOneOfSymbols(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, IImmutableSet symbols) { ISymbol? symbol = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(expression), context.CancellationToken).Symbol; diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index afdbe4554..80847626f 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -148,11 +148,16 @@ private static void InspectMemberAccess( ParameterSyntax? firstParameter = GetFirstParameter(anonymousFunctionSyntax); if (firstParameter is object - && context.SemanticModel.GetDeclaredSymbol(firstParameter, context.CancellationToken) is IParameterSymbol completedTask - && SymbolEqualityComparer.Default.Equals(GetTaskReceiverSymbol(context, memberAccessSyntax), completedTask) - && !IsTaskReassignedInContinuation(context, anonymousFunctionSyntax, memberAccessSyntax, completedTask)) + && context.SemanticModel.GetDeclaredSymbol(firstParameter, context.CancellationToken) is IParameterSymbol completedTask) { - return; + ImmutableHashSet taskSymbols = CSharpCommonInterest.GetSymbolAndRefAliases(context, anonymousFunctionSyntax, completedTask); + ISymbol? receiverSymbol = GetTaskReceiverSymbol(context, memberAccessSyntax); + if (receiverSymbol is object + && taskSymbols.Contains(receiverSymbol) + && !IsTaskReassignedInContinuation(context, anonymousFunctionSyntax, memberAccessSyntax, taskSymbols)) + { + return; + } } } @@ -252,15 +257,16 @@ private static bool IsTaskReassignedInContinuation( SyntaxNodeAnalysisContext context, AnonymousFunctionExpressionSyntax continuation, MemberAccessExpressionSyntax memberAccess, - IParameterSymbol taskParameter) + ImmutableHashSet taskSymbols) { bool accessIsNested = memberAccess.Ancestors().TakeWhile(node => node != continuation).Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax); int beforePosition = accessIsNested ? continuation.Span.End + 1 : memberAccess.SpanStart; - ImmutableHashSet taskSymbols = CSharpCommonInterest.GetSymbolAndRefAliases(context, continuation, taskParameter); foreach (AssignmentExpressionSyntax assignment in continuation.DescendantNodes().OfType()) { - if (assignment.SpanStart < beforePosition + bool isDeferredWrite = assignment.Ancestors().TakeWhile(node => node != continuation) + .Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax); + if ((assignment.SpanStart < beforePosition || isDeferredWrite) && IsAssignmentToParameter(context, assignment.Left, taskSymbols)) { return true; @@ -270,7 +276,9 @@ private static bool IsTaskReassignedInContinuation( foreach (ArgumentSyntax argument in continuation.DescendantNodes().OfType()) { ISymbol? argumentSymbol = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(argument.Expression), context.CancellationToken).Symbol; - if (argument.SpanStart < beforePosition + bool isDeferredWrite = argument.Ancestors().TakeWhile(node => node != continuation) + .Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax); + if ((argument.SpanStart < beforePosition || isDeferredWrite) && (argument.RefKindKeyword.IsKind(SyntaxKind.RefKeyword) || argument.RefKindKeyword.IsKind(SyntaxKind.OutKeyword)) && argumentSymbol is object && taskSymbols.Contains(argumentSymbol)) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index 4711e1915..c3a9ac47b 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -131,8 +131,12 @@ private static bool TryFindNodeAtSource(Diagnostic diagnostic, SyntaxNode root, getAwaiterInvocationExpression = parenthesized.Expression; } - var getAwaiterAccess = (getAwaiterInvocationExpression as InvocationExpressionSyntax)?.Expression as MemberAccessExpressionSyntax; - return getAwaiterAccess?.Name.Identifier.ValueText == "GetAwaiter" ? getAwaiterAccess.Expression : null; + 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)) diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 2e5f09821..7bf9e8e93 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -586,6 +586,10 @@ void F() { 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; + }); } void ContinueWith(Func, int> del) { } @@ -664,6 +668,12 @@ void F() { alias = Task.Run(() => 7); return t.[|Result|]; }); + task.ContinueWith(t => { + Replace(); + return t.[|Result|]; + + void Replace() => t = Task.Run(() => 8); + }); } } "; @@ -918,6 +928,21 @@ void ReassignedThroughReboundRefAliasInsideGuard(ref Task task, ref 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(); @@ -1035,6 +1060,27 @@ void F(Task task) { 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 CodeFixIsNotOfferedWhenBlockingExpressionHasCompileErrors() { 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() { From 9e0d607af37f980956ddebc546792ab6f7f75b26 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 13:09:21 -0600 Subject: [PATCH 09/26] Harden completion branches and awaiter fixes Recognize nested conditional branches that definitely await the same task, and suppress invalid await code fixes for unrelated static GetAwaiter factories. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 9 +++-- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 19 +++++++++++ .../VSTHRD002UseJtfRunAnalyzerTests.cs | 34 +++++++++++++++++++ 3 files changed, 60 insertions(+), 2 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index 086ddf37d..feb5cf966 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -449,12 +449,17 @@ private static bool StatementDefinitelyAwaitsTask( return !MayReassignTask(context, statement, taskSymbols, awaitExpression.Span.End, statement.Span.End + 1); } + if (statement is IfStatementSyntax { Else: { } elseClause } ifStatement) + { + return StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbols) + && StatementDefinitelyAwaitsTask(context, elseClause.Statement, taskSymbols); + } + if (statement is BlockSyntax block) { for (int i = block.Statements.Count - 1; i >= 0; i--) { - if (TryGetAwaitExpression(context, block.Statements[i], taskSymbols, out awaitExpression) - && !MayReassignTask(context, block.Statements[i], taskSymbols, awaitExpression.Span.End, block.Statements[i].Span.End + 1)) + if (StatementDefinitelyAwaitsTask(context, block.Statements[i], taskSymbols)) { return true; } diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index c3a9ac47b..87e9e8bc0 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -198,6 +198,25 @@ private static bool CanUseAwaitCodeFix(SemanticModel semanticModel, ExpressionSy return method.Parameters.IsEmpty; } + if (method.Name == nameof(TaskAwaiter.GetResult)) + { + ExpressionSyntax? getAwaiterExpression = (invocation.Expression as MemberAccessExpressionSyntax)?.Expression; + while (getAwaiterExpression is ParenthesizedExpressionSyntax parenthesized) + { + getAwaiterExpression = parenthesized.Expression; + } + + ExpressionSyntax? receiver = ((getAwaiterExpression as InvocationExpressionSyntax)?.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)) { diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 7bf9e8e93..9e7cd59c4 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -777,6 +777,19 @@ async void GuardedWithElse() { _ = task.Result; } + async void GuardedWithNestedConditional(bool condition) { + var task = Task.Run(() => 1); + if (!task.IsCompleted) { + if (condition) { + await task; + } else { + await task; + } + } + + _ = task.Result; + } + async void AwaitedBeforeNestedBlock(bool condition) { var task = Task.Run(() => 1); await task; @@ -1081,6 +1094,27 @@ class CustomAwaitable { await CSVerify.VerifyCodeFixAsync(test, test); } + [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() { From 38bd17cd6a9c4826e60da4fe00c6fb040b2b1e1a Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 13:36:16 -0600 Subject: [PATCH 10/26] Make VSTHRD002 alias analysis flow-aware Separate definite ref aliases used for completion proofs from potential aliases used for reassignment invalidation, and include closure writes from accessors and top-level statements. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 177 ++++++++++++------ .../VSTHRD002UseJtfRunAnalyzer.cs | 5 +- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 64 +++++++ 3 files changed, 182 insertions(+), 64 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index feb5cf966..14fa02a61 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -30,63 +30,108 @@ internal static class CSharpCommonInterest SyntaxKind.RemoveAccessorDeclaration); /// - /// Gets a symbol and any ref locals that alias it within the containing function. + /// Gets a symbol and ref locals that definitely or potentially alias it at the specified syntax node. /// - internal static ImmutableHashSet GetSymbolAndRefAliases(SyntaxNodeAnalysisContext context, SyntaxNode node, ISymbol symbol) + internal static (ImmutableHashSet Definite, ImmutableHashSet Potential) GetSymbolAndRefAliases( + SyntaxNodeAnalysisContext context, + SyntaxNode node, + ISymbol symbol) { - ImmutableHashSet.Builder symbols = ImmutableHashSet.CreateBuilder(SymbolEqualityComparer.Default); - symbols.Add(symbol); - SyntaxNode searchRoot = node.AncestorsAndSelf().FirstOrDefault( ancestor => ancestor is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax or BaseMethodDeclarationSyntax - or AccessorDeclarationSyntax) ?? node; + or AccessorDeclarationSyntax) + ?? node.FirstAncestorOrSelf()?.Parent + ?? node; + + var refTargets = new Dictionary>(SymbolEqualityComparer.Default); + bool DescendIntoChildren(SyntaxNode child) => + child == searchRoot || child is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; - bool addedAlias; - do + HashSet GetRefTargets(ISymbol candidate) { - addedAlias = false; - foreach (VariableDeclaratorSyntax variable in searchRoot.DescendantNodes().OfType()) + if (refTargets.TryGetValue(candidate, out HashSet? targets)) { - if (variable.Initializer is null - || context.SemanticModel.GetDeclaredSymbol(variable, context.CancellationToken) is not ILocalSymbol { RefKind: not RefKind.None } local) - { - continue; - } + return new HashSet(targets, SymbolEqualityComparer.Default); + } + + return new HashSet(SymbolEqualityComparer.Default) { candidate }; + } - ExpressionSyntax initializer = variable.Initializer.Value is RefExpressionSyntax refExpression - ? refExpression.Expression + 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 => 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; - ISymbol? initializedFrom = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(initializer), context.CancellationToken).Symbol; - if (initializedFrom is object - && (symbols.Contains(initializedFrom) || symbols.Contains(local)) - && (symbols.Add(initializedFrom) | symbols.Add(local))) + if (context.SemanticModel.GetSymbolInfo(UnwrapParentheses(initializer), context.CancellationToken).Symbol is ISymbol initializedFrom) { - addedAlias = true; + refTargets[local] = GetRefTargets(initializedFrom); } } - - foreach (AssignmentExpressionSyntax assignment in searchRoot.DescendantNodes().OfType()) + else if (candidate is AssignmentExpressionSyntax { Right: RefExpressionSyntax refAssignment } assignment + && context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol is ILocalSymbol { RefKind: not RefKind.None } reboundLocal + && context.SemanticModel.GetSymbolInfo(UnwrapParentheses(refAssignment.Expression), context.CancellationToken).Symbol is ISymbol assignedFrom) { - if (assignment.Right is not RefExpressionSyntax refExpression - || context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol is not ILocalSymbol { RefKind: not RefKind.None } local) + HashSet assignedTargets = GetRefTargets(assignedFrom); + if (DefinitelyPrecedesNode(candidate) || !refTargets.TryGetValue(reboundLocal, out HashSet? existingTargets)) { - continue; + refTargets[reboundLocal] = assignedTargets; } - - ISymbol? assignedFrom = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(refExpression.Expression), context.CancellationToken).Symbol; - if (assignedFrom is object - && (symbols.Contains(assignedFrom) || symbols.Contains(local)) - && (symbols.Add(assignedFrom) | symbols.Add(local))) + else { - addedAlias = true; + existingTargets.UnionWith(assignedTargets); } } } - while (addedAlias); - return symbols.ToImmutable(); + 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)) + { + definiteSymbols.Add(refTarget.Key); + } + + if (refTarget.Value.Overlaps(symbolTargets)) + { + potentialSymbols.Add(refTarget.Key); + } + } + + return (definiteSymbols.ToImmutable(), potentialSymbols.ToImmutable()); } /// @@ -225,13 +270,14 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc return false; } - ImmutableHashSet taskSymbols = GetSymbolAndRefAliases(context, memberAccessSyntax, taskSymbol); - if (NestedFunctionMayReassignTask(context, memberAccessSyntax, taskSymbols)) + (ImmutableHashSet taskSymbols, ImmutableHashSet potentialTaskSymbols) = + GetSymbolAndRefAliases(context, memberAccessSyntax, taskSymbol); + if (NestedFunctionMayReassignTask(context, memberAccessSyntax, potentialTaskSymbols)) { return false; } - if (IsWithinCompletedTaskBranch(context, memberAccessSyntax, taskSymbols)) + if (IsWithinCompletedTaskBranch(context, memberAccessSyntax, taskSymbols, potentialTaskSymbols)) { return true; } @@ -243,14 +289,14 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc } if (TryGetAwaitExpression(context, containingStatement, taskSymbols, memberAccessSyntax.SpanStart, out AwaitExpressionSyntax? precedingAwait) - && !MayReassignTask(context, containingStatement, taskSymbols, precedingAwait.Span.End, memberAccessSyntax.SpanStart)) + && !MayReassignTask(context, containingStatement, potentialTaskSymbols, precedingAwait.Span.End, memberAccessSyntax.SpanStart)) { return true; } while (containingStatement.Parent is BlockSyntax block) { - if (MayReassignTask(context, containingStatement, taskSymbols, containingStatement.SpanStart - 1, memberAccessSyntax.SpanStart)) + if (MayReassignTask(context, containingStatement, potentialTaskSymbols, containingStatement.SpanStart - 1, memberAccessSyntax.SpanStart)) { return false; } @@ -259,12 +305,12 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc for (int i = statementIndex - 1; i >= 0; i--) { StatementSyntax statement = block.Statements[i]; - if (StatementCompletesTask(context, statement, taskSymbols)) + if (StatementCompletesTask(context, statement, taskSymbols, potentialTaskSymbols)) { return true; } - if (MayReassignTask(context, statement, taskSymbols)) + if (MayReassignTask(context, statement, potentialTaskSymbols)) { return false; } @@ -276,7 +322,7 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc return false; } - if (MayReassignTask(context, outerStatement, taskSymbols, outerStatement.SpanStart - 1, containingStatement.SpanStart)) + if (MayReassignTask(context, outerStatement, potentialTaskSymbols, outerStatement.SpanStart - 1, containingStatement.SpanStart)) { return false; } @@ -290,7 +336,8 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc private static bool IsWithinCompletedTaskBranch( SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax, - IImmutableSet taskSymbols) + IImmutableSet taskSymbols, + IImmutableSet potentialTaskSymbols) { foreach (IfStatementSyntax ifStatement in memberAccessSyntax.Ancestors().OfType()) { @@ -302,14 +349,14 @@ private static bool IsWithinCompletedTaskBranch( if (ifStatement.Statement.FullSpan.Contains(memberAccessSyntax.Span) && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: true) - && !MayReassignTask(context, ifStatement, taskSymbols, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) + && !MayReassignTask(context, ifStatement, potentialTaskSymbols, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) { return true; } if (ifStatement.Else?.Statement.FullSpan.Contains(memberAccessSyntax.Span) is true && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: false) - && !MayReassignTask(context, ifStatement, taskSymbols, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) + && !MayReassignTask(context, ifStatement, potentialTaskSymbols, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) { return true; } @@ -402,17 +449,18 @@ private static bool IsRanToCompletion(SyntaxNodeAnalysisContext context, Express private static bool StatementCompletesTask( SyntaxNodeAnalysisContext context, StatementSyntax statement, - IImmutableSet taskSymbols) + IImmutableSet taskSymbols, + IImmutableSet potentialTaskSymbols) { if (TryGetAwaitExpression(context, statement, taskSymbols, out AwaitExpressionSyntax? awaitExpression) - && !MayReassignTask(context, statement, taskSymbols, awaitExpression.Span.End, statement.Span.End + 1)) + && !MayReassignTask(context, statement, potentialTaskSymbols, awaitExpression.Span.End, statement.Span.End + 1)) { return true; } if (statement is BlockSyntax block) { - return StatementDefinitelyAwaitsTask(context, block, taskSymbols); + return StatementDefinitelyAwaitsTask(context, block, taskSymbols, potentialTaskSymbols); } if (statement is not IfStatementSyntax ifStatement) @@ -420,7 +468,7 @@ private static bool StatementCompletesTask( return false; } - if (MayReassignTask(context, ifStatement.Condition, taskSymbols)) + if (MayReassignTask(context, ifStatement.Condition, potentialTaskSymbols)) { return false; } @@ -428,43 +476,44 @@ private static bool StatementCompletesTask( if (ifStatement.Else is null) { return ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: false) - && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbols); + && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbols, potentialTaskSymbols); } return (ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: true) - && !MayReassignTask(context, ifStatement.Statement, taskSymbols) - && StatementDefinitelyAwaitsTask(context, ifStatement.Else.Statement, taskSymbols)) + && !MayReassignTask(context, ifStatement.Statement, potentialTaskSymbols) + && StatementDefinitelyAwaitsTask(context, ifStatement.Else.Statement, taskSymbols, potentialTaskSymbols)) || (ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: false) - && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbols) - && !MayReassignTask(context, ifStatement.Else.Statement, taskSymbols)); + && StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbols, potentialTaskSymbols) + && !MayReassignTask(context, ifStatement.Else.Statement, potentialTaskSymbols)); } private static bool StatementDefinitelyAwaitsTask( SyntaxNodeAnalysisContext context, StatementSyntax statement, - IImmutableSet taskSymbols) + IImmutableSet taskSymbols, + IImmutableSet potentialTaskSymbols) { if (TryGetAwaitExpression(context, statement, taskSymbols, out AwaitExpressionSyntax? awaitExpression)) { - return !MayReassignTask(context, statement, taskSymbols, awaitExpression.Span.End, statement.Span.End + 1); + return !MayReassignTask(context, statement, potentialTaskSymbols, awaitExpression.Span.End, statement.Span.End + 1); } if (statement is IfStatementSyntax { Else: { } elseClause } ifStatement) { - return StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbols) - && StatementDefinitelyAwaitsTask(context, elseClause.Statement, taskSymbols); + return 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)) + if (StatementDefinitelyAwaitsTask(context, block.Statements[i], taskSymbols, potentialTaskSymbols)) { return true; } - if (MayReassignTask(context, block.Statements[i], taskSymbols)) + if (MayReassignTask(context, block.Statements[i], potentialTaskSymbols)) { return false; } @@ -603,7 +652,11 @@ private static bool NestedFunctionMayReassignTask( IImmutableSet taskSymbols) { SyntaxNode? containingFunction = node.Ancestors().FirstOrDefault( - ancestor => ancestor is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax or BaseMethodDeclarationSyntax); + ancestor => ancestor is AnonymousFunctionExpressionSyntax + or LocalFunctionStatementSyntax + or BaseMethodDeclarationSyntax + or AccessorDeclarationSyntax) + ?? node.FirstAncestorOrSelf()?.Parent; return containingFunction?.DescendantNodes() .Where(descendant => descendant is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax) .Any(nestedFunction => MayReassignTask(context, nestedFunction, taskSymbols)) is true; diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index 80847626f..e824597fd 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -150,11 +150,12 @@ private static void InspectMemberAccess( if (firstParameter is object && context.SemanticModel.GetDeclaredSymbol(firstParameter, context.CancellationToken) is IParameterSymbol completedTask) { - ImmutableHashSet taskSymbols = CSharpCommonInterest.GetSymbolAndRefAliases(context, anonymousFunctionSyntax, completedTask); + (ImmutableHashSet taskSymbols, ImmutableHashSet potentialTaskSymbols) = + CSharpCommonInterest.GetSymbolAndRefAliases(context, memberAccessSyntax, completedTask); ISymbol? receiverSymbol = GetTaskReceiverSymbol(context, memberAccessSyntax); if (receiverSymbol is object && taskSymbols.Contains(receiverSymbol) - && !IsTaskReassignedInContinuation(context, anonymousFunctionSyntax, memberAccessSyntax, taskSymbols)) + && !IsTaskReassignedInContinuation(context, anonymousFunctionSyntax, memberAccessSyntax, potentialTaskSymbols)) { return; } diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 9e7cd59c4..8c410bf5f 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -641,6 +641,12 @@ void F() { t = Task.Run(() => 6); useResultLater(); }); + task.ContinueWith(t => { + Task other = Task.Run(() => 7); + ref Task alias = ref t; + alias = ref other; + return alias.[|Result|]; + }); } } "; @@ -890,6 +896,39 @@ async void ReassignedByClosureDeclaredBeforeAwait() { _ = 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; + } + } + async void AwaitInDoWhileConditionIsNotDefinite() { var task = Task.Run(() => true); do { @@ -986,6 +1025,31 @@ bool ReplaceFlag(ref Task task) { }.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() { From 2795a2beb0db4dcca2267a5468e25d4d8e2f9198 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 13:56:51 -0600 Subject: [PATCH 11/26] Harden deferred VSTHRD002 analysis Include continuation-scope aliases in deferred-write checks, validate awaiter chains semantically, preserve closure ordering, recognize WhenAll arrays, and avoid configured-method duplicates with VSTHRD103. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 78 ++++++++++++++++--- .../VSTHRD002UseJtfRunAnalyzer.cs | 59 ++++++++++---- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 18 ++++- .../MultiAnalyzerTests.cs | 32 ++++++++ .../VSTHRD002UseJtfRunAnalyzerTests.cs | 75 ++++++++++++++++++ 5 files changed, 231 insertions(+), 31 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index 14fa02a61..b01b93798 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -35,9 +35,11 @@ internal static class CSharpCommonInterest internal static (ImmutableHashSet Definite, ImmutableHashSet Potential) GetSymbolAndRefAliases( SyntaxNodeAnalysisContext context, SyntaxNode node, - ISymbol symbol) + ISymbol symbol, + SyntaxNode? aliasSearchRoot = null, + bool includeAllCandidates = false) { - SyntaxNode searchRoot = node.AncestorsAndSelf().FirstOrDefault( + SyntaxNode searchRoot = aliasSearchRoot ?? node.AncestorsAndSelf().FirstOrDefault( ancestor => ancestor is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax or BaseMethodDeclarationSyntax @@ -76,7 +78,8 @@ bool DefinitelyPrecedesNode(SyntaxNode candidate) } foreach (SyntaxNode candidate in searchRoot.DescendantNodes(DescendIntoChildren) - .Where(candidate => candidate.SpanStart < node.SpanStart && candidate is VariableDeclaratorSyntax or AssignmentExpressionSyntax) + .Where(candidate => (includeAllCandidates || candidate.SpanStart < node.SpanStart) + && candidate is VariableDeclaratorSyntax or AssignmentExpressionSyntax) .OrderBy(candidate => candidate.SpanStart)) { if (candidate is VariableDeclaratorSyntax variable @@ -96,7 +99,8 @@ bool DefinitelyPrecedesNode(SyntaxNode candidate) && context.SemanticModel.GetSymbolInfo(UnwrapParentheses(refAssignment.Expression), context.CancellationToken).Symbol is ISymbol assignedFrom) { HashSet assignedTargets = GetRefTargets(assignedFrom); - if (DefinitelyPrecedesNode(candidate) || !refTargets.TryGetValue(reboundLocal, out HashSet? existingTargets)) + if ((!includeAllCandidates && DefinitelyPrecedesNode(candidate)) + || !refTargets.TryGetValue(reboundLocal, out HashSet? existingTargets)) { refTargets[reboundLocal] = assignedTargets; } @@ -220,6 +224,12 @@ internal static void InspectMemberAccess( } } + internal static ISymbol? GetTaskReceiverSymbol(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) + { + ExpressionSyntax receiver = GetTaskReceiver(context, memberAccessSyntax); + return context.SemanticModel.GetSymbolInfo(receiver, context.CancellationToken).Symbol; + } + private static ExpressionSyntax UnwrapParentheses(ExpressionSyntax expression) { while (expression is ParenthesizedExpressionSyntax parenthesized) @@ -230,17 +240,19 @@ private static ExpressionSyntax UnwrapParentheses(ExpressionSyntax expression) return expression; } - private static ExpressionSyntax GetTaskReceiver(MemberAccessExpressionSyntax memberAccessSyntax) + 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) + && 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) + && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess + && IsSupportedConfigureAwaitInvocation(context, configureAwaitInvocation)) { receiver = UnwrapParentheses(configureAwaitAccess.Expression); } @@ -248,12 +260,47 @@ private static ExpressionSyntax GetTaskReceiver(MemberAccessExpressionSyntax mem 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; + } + + if (IsTaskLike(method.ContainingType)) + { + return true; + } + + receiver = UnwrapParentheses(receiver); + return receiver is InvocationExpressionSyntax configureAwaitInvocation + && IsSupportedConfigureAwaitInvocation(context, configureAwaitInvocation); + } + + 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 HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) { - ExpressionSyntax taskReceiver = GetTaskReceiver(memberAccessSyntax); + ExpressionSyntax taskReceiver = GetTaskReceiver(context, memberAccessSyntax); ITypeSymbol? taskType = context.SemanticModel.GetTypeInfo(taskReceiver, context.CancellationToken).Type; - if (!Utils.IsTask(taskType) - && !(taskType?.Name == nameof(ValueTask) && taskType.BelongsToNamespace(Namespaces.SystemThreadingTasks))) + if (!IsTaskLike(taskType)) { return false; } @@ -559,6 +606,7 @@ or ConditionalExpressionSyntax or SwitchExpressionSyntax or ConditionalAccessExpressionSyntax or WhenClauseSyntax + or CatchFilterClauseSyntax || (node is BinaryExpressionSyntax binary && (binary.IsKind(SyntaxKind.LogicalAndExpression) || binary.IsKind(SyntaxKind.LogicalOrExpression) @@ -603,7 +651,12 @@ private static bool AwaitCompletesTask( && whenAllMethod.ContainingType.Name == nameof(Task) && whenAllMethod.ContainingType.BelongsToNamespace(Namespaces.SystemThreadingTasks)) { - return whenAllInvocation.ArgumentList.Arguments.Any(argument => IsOneOfSymbols(context, argument.Expression, taskSymbols)); + 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)))); } return false; @@ -658,7 +711,8 @@ or BaseMethodDeclarationSyntax or AccessorDeclarationSyntax) ?? node.FirstAncestorOrSelf()?.Parent; return containingFunction?.DescendantNodes() - .Where(descendant => descendant is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax) + .Where(descendant => descendant is LocalFunctionStatementSyntax + || (descendant is AnonymousFunctionExpressionSyntax && descendant.SpanStart < node.SpanStart)) .Any(nestedFunction => MayReassignTask(context, nestedFunction, taskSymbols)) is true; } diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index e824597fd..dde781a2e 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -152,6 +152,17 @@ private static void InspectMemberAccess( { (ImmutableHashSet taskSymbols, ImmutableHashSet potentialTaskSymbols) = CSharpCommonInterest.GetSymbolAndRefAliases(context, memberAccessSyntax, completedTask); + if (memberAccessSyntax.Ancestors().TakeWhile(node => node != anonymousFunctionSyntax) + .Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax)) + { + potentialTaskSymbols = CSharpCommonInterest.GetSymbolAndRefAliases( + context, + memberAccessSyntax, + completedTask, + anonymousFunctionSyntax, + includeAllCandidates: true).Potential; + } + ISymbol? receiverSymbol = GetTaskReceiverSymbol(context, memberAccessSyntax); if (receiverSymbol is object && taskSymbols.Contains(receiverSymbol) @@ -190,7 +201,10 @@ private static void AnalyzeInvocation( IMethodSymbol methodDefinition = invokedMethod.ReducedFrom ?? invokedMethod; bool isBuiltInSyncBlockingMethod = CommonInterest.ProblematicSyncBlockingMethods.Any( method => method.Method.IsMatch(invokedMethod) || method.Method.IsMatch(methodDefinition)); + bool coveredByVSTHRD103 = !ShouldAnalyze(context, analyzeWholeCodeBlock) + && HasAsyncAlternative(context, invocationExpressionSyntax, invokedMethod); if (!isBuiltInSyncBlockingMethod + && !coveredByVSTHRD103 && configuredSyncBlockingMethods.Any(method => method.IsMatch(invokedMethod) || method.IsMatch(methodDefinition))) { SimpleNameSyntax? methodName = invocationExpressionSyntax.Expression switch @@ -211,6 +225,34 @@ private static void AnalyzeInvocation( } } + private static bool HasAsyncAlternative( + SyntaxNodeAnalysisContext context, + InvocationExpressionSyntax invocation, + IMethodSymbol invokedMethod) + { + string asyncMethodName = invokedMethod.Name + VSTHRD200UseAsyncNamingConventionAnalyzer.MandatoryAsyncSuffix; + INamespaceOrTypeSymbol lookupContainer = invokedMethod.ReducedFrom is { Parameters.Length: > 0 } reducedFrom + ? (INamespaceOrTypeSymbol)reducedFrom.Parameters[0].Type + : invokedMethod.ContainingType; + 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() + && HasSupersetOfParameterTypes(candidate, invokedMethod)); + } + + private static bool HasSupersetOfParameterTypes(IMethodSymbol candidateMethod, IMethodSymbol baselineMethod) + => baselineMethod.Parameters.Length <= candidateMethod.Parameters.Length + && baselineMethod.Parameters.All( + baselineParameter => candidateMethod.Parameters.Any( + candidateParameter => SymbolEqualityComparer.Default.Equals(baselineParameter.Type, candidateParameter.Type))); + private static void AnalyzeMemberAccess(SyntaxNodeAnalysisContext context, INamedTypeSymbol taskSymbol, bool analyzeWholeCodeBlock) { if (!ShouldAnalyze(context, analyzeWholeCodeBlock)) @@ -227,22 +269,7 @@ private static void AnalyzeMemberAccess(SyntaxNodeAnalysisContext context, IName } private static ISymbol? GetTaskReceiverSymbol(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) - { - ExpressionSyntax receiver = UnwrapParentheses(memberAccessSyntax.Expression); - if (receiver is InvocationExpressionSyntax getAwaiterInvocation - && getAwaiterInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: "GetAwaiter" } getAwaiterAccess) - { - receiver = UnwrapParentheses(getAwaiterAccess.Expression); - } - - if (receiver is InvocationExpressionSyntax configureAwaitInvocation - && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess) - { - receiver = UnwrapParentheses(configureAwaitAccess.Expression); - } - - return context.SemanticModel.GetSymbolInfo(receiver, context.CancellationToken).Symbol; - } + => CSharpCommonInterest.GetTaskReceiverSymbol(context, memberAccessSyntax); private static ExpressionSyntax UnwrapParentheses(ExpressionSyntax expression) { diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index 87e9e8bc0..46e7310bd 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -193,9 +193,12 @@ private static bool CanUseAwaitCodeFix(SemanticModel semanticModel, ExpressionSy return true; } - if (method.Name == nameof(Task.Wait) && Utils.IsTask(method.ContainingType)) + if (method.Name == nameof(Task.Wait)) { - return method.Parameters.IsEmpty; + return method.ReducedFrom is null + && !method.IsStatic + && Utils.IsTask(method.ContainingType) + && method.Parameters.IsEmpty; } if (method.Name == nameof(TaskAwaiter.GetResult)) @@ -206,7 +209,16 @@ private static bool CanUseAwaitCodeFix(SemanticModel semanticModel, ExpressionSy getAwaiterExpression = parenthesized.Expression; } - ExpressionSyntax? receiver = ((getAwaiterExpression as InvocationExpressionSyntax)?.Expression as MemberAccessExpressionSyntax)?.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) { diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs index 61d2f3cda..a27e9f04b 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs @@ -62,6 +62,38 @@ 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(); + } + /// /// 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 8c410bf5f..23039bcd1 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -647,6 +647,12 @@ void F() { alias = ref other; return alias.[|Result|]; }); + task.ContinueWith(t => { + ref Task alias = ref t; + alias = Task.Run(() => 8); + Func useResultLater = () => t.[|Result|]; + return useResultLater(); + }); } } "; @@ -747,6 +753,12 @@ async void AwaitedBeforeOtherTaskInSameStatement() { _ = task.Result; } + async void AwaitedInWhenAllArray() { + var task = Task.Run(() => 1); + await Task.WhenAll(new[] { task }); + _ = task.Result; + } + async void AwaitedInLeftShortCircuitOperand(bool condition) { var task = Task.Run(() => true); if (await task && condition) { @@ -929,6 +941,23 @@ int ReassignedByClosureInGetter { } } + 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() => task = Task.Run(() => 5); + } + async void AwaitInDoWhileConditionIsNotDefinite() { var task = Task.Run(() => true); do { @@ -1158,6 +1187,52 @@ class CustomAwaitable { 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 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) { } +} +"; + + await CSVerify.VerifyCodeFixAsync(test, test); + } + [Fact] public async Task StaticGetAwaiterFactoryDoesNotOfferCodeFix() { From 4dce6ff4ccb242ba4d5ea42fcd9c20100ede3282 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 14:20:15 -0600 Subject: [PATCH 12/26] Respect VSTHRD103 exclusions in VSTHRD002 Keep configured blockers diagnosed when VSTHRD103 is excluded, and harden completion proofs for guarded conditions, nested closures, nested guards, conditional awaits, and custom ConfigureAwait methods. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 52 ++++++++++++++++-- .../VSTHRD002UseJtfRunAnalyzer.cs | 17 ++++-- .../MultiAnalyzerTests.cs | 32 +++++++++++ .../VSTHRD002UseJtfRunAnalyzerTests.cs | 55 ++++++++++++++++++- 4 files changed, 146 insertions(+), 10 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index b01b93798..70e1067c1 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -224,6 +224,9 @@ internal static void InspectMemberAccess( } } + /// + /// 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); @@ -394,6 +397,22 @@ private static bool IsWithinCompletedTaskBranch( continue; } + foreach (BinaryExpressionSyntax binary in memberAccessSyntax.Ancestors() + .TakeWhile(node => node != ifStatement) + .OfType()) + { + bool leftProvesCompletion = binary.Right.FullSpan.Contains(memberAccessSyntax.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.Span.End, memberAccessSyntax.SpanStart)) + { + return true; + } + } + if (ifStatement.Statement.FullSpan.Contains(memberAccessSyntax.Span) && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: true) && !MayReassignTask(context, ifStatement, potentialTaskSymbols, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) @@ -545,9 +564,15 @@ private static bool StatementDefinitelyAwaitsTask( return !MayReassignTask(context, statement, potentialTaskSymbols, awaitExpression.Span.End, statement.Span.End + 1); } - if (statement is IfStatementSyntax { Else: { } elseClause } ifStatement) + if (statement is IfStatementSyntax ifStatement) { - return StatementDefinitelyAwaitsTask(context, ifStatement.Statement, taskSymbols, potentialTaskSymbols) + 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); } @@ -602,11 +627,12 @@ private static bool TryGetAwaitExpression( IEnumerable ancestorsWithinStatement = candidate.Ancestors().TakeWhile(node => node != statement); bool isConditionallyExecuted = ancestorsWithinStatement.Any( node => node is StatementSyntax - or ConditionalExpressionSyntax or SwitchExpressionSyntax or ConditionalAccessExpressionSyntax or WhenClauseSyntax or CatchFilterClauseSyntax + || (node is ConditionalExpressionSyntax conditional + && !conditional.Condition.FullSpan.Contains(candidate.Span)) || (node is BinaryExpressionSyntax binary && (binary.IsKind(SyntaxKind.LogicalAndExpression) || binary.IsKind(SyntaxKind.LogicalOrExpression) @@ -635,7 +661,8 @@ private static bool AwaitCompletesTask( ExpressionSyntax awaitedExpression = UnwrapParentheses(awaitExpression.Expression); if (awaitedExpression is InvocationExpressionSyntax configureAwaitInvocation - && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess) + && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess + && IsSupportedConfigureAwaitInvocation(context, configureAwaitInvocation)) { awaitedExpression = UnwrapParentheses(configureAwaitAccess.Expression); } @@ -713,7 +740,22 @@ or BaseMethodDeclarationSyntax return containingFunction?.DescendantNodes() .Where(descendant => descendant is LocalFunctionStatementSyntax || (descendant is AnonymousFunctionExpressionSyntax && descendant.SpanStart < node.SpanStart)) - .Any(nestedFunction => MayReassignTask(context, nestedFunction, taskSymbols)) is true; + .Any(nestedFunction => + { + ImmutableHashSet.Builder nestedTaskSymbols = ImmutableHashSet.CreateBuilder(SymbolEqualityComparer.Default); + nestedTaskSymbols.UnionWith(taskSymbols); + foreach (ISymbol taskSymbol in taskSymbols) + { + nestedTaskSymbols.UnionWith(GetSymbolAndRefAliases( + context, + nestedFunction, + taskSymbol, + nestedFunction, + includeAllCandidates: true).Potential); + } + + return MayReassignTask(context, nestedFunction, nestedTaskSymbols.ToImmutable()); + }) is true; } private static bool IsAssignmentToTask(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, IImmutableSet taskSymbols) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index dde781a2e..92d72b2f3 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -66,6 +66,10 @@ public override void Initialize(AnalysisContext context) 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) { compilationContext.RegisterCodeBlockStartAction(codeBlockContext => @@ -75,7 +79,9 @@ 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, configuredSyncBlockingMethods, analyzeWholeCodeBlock)), SyntaxKind.InvocationExpression); + codeBlockContext.RegisterSyntaxNodeAction( + Utils.DebuggableWrapper(c => AnalyzeInvocation(c, taskSymbol, configuredSyncBlockingMethods, methodsExcludedFromVSTHRD103, analyzeWholeCodeBlock)), + SyntaxKind.InvocationExpression); codeBlockContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(c => AnalyzeMemberAccess(c, taskSymbol, analyzeWholeCodeBlock)), SyntaxKind.SimpleMemberAccessExpression); } }); @@ -155,12 +161,12 @@ private static void InspectMemberAccess( if (memberAccessSyntax.Ancestors().TakeWhile(node => node != anonymousFunctionSyntax) .Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax)) { - potentialTaskSymbols = CSharpCommonInterest.GetSymbolAndRefAliases( + potentialTaskSymbols = potentialTaskSymbols.Union(CSharpCommonInterest.GetSymbolAndRefAliases( context, memberAccessSyntax, completedTask, anonymousFunctionSyntax, - includeAllCandidates: true).Potential; + includeAllCandidates: true).Potential); } ISymbol? receiverSymbol = GetTaskReceiverSymbol(context, memberAccessSyntax); @@ -180,6 +186,7 @@ private static void AnalyzeInvocation( SyntaxNodeAnalysisContext context, INamedTypeSymbol taskSymbol, ImmutableArray configuredSyncBlockingMethods, + ImmutableArray methodsExcludedFromVSTHRD103, bool analyzeWholeCodeBlock) { var invocationExpressionSyntax = (InvocationExpressionSyntax)context.Node; @@ -201,7 +208,9 @@ private static void AnalyzeInvocation( IMethodSymbol methodDefinition = invokedMethod.ReducedFrom ?? invokedMethod; bool isBuiltInSyncBlockingMethod = CommonInterest.ProblematicSyncBlockingMethods.Any( method => method.Method.IsMatch(invokedMethod) || method.Method.IsMatch(methodDefinition)); - bool coveredByVSTHRD103 = !ShouldAnalyze(context, analyzeWholeCodeBlock) + bool coveredByVSTHRD103 = !methodsExcludedFromVSTHRD103.Contains(invokedMethod) + && !methodsExcludedFromVSTHRD103.Contains(methodDefinition) + && !ShouldAnalyze(context, analyzeWholeCodeBlock) && HasAsyncAlternative(context, invocationExpressionSyntax, invokedMethod); if (!isBuiltInSyncBlockingMethod && !coveredByVSTHRD103 diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs index a27e9f04b..36f078377 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs @@ -94,6 +94,38 @@ Task FAsync(CustomWaiter waiter) { 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(); + } + /// /// 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 23039bcd1..8f0dee377 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -653,6 +653,14 @@ void F() { 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(); + }); } } "; @@ -759,6 +767,12 @@ async void AwaitedInWhenAllArray() { _ = task.Result; } + async void AwaitedInConditionalCondition() { + var task = Task.Run(() => true); + _ = (await task) ? 1 : 0; + _ = task.Result; + } + async void AwaitedInLeftShortCircuitOperand(bool condition) { var task = Task.Run(() => true); if (await task && condition) { @@ -808,6 +822,17 @@ async void GuardedWithNestedConditional(bool condition) { _ = task.Result; } + 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; @@ -840,6 +865,9 @@ void CompletionProperties(Task task) { if (task.Status == TaskStatus.RanToCompletion) { _ = task.Result; } + + if (task.IsCompleted && task.Result == 1) { + } } async void ConditionalAwait(bool condition) { @@ -955,7 +983,10 @@ void LaterLocalFunctionStillInvalidatesEarlierGuard(Task task) { _ = task.[|Result|]; } - void Replace() => task = Task.Run(() => 5); + void Replace() { + ref Task alias = ref task; + alias = Task.Run(() => 5); + } } async void AwaitInDoWhileConditionIsNotDefinite() { @@ -1213,6 +1244,28 @@ public static TaskAwaiter GetAwaiter(this Task task, int mode) 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() { From 0bbc37a330973ac994ab21a860f05ffc9c34eb4c Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 14:39:49 -0600 Subject: [PATCH 13/26] Harden VSTHRD002 review edge cases Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 73 ++++++++++++++----- .../VSTHRD002UseJtfRunAnalyzer.cs | 34 ++++++++- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 3 +- .../MultiAnalyzerTests.cs | 35 +++++++++ .../VSTHRD002UseJtfRunAnalyzerTests.cs | 59 ++++++++++++++- 5 files changed, 179 insertions(+), 25 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index 70e1067c1..e64c0d1e9 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -204,13 +204,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)) { @@ -243,6 +247,27 @@ private static ExpressionSyntax UnwrapParentheses(ExpressionSyntax expression) return expression; } + private static bool IsBuiltInBlockingMember( + SyntaxNodeAnalysisContext context, + ISymbol accessedSymbol, + CommonInterest.QualifiedMember expectedMember) + { + if (accessedSymbol is not IMethodSymbol { ReducedFrom: not null } reducedMethod) + { + return true; + } + + if (expectedMember.Name != nameof(Task.Wait) + || context.Compilation.GetTypeByMetadataName(Types.Task.FullName) is not INamedTypeSymbol taskType) + { + return false; + } + + return reducedMethod.Parameters.IsEmpty + && reducedMethod.ReturnsVoid + && Utils.IsEqualToOrDerivedFrom(reducedMethod.ReceiverType, taskType); + } + private static ExpressionSyntax GetTaskReceiver(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) { ExpressionSyntax receiver = UnwrapParentheses(memberAccessSyntax.Expression); @@ -737,25 +762,37 @@ 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 => - { - ImmutableHashSet.Builder nestedTaskSymbols = ImmutableHashSet.CreateBuilder(SymbolEqualityComparer.Default); - nestedTaskSymbols.UnionWith(taskSymbols); - foreach (ISymbol taskSymbol in taskSymbols) - { - nestedTaskSymbols.UnionWith(GetSymbolAndRefAliases( - context, - nestedFunction, - taskSymbol, - nestedFunction, - includeAllCandidates: true).Potential); - } - - return MayReassignTask(context, nestedFunction, nestedTaskSymbols.ToImmutable()); - }) is true; + .Any(nestedFunction => FunctionMayReassignTask(nestedFunction, taskSymbols)) is true; } private static bool IsAssignmentToTask(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, IImmutableSet taskSymbols) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index 92d72b2f3..a846f507c 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -210,7 +210,7 @@ private static void AnalyzeInvocation( method => method.Method.IsMatch(invokedMethod) || method.Method.IsMatch(methodDefinition)); bool coveredByVSTHRD103 = !methodsExcludedFromVSTHRD103.Contains(invokedMethod) && !methodsExcludedFromVSTHRD103.Contains(methodDefinition) - && !ShouldAnalyze(context, analyzeWholeCodeBlock) + && IsInTaskReturningMethodOrDelegate(context) && HasAsyncAlternative(context, invocationExpressionSyntax, invokedMethod); if (!isBuiltInSyncBlockingMethod && !coveredByVSTHRD103 @@ -240,9 +240,21 @@ private static bool HasAsyncAlternative( IMethodSymbol invokedMethod) { string asyncMethodName = invokedMethod.Name + VSTHRD200UseAsyncNamingConventionAnalyzer.MandatoryAsyncSuffix; - INamespaceOrTypeSymbol lookupContainer = invokedMethod.ReducedFrom is { Parameters.Length: > 0 } reducedFrom - ? (INamespaceOrTypeSymbol)reducedFrom.Parameters[0].Type - : invokedMethod.ContainingType; + INamespaceOrTypeSymbol lookupContainer = invokedMethod.ContainingType; + if (invokedMethod.ReducedFrom is object) + { + 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; + } + string? declaringMethodName = invocation.FirstAncestorOrSelf()?.Identifier.Text; return context.SemanticModel.LookupSymbols( invocation.Expression.SpanStart, @@ -262,6 +274,20 @@ private static bool HasSupersetOfParameterTypes(IMethodSymbol candidateMethod, I baselineParameter => candidateMethod.Parameters.Any( candidateParameter => SymbolEqualityComparer.Default.Equals(baselineParameter.Type, candidateParameter.Type))); + 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) { if (!ShouldAnalyze(context, analyzeWholeCodeBlock)) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index 46e7310bd..bbe51fbb4 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -43,7 +43,8 @@ public override async Task RegisterCodeFixesAsync(CodeFixContext context) if (TryFindNodeAtSource(diagnostic, root, out ExpressionSyntax? target, out _)) { SemanticModel? semanticModel = await context.Document.GetSemanticModelAsync(context.CancellationToken).ConfigureAwait(false); - if (semanticModel is null + if (target.FirstAncestorOrSelf() is null + || semanticModel is null || semanticModel.GetDiagnostics(target.FullSpan, context.CancellationToken).Any(d => d.Severity == DiagnosticSeverity.Error) || !CanUseAwaitCodeFix(semanticModel, target, context.CancellationToken)) { diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs index 36f078377..0d68f1af2 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs @@ -126,6 +126,41 @@ Task FAsync(CustomWaiter waiter) { 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 8f0dee377..74535f1e5 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -984,8 +984,8 @@ void LaterLocalFunctionStillInvalidatesEarlierGuard(Task task) { } void Replace() { - ref Task alias = ref task; - alias = Task.Run(() => 5); + Action nestedReplace = () => task = Task.Run(() => 5); + nestedReplace(); } } @@ -1153,6 +1153,33 @@ Task FAsync(Task task) { 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 ConfiguredSyncBlockingMethodInsideNameOfDoesNotReport() { @@ -1281,6 +1308,34 @@ void F(Task task) { 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); From 9aaa9619661d0e06bdabd6645160876d9b18f888 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 14:57:46 -0600 Subject: [PATCH 14/26] Restrict await fixes to convertible methods Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 33 ++++++++++- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 58 +++++++++++++++++++ 2 files changed, 89 insertions(+), 2 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index bbe51fbb4..ac805477a 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -43,8 +43,10 @@ public override async Task RegisterCodeFixesAsync(CodeFixContext context) if (TryFindNodeAtSource(diagnostic, root, out ExpressionSyntax? target, out _)) { SemanticModel? semanticModel = await context.Document.GetSemanticModelAsync(context.CancellationToken).ConfigureAwait(false); - if (target.FirstAncestorOrSelf() is null - || semanticModel is null + MethodDeclarationSyntax? containingMethod = target.FirstAncestorOrSelf(); + if (semanticModel is null + || containingMethod is null + || !CanConvertToAsync(semanticModel, containingMethod, context.CancellationToken) || semanticModel.GetDiagnostics(target.FullSpan, context.CancellationToken).Any(d => d.Severity == DiagnosticSeverity.Error) || !CanUseAwaitCodeFix(semanticModel, target, context.CancellationToken)) { @@ -81,6 +83,33 @@ public override async Task RegisterCodeFixesAsync(CodeFixContext context) /// public override FixAllProvider GetFixAllProvider() => WellKnownFixAllProviders.BatchFixer; + private static bool CanConvertToAsync(SemanticModel semanticModel, MethodDeclarationSyntax method, CancellationToken cancellationToken) + { + if (method.Modifiers.Any(SyntaxKind.AsyncKeyword)) + { + return true; + } + + IMethodSymbol? methodSymbol = semanticModel.GetDeclaredSymbol(method, cancellationToken); + if (methodSymbol is null + || methodSymbol.Parameters.Any(parameter => parameter.RefKind != RefKind.None) + || methodSymbol.ReturnsByRef + || methodSymbol.ReturnsByRefReadonly + || method.DescendantNodes( + node => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax) + .OfType() + .Any()) + { + return false; + } + + bool changesContract = !methodSymbol.HasAsyncCompatibleReturnType(); + return !changesContract + || (!methodSymbol.IsVirtual + && !methodSymbol.IsOverride + && !methodSymbol.FindInterfacesImplemented().Any()); + } + private static bool TryFindNodeAtSource(Diagnostic diagnostic, SyntaxNode root, [NotNullWhen(true)] out ExpressionSyntax? target, [NotNullWhen(true)] out Func? transform) { transform = null; diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 74535f1e5..feff4663c 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -1341,6 +1341,64 @@ int Value { await CSVerify.VerifyCodeFixAsync(test, test); } + [Fact] + public async Task CodeFixIsNotOfferedForMethodsThatCannotBeAsync() + { + var test = @" +using System.Collections.Generic; +using System.Threading.Tasks; + +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|]; + } +} +"; + + 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 StaticGetAwaiterFactoryDoesNotOfferCodeFix() { From 9232e0c5bf5e722be6f18b48df33ce83c32badf0 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 19:45:18 -0600 Subject: [PATCH 15/26] Align completion proofs across analyzers Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 14 ++++- .../VSTHRD002UseJtfRunAnalyzer.cs | 12 ++++- .../VSTHRD103UseAsyncOptionAnalyzer.cs | 6 +++ .../VSTHRD002UseJtfRunAnalyzerTests.cs | 51 +++++++++++++++++++ .../VSTHRD103UseAsyncOptionAnalyzerTests.cs | 17 +++++++ 5 files changed, 97 insertions(+), 3 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index e64c0d1e9..bb4e2d2f5 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -237,6 +237,12 @@ internal static void InspectMemberAccess( return context.SemanticModel.GetSymbolInfo(receiver, context.CancellationToken).Symbol; } + /// + /// Determines whether a blocking member access has a receiver that is provably complete. + /// + internal static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) + => HasTaskCompletedCore(context, memberAccessSyntax); + private static ExpressionSyntax UnwrapParentheses(ExpressionSyntax expression) { while (expression is ParenthesizedExpressionSyntax parenthesized) @@ -324,7 +330,7 @@ private static bool IsTaskLike(ITypeSymbol? type) => Utils.IsTask(type) || (type?.Name == nameof(ValueTask) && type.BelongsToNamespace(Namespaces.SystemThreadingTasks)); - private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) + private static bool HasTaskCompletedCore(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) { ExpressionSyntax taskReceiver = GetTaskReceiver(context, memberAccessSyntax); ITypeSymbol? taskType = context.SemanticModel.GetTypeInfo(taskReceiver, context.CancellationToken).Type; @@ -357,6 +363,12 @@ private static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAc return true; } + // Awaiting an IValueTaskSource-backed ValueTask consumes it, so a later Result access is not safe. + if (!Utils.IsTask(taskType)) + { + return false; + } + StatementSyntax? containingStatement = memberAccessSyntax.FirstAncestorOrSelf(); if (containingStatement is null) { diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index a846f507c..9301ad90f 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -206,15 +206,23 @@ private static void AnalyzeInvocation( } IMethodSymbol methodDefinition = invokedMethod.ReducedFrom ?? invokedMethod; + bool isConfiguredSyncBlockingMethod = configuredSyncBlockingMethods.Any( + method => method.IsMatch(invokedMethod) || method.IsMatch(methodDefinition)); + if (!isConfiguredSyncBlockingMethod) + { + return; + } + 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 - && configuredSyncBlockingMethods.Any(method => method.IsMatch(invokedMethod) || method.IsMatch(methodDefinition))) + && !coveredByVSTHRD103) { SimpleNameSyntax? methodName = invocationExpressionSyntax.Expression switch { diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs index 2578f59b8..839833533 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs @@ -240,6 +240,12 @@ private bool InspectMemberAccess(SyntaxNodeAnalysisContext context, ExpressionSy { if (item.Method.IsMatch(memberSymbol)) { + if (memberName.Parent is MemberAccessExpressionSyntax memberAccess + && CSharpCommonInterest.HasTaskCompleted(context, memberAccess)) + { + return false; + } + // Check if this method is excluded from VSTHRD103 diagnostics if (this.excludedMethods.Contains(memberSymbol)) { diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index feff4663c..2d2918d75 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -547,6 +547,29 @@ 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 TaskResultShouldReportWarning_WithinAnonymousDelegate() { @@ -1180,6 +1203,34 @@ public static void WaitSynchronously(this T value) { } 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() { diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs index 8d681887b..92662ac02 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs @@ -124,6 +124,23 @@ 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 JTFRunOfTInTaskReturningMethodGeneratesWarning() { From edd04e818cc5cdade701d3314eb6f936f8f9ae9c Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 20:02:47 -0600 Subject: [PATCH 16/26] Handle completed waits and forbidden awaits Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../VSTHRD103UseAsyncOptionAnalyzer.cs | 2 +- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 2 ++ .../VSTHRD002UseJtfRunAnalyzerTests.cs | 23 +++++++++++++++++++ .../VSTHRD103UseAsyncOptionAnalyzerTests.cs | 23 +++++++++++++++++++ 4 files changed, 49 insertions(+), 1 deletion(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs index 839833533..d15433224 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs @@ -243,7 +243,7 @@ private bool InspectMemberAccess(SyntaxNodeAnalysisContext context, ExpressionSy if (memberName.Parent is MemberAccessExpressionSyntax memberAccess && CSharpCommonInterest.HasTaskCompleted(context, memberAccess)) { - return false; + return true; } // Check if this method is excluded from VSTHRD103 diagnostics diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index ac805477a..16f9d732c 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -46,6 +46,8 @@ public override async Task RegisterCodeFixesAsync(CodeFixContext context) MethodDeclarationSyntax? containingMethod = target.FirstAncestorOrSelf(); if (semanticModel is null || containingMethod is null + || target.Ancestors().TakeWhile(node => node != containingMethod) + .Any(node => node is LockStatementSyntax or CatchFilterClauseSyntax) || !CanConvertToAsync(semanticModel, containingMethod, context.CancellationToken) || semanticModel.GetDiagnostics(target.FullSpan, context.CancellationToken).Any(d => d.Severity == DiagnosticSeverity.Error) || !CanUseAwaitCodeFix(semanticModel, target, context.CancellationToken)) diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 2d2918d75..70eafdbc8 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -1418,6 +1418,29 @@ IEnumerable Iterator(Task task) { 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 CodeFixIsNotOfferedWhenChangingMethodContract() { diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs index 92662ac02..602a98ec6 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs @@ -141,6 +141,29 @@ async Task T(Task task) { await CSVerify.VerifyAnalyzerAsync(test); } + [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() { From 405299ed7709b9f83fd00c2879011782ecf08d8d Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 20:24:52 -0600 Subject: [PATCH 17/26] Avoid await fixes for ref-like signatures Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 3 ++- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 12 ++++++++++++ 2 files changed, 14 insertions(+), 1 deletion(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index 16f9d732c..df55509ee 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -94,9 +94,10 @@ private static bool CanConvertToAsync(SemanticModel semanticModel, MethodDeclara IMethodSymbol? methodSymbol = semanticModel.GetDeclaredSymbol(method, cancellationToken); if (methodSymbol is null - || methodSymbol.Parameters.Any(parameter => parameter.RefKind != RefKind.None) + || methodSymbol.Parameters.Any(parameter => parameter.RefKind != RefKind.None || parameter.Type.IsRefLikeType) || methodSymbol.ReturnsByRef || methodSymbol.ReturnsByRefReadonly + || methodSymbol.ReturnType.IsRefLikeType || method.DescendantNodes( node => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax) .OfType() diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 70eafdbc8..1ec4b1857 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -1399,6 +1399,10 @@ public async Task CodeFixIsNotOfferedForMethodsThatCannotBeAsync() using System.Collections.Generic; using System.Threading.Tasks; +ref struct RefLike { + internal int Value; +} + class Test { int RefParameter(Task task, ref int value) { return task.[|Result|]; @@ -1412,6 +1416,14 @@ int OutParameter(Task task, out int value) { 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|] }; + } } "; From d63118f018dabfe5a4500bd343038859671a17e0 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 20:44:57 -0600 Subject: [PATCH 18/26] Harden VSTHRD002 completion analysis Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 246 +++++++++++++++--- .../VSTHRD002UseJtfRunAnalyzer.cs | 137 +--------- .../VSTHRD103UseAsyncOptionAnalyzer.cs | 29 ++- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 9 +- .../MultiAnalyzerTests.cs | 34 +++ .../VSTHRD002UseJtfRunAnalyzerTests.cs | 40 ++- .../VSTHRD103UseAsyncOptionAnalyzerTests.cs | 117 +++++++++ 7 files changed, 442 insertions(+), 170 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index bb4e2d2f5..c24977038 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -46,8 +46,16 @@ 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; @@ -91,7 +99,21 @@ bool DefinitelyPrecedesNode(SyntaxNode candidate) : variable.Initializer.Value; if (context.SemanticModel.GetSymbolInfo(UnwrapParentheses(initializer), context.CancellationToken).Symbol is ISymbol initializedFrom) { - refTargets[local] = GetRefTargets(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 (candidate is AssignmentExpressionSyntax { Right: RefExpressionSyntax refAssignment } assignment @@ -103,10 +125,22 @@ bool DefinitelyPrecedesNode(SyntaxNode candidate) || !refTargets.TryGetValue(reboundLocal, out HashSet? existingTargets)) { refTargets[reboundLocal] = assignedTargets; + if (potentialOnlyRefLocals.Contains(assignedFrom)) + { + potentialOnlyRefLocals.Add(reboundLocal); + } + else + { + potentialOnlyRefLocals.Remove(reboundLocal); + } } else { existingTargets.UnionWith(assignedTargets); + if (potentialOnlyRefLocals.Contains(assignedFrom)) + { + potentialOnlyRefLocals.Add(reboundLocal); + } } } } @@ -124,7 +158,9 @@ bool DefinitelyPrecedesNode(SyntaxNode candidate) foreach (KeyValuePair> refTarget in refTargets) { - if (symbolTargets.Count == 1 && refTarget.Value.SetEquals(symbolTargets)) + if (symbolTargets.Count == 1 + && refTarget.Value.SetEquals(symbolTargets) + && !potentialOnlyRefLocals.Contains(refTarget.Key)) { definiteSymbols.Add(refTarget.Key); } @@ -241,7 +277,118 @@ internal static void InspectMemberAccess( /// Determines whether a blocking member access has a receiver that is provably complete. /// internal static bool HasTaskCompleted(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) - => HasTaskCompletedCore(context, 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()) + { + if (anonymousFunction.Parent is not ArgumentSyntax anonymousFunctionArgument + || anonymousFunctionArgument.Parent?.Parent is not InvocationExpressionSyntax continuationInvocation + || continuationInvocation.ArgumentList.Arguments.FirstOrDefault() != anonymousFunctionArgument) + { + continue; + } + + if (context.SemanticModel.GetSymbolInfo(continuationInvocation, context.CancellationToken).Symbol is not IMethodSymbol invokedMethod + || invokedMethod.Name != nameof(Task.ContinueWith) + || !Utils.IsTask(invokedMethod.ContainingType)) + { + 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; + } + + (ImmutableHashSet taskSymbols, ImmutableHashSet potentialTaskSymbols) = + GetSymbolAndRefAliases(context, accessSyntax, completedTask); + if (accessSyntax.Ancestors().TakeWhile(node => node != anonymousFunction) + .Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax)) + { + potentialTaskSymbols = potentialTaskSymbols.Union(GetSymbolAndRefAliases( + context, + accessSyntax, + completedTask, + anonymousFunction, + includeAllCandidates: true).Potential); + } + + ISymbol? receiverSymbol = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(taskReceiver), context.CancellationToken).Symbol; + if (receiverSymbol is object + && taskSymbols.Contains(receiverSymbol) + && !IsTaskReassignedInContinuation(context, anonymousFunction, accessSyntax, potentialTaskSymbols)) + { + return true; + } + } + + return false; + } + + private static bool IsTaskReassignedInContinuation( + SyntaxNodeAnalysisContext context, + AnonymousFunctionExpressionSyntax continuation, + SyntaxNode accessSyntax, + IImmutableSet taskSymbols) + { + 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()) + { + 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; + } + } + + foreach (ArgumentSyntax argument in continuation.DescendantNodes().OfType()) + { + 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)) + && IsOneOfSymbols(context, argument.Expression, taskSymbols)) + { + return true; + } + } + + return false; + } private static ExpressionSyntax UnwrapParentheses(ExpressionSyntax expression) { @@ -330,9 +477,12 @@ private static bool IsTaskLike(ITypeSymbol? type) => Utils.IsTask(type) || (type?.Name == nameof(ValueTask) && type.BelongsToNamespace(Namespaces.SystemThreadingTasks)); - private static bool HasTaskCompletedCore(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) + private static bool HasTaskCompletedCore( + SyntaxNodeAnalysisContext context, + ExpressionSyntax taskReceiver, + SyntaxNode accessSyntax) { - ExpressionSyntax taskReceiver = GetTaskReceiver(context, memberAccessSyntax); + taskReceiver = UnwrapParentheses(taskReceiver); ITypeSymbol? taskType = context.SemanticModel.GetTypeInfo(taskReceiver, context.CancellationToken).Type; if (!IsTaskLike(taskType)) { @@ -345,22 +495,22 @@ private static bool HasTaskCompletedCore(SyntaxNodeAnalysisContext context, Memb return false; } - if (context.SemanticModel.GetEnclosingSymbol(memberAccessSyntax.SpanStart, context.CancellationToken) is not IMethodSymbol enclosingMethod + 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, memberAccessSyntax, taskSymbol); - if (NestedFunctionMayReassignTask(context, memberAccessSyntax, potentialTaskSymbols)) + GetSymbolAndRefAliases(context, accessSyntax, taskSymbol); + if (NestedFunctionMayReassignTask(context, accessSyntax, potentialTaskSymbols)) { return false; } - if (IsWithinCompletedTaskBranch(context, memberAccessSyntax, taskSymbols, potentialTaskSymbols)) + if (IsWithinCompletedTaskBranch(context, accessSyntax, taskSymbols, potentialTaskSymbols)) { - return true; + return Utils.IsTask(taskType) || !MayHaveAwaitedTaskBefore(context, accessSyntax, taskSymbols); } // Awaiting an IValueTaskSource-backed ValueTask consumes it, so a later Result access is not safe. @@ -369,29 +519,35 @@ private static bool HasTaskCompletedCore(SyntaxNodeAnalysisContext context, Memb return false; } - StatementSyntax? containingStatement = memberAccessSyntax.FirstAncestorOrSelf(); + StatementSyntax? containingStatement = accessSyntax.FirstAncestorOrSelf(); if (containingStatement is null) { return false; } - if (TryGetAwaitExpression(context, containingStatement, taskSymbols, memberAccessSyntax.SpanStart, out AwaitExpressionSyntax? precedingAwait) - && !MayReassignTask(context, containingStatement, potentialTaskSymbols, precedingAwait.Span.End, memberAccessSyntax.SpanStart)) + if (TryGetAwaitExpression(context, containingStatement, taskSymbols, accessSyntax.SpanStart, out AwaitExpressionSyntax? precedingAwait) + && !MayReassignTask(context, containingStatement, potentialTaskSymbols, precedingAwait.Span.End, accessSyntax.SpanStart)) { return true; } - while (containingStatement.Parent is BlockSyntax block) + while (true) { - if (MayReassignTask(context, containingStatement, potentialTaskSymbols, containingStatement.SpanStart - 1, memberAccessSyntax.SpanStart)) + if (MayReassignTask(context, containingStatement, potentialTaskSymbols, containingStatement.SpanStart - 1, accessSyntax.SpanStart)) { return false; } - int statementIndex = block.Statements.IndexOf(containingStatement); + SyntaxList statements = containingStatement.Parent switch + { + BlockSyntax block => block.Statements, + SwitchSectionSyntax switchSection => switchSection.Statements, + _ => default, + }; + int statementIndex = statements.IndexOf(containingStatement); for (int i = statementIndex - 1; i >= 0; i--) { - StatementSyntax statement = block.Statements[i]; + StatementSyntax statement = statements[i]; if (StatementCompletesTask(context, statement, taskSymbols, potentialTaskSymbols)) { return true; @@ -403,8 +559,19 @@ private static bool HasTaskCompletedCore(SyntaxNodeAnalysisContext context, Memb } } - StatementSyntax? outerStatement = containingStatement.Ancestors().OfType().FirstOrDefault(statement => statement.Parent is BlockSyntax); - if (outerStatement is not IfStatementSyntax and not BlockSyntax) + StatementSyntax? outerStatement = containingStatement.Ancestors().OfType() + .FirstOrDefault(statement => statement.Parent is BlockSyntax or SwitchSectionSyntax); + if (outerStatement is null) + { + return false; + } + + if (outerStatement is WhileStatementSyntax + or DoStatementSyntax + or ForStatementSyntax + or ForEachStatementSyntax + or ForEachVariableStatementSyntax + && MayReassignTask(context, outerStatement, potentialTaskSymbols)) { return false; } @@ -416,50 +583,69 @@ private static bool HasTaskCompletedCore(SyntaxNodeAnalysisContext context, Memb containingStatement = outerStatement; } + } - return false; + private static bool MayHaveAwaitedTaskBefore( + 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; + + return searchRoot.DescendantNodes(DescendIntoChildren) + .OfType() + .Any(awaitExpression => awaitExpression.SpanStart < accessSyntax.SpanStart + && AwaitCompletesTask(context, awaitExpression, taskSymbols)); } private static bool IsWithinCompletedTaskBranch( SyntaxNodeAnalysisContext context, - MemberAccessExpressionSyntax memberAccessSyntax, + SyntaxNode accessSyntax, IImmutableSet taskSymbols, IImmutableSet potentialTaskSymbols) { - foreach (IfStatementSyntax ifStatement in memberAccessSyntax.Ancestors().OfType()) + foreach (IfStatementSyntax ifStatement in accessSyntax.Ancestors().OfType()) { - IEnumerable nodesBetweenAccessAndCondition = memberAccessSyntax.Ancestors().TakeWhile(node => node != ifStatement); + IEnumerable nodesBetweenAccessAndCondition = accessSyntax.Ancestors().TakeWhile(node => node != ifStatement); if (nodesBetweenAccessAndCondition.Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax)) { continue; } - foreach (BinaryExpressionSyntax binary in memberAccessSyntax.Ancestors() + foreach (BinaryExpressionSyntax binary in accessSyntax.Ancestors() .TakeWhile(node => node != ifStatement) .OfType()) { - bool leftProvesCompletion = binary.Right.FullSpan.Contains(memberAccessSyntax.Span) + 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.Span.End, memberAccessSyntax.SpanStart)) + && !MayReassignTask(context, binary, potentialTaskSymbols, binary.Left.Span.End, accessSyntax.SpanStart)) { return true; } } - if (ifStatement.Statement.FullSpan.Contains(memberAccessSyntax.Span) + if (ifStatement.Statement.FullSpan.Contains(accessSyntax.Span) && ConditionProvesCompletion(context, ifStatement.Condition, taskSymbols, conditionValue: true) - && !MayReassignTask(context, ifStatement, potentialTaskSymbols, ifStatement.Condition.SpanStart - 1, memberAccessSyntax.SpanStart)) + && !MayReassignTask(context, ifStatement, potentialTaskSymbols, ifStatement.Condition.SpanStart - 1, accessSyntax.SpanStart)) { return true; } - if (ifStatement.Else?.Statement.FullSpan.Contains(memberAccessSyntax.Span) is 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, memberAccessSyntax.SpanStart)) + && !MayReassignTask(context, ifStatement, potentialTaskSymbols, ifStatement.Condition.SpanStart - 1, accessSyntax.SpanStart)) { return true; } diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index 9301ad90f..fdd14399f 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -80,9 +80,9 @@ public override void Initialize(AnalysisContext context) { bool analyzeWholeCodeBlock = propertySymbol is object || !methodSymbol!.HasAsyncCompatibleReturnType(); codeBlockContext.RegisterSyntaxNodeAction( - Utils.DebuggableWrapper(c => AnalyzeInvocation(c, taskSymbol, configuredSyncBlockingMethods, methodsExcludedFromVSTHRD103, analyzeWholeCodeBlock)), + Utils.DebuggableWrapper(c => AnalyzeInvocation(c, configuredSyncBlockingMethods, methodsExcludedFromVSTHRD103, analyzeWholeCodeBlock)), SyntaxKind.InvocationExpression); - codeBlockContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(c => AnalyzeMemberAccess(c, taskSymbol, analyzeWholeCodeBlock)), SyntaxKind.SimpleMemberAccessExpression); + codeBlockContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(c => AnalyzeMemberAccess(c, analyzeWholeCodeBlock)), SyntaxKind.SimpleMemberAccessExpression); } }); } @@ -109,82 +109,21 @@ private static bool ShouldAnalyze(SyntaxNodeAnalysisContext context, bool analyz && !containingMethod.HasAsyncCompatibleReturnType(); } - private static ParameterSyntax? GetFirstParameter(AnonymousFunctionExpressionSyntax? anonymousFunctionSyntax) - { - switch (anonymousFunctionSyntax) - { - case SimpleLambdaExpressionSyntax lambda: - return lambda.Parameter; - case ParenthesizedLambdaExpressionSyntax lambda: - return lambda.ParameterList.Parameters.FirstOrDefault(); - case AnonymousMethodExpressionSyntax anonymousMethod: - return anonymousMethod.ParameterList?.Parameters.FirstOrDefault(); - } - - return null; - } - private static void InspectMemberAccess( SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax? memberAccessSyntax, - IEnumerable problematicMethods, - INamedTypeSymbol taskSymbol) + IEnumerable problematicMethods) { if (memberAccessSyntax is null) { return; } - // A continuation's antecedent is complete throughout its delegate, including nested delegates that capture it. - foreach (AnonymousFunctionExpressionSyntax anonymousFunctionSyntax in context.Node.Ancestors().OfType()) - { - var anonymousFunctionArgument = anonymousFunctionSyntax.Parent as ArgumentSyntax; - var continuationInvocation = anonymousFunctionArgument?.Parent?.Parent as InvocationExpressionSyntax; - if (continuationInvocation is null || continuationInvocation.ArgumentList.Arguments.FirstOrDefault() != anonymousFunctionArgument) - { - continue; - } - - var invokedMemberSymbol = context.SemanticModel.GetSymbolInfo(continuationInvocation, context.CancellationToken).Symbol as IMethodSymbol; - if (invokedMemberSymbol?.Name != nameof(Task.ContinueWith) - || !Utils.IsEqualToOrDerivedFrom(invokedMemberSymbol.ContainingType, taskSymbol)) - { - continue; - } - - ParameterSyntax? firstParameter = GetFirstParameter(anonymousFunctionSyntax); - if (firstParameter is object - && context.SemanticModel.GetDeclaredSymbol(firstParameter, context.CancellationToken) is IParameterSymbol completedTask) - { - (ImmutableHashSet taskSymbols, ImmutableHashSet potentialTaskSymbols) = - CSharpCommonInterest.GetSymbolAndRefAliases(context, memberAccessSyntax, completedTask); - if (memberAccessSyntax.Ancestors().TakeWhile(node => node != anonymousFunctionSyntax) - .Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax)) - { - potentialTaskSymbols = potentialTaskSymbols.Union(CSharpCommonInterest.GetSymbolAndRefAliases( - context, - memberAccessSyntax, - completedTask, - anonymousFunctionSyntax, - includeAllCandidates: true).Potential); - } - - ISymbol? receiverSymbol = GetTaskReceiverSymbol(context, memberAccessSyntax); - if (receiverSymbol is object - && taskSymbols.Contains(receiverSymbol) - && !IsTaskReassignedInContinuation(context, anonymousFunctionSyntax, memberAccessSyntax, potentialTaskSymbols)) - { - return; - } - } - } - CSharpCommonInterest.InspectMemberAccess(context, memberAccessSyntax, Descriptor, problematicMethods); } private static void AnalyzeInvocation( SyntaxNodeAnalysisContext context, - INamedTypeSymbol taskSymbol, ImmutableArray configuredSyncBlockingMethods, ImmutableArray methodsExcludedFromVSTHRD103, bool analyzeWholeCodeBlock) @@ -195,8 +134,7 @@ private static void AnalyzeInvocation( InspectMemberAccess( context, invocationExpressionSyntax.Expression as MemberAccessExpressionSyntax, - CommonInterest.ProblematicSyncBlockingMethods, - taskSymbol); + CommonInterest.ProblematicSyncBlockingMethods); } if (configuredSyncBlockingMethods.IsEmpty @@ -296,7 +234,7 @@ private static bool IsInTaskReturningMethodOrDelegate(SyntaxNodeAnalysisContext 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)) { @@ -307,69 +245,6 @@ private static void AnalyzeMemberAccess(SyntaxNodeAnalysisContext context, IName InspectMemberAccess( context, memberAccessSyntax, - CommonInterest.SyncBlockingProperties, - taskSymbol); - } - - private static ISymbol? GetTaskReceiverSymbol(SyntaxNodeAnalysisContext context, MemberAccessExpressionSyntax memberAccessSyntax) - => CSharpCommonInterest.GetTaskReceiverSymbol(context, memberAccessSyntax); - - private static ExpressionSyntax UnwrapParentheses(ExpressionSyntax expression) - { - while (expression is ParenthesizedExpressionSyntax parenthesized) - { - expression = parenthesized.Expression; - } - - return expression; - } - - private static bool IsTaskReassignedInContinuation( - SyntaxNodeAnalysisContext context, - AnonymousFunctionExpressionSyntax continuation, - MemberAccessExpressionSyntax memberAccess, - ImmutableHashSet taskSymbols) - { - bool accessIsNested = memberAccess.Ancestors().TakeWhile(node => node != continuation).Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax); - int beforePosition = accessIsNested ? continuation.Span.End + 1 : memberAccess.SpanStart; - - foreach (AssignmentExpressionSyntax assignment in continuation.DescendantNodes().OfType()) - { - bool isDeferredWrite = assignment.Ancestors().TakeWhile(node => node != continuation) - .Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax); - if ((assignment.SpanStart < beforePosition || isDeferredWrite) - && IsAssignmentToParameter(context, assignment.Left, taskSymbols)) - { - return true; - } - } - - foreach (ArgumentSyntax argument in continuation.DescendantNodes().OfType()) - { - ISymbol? argumentSymbol = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(argument.Expression), context.CancellationToken).Symbol; - bool isDeferredWrite = argument.Ancestors().TakeWhile(node => node != continuation) - .Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax); - if ((argument.SpanStart < beforePosition || isDeferredWrite) - && (argument.RefKindKeyword.IsKind(SyntaxKind.RefKeyword) || argument.RefKindKeyword.IsKind(SyntaxKind.OutKeyword)) - && argumentSymbol is object - && taskSymbols.Contains(argumentSymbol)) - { - return true; - } - } - - return false; - } - - private static bool IsAssignmentToParameter(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, IImmutableSet taskSymbols) - { - expression = UnwrapParentheses(expression); - if (expression is TupleExpressionSyntax tuple) - { - return tuple.Arguments.Any(argument => IsAssignmentToParameter(context, argument.Expression, taskSymbols)); - } - - ISymbol? symbol = context.SemanticModel.GetSymbolInfo(expression, context.CancellationToken).Symbol; - return symbol is object && taskSymbols.Contains(symbol); + CommonInterest.SyncBlockingProperties); } } diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs index d15433224..6439b536f 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); } } @@ -158,7 +163,7 @@ internal void AnalyzeInvocation(SyntaxNodeAnalysisContext context) && m.HasAsyncCompatibleReturnType()) { // Check if this method is excluded from VSTHRD103 diagnostics - if (this.excludedMethods.Contains(methodSymbol)) + if (this.IsExcluded(methodSymbol)) { return; } @@ -231,7 +236,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,14 +250,16 @@ private bool InspectMemberAccess(SyntaxNodeAnalysisContext context, ExpressionSy { if (item.Method.IsMatch(memberSymbol)) { - if (memberName.Parent is MemberAccessExpressionSyntax memberAccess - && CSharpCommonInterest.HasTaskCompleted(context, memberAccess)) + if ((memberName.Parent is MemberAccessExpressionSyntax memberAccess + && CSharpCommonInterest.HasTaskCompleted(context, memberAccess)) + || (taskReceiver is object + && CSharpCommonInterest.HasTaskCompleted(context, taskReceiver, accessSyntax ?? memberName))) { return true; } // Check if this method is excluded from VSTHRD103 diagnostics - if (this.excludedMethods.Contains(memberSymbol)) + if (this.IsExcluded(memberSymbol)) { return false; } @@ -279,5 +291,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 df55509ee..2c1d82138 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -101,14 +101,19 @@ private static bool CanConvertToAsync(SemanticModel semanticModel, MethodDeclara || method.DescendantNodes( node => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax) .OfType() - .Any()) + .Any() + || method.DescendantNodes( + node => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax) + .OfType() + .Any(variable => semanticModel.GetDeclaredSymbol(variable, cancellationToken) is ILocalSymbol { RefKind: not RefKind.None })) { return false; } bool changesContract = !methodSymbol.HasAsyncCompatibleReturnType(); return !changesContract - || (!methodSymbol.IsVirtual + || (!method.Modifiers.Any(SyntaxKind.PartialKeyword) + && !methodSymbol.IsVirtual && !methodSymbol.IsOverride && !methodSymbol.FindInterfacesImplemented().Any()); } diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs index 0d68f1af2..5f2a2e609 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/MultiAnalyzerTests.cs @@ -126,6 +126,40 @@ Task FAsync(CustomWaiter waiter) { 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() { diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 1ec4b1857..b57bc0572 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -613,6 +613,10 @@ void F() { 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) { } @@ -691,6 +695,28 @@ void F() { await CSVerify.VerifyCodeFixAsync(test, test); } + [Fact] + public async Task RefReturningInvocationCreatesPotentialAlias() + { + var test = @" +using System.Threading.Tasks; + +class Test { + 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, test); + } + [Fact] public async Task ContinuationParameterDeconstructionReportsWarning() { @@ -1403,7 +1429,7 @@ ref struct RefLike { internal int Value; } -class Test { +partial class Test { int RefParameter(Task task, ref int value) { return task.[|Result|]; } @@ -1424,6 +1450,18 @@ int RefLikeParameter(Task task, RefLike value) { 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|]; + } + + private partial int PartialMethod(Task task); + + private partial int PartialMethod(Task task) { + return task.[|Result|]; + } } "; diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs index 602a98ec6..35d7d1eed 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs @@ -141,6 +141,103 @@ async Task T(Task task) { 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() { @@ -830,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 TaskGetAwaiterGetResultInTaskReturningMethodGeneratesWarning() { From b7e4be0e543d26cd93320078564c130e52dfe6f9 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 20:50:00 -0600 Subject: [PATCH 19/26] Guard await fixes against delegate references Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 5 ++ .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 57 +++++++++++++++++-- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 14 +++++ 3 files changed, 72 insertions(+), 4 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index c24977038..8bc31e8bd 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -410,6 +410,11 @@ private static bool IsBuiltInBlockingMember( return true; } + if (expectedMember.IsMatch(reducedMethod.ReducedFrom)) + { + return true; + } + if (expectedMember.Name != nameof(Task.Wait) || context.Compilation.GetTypeByMetadataName(Types.Task.FullName) is not INamedTypeSymbol taskType) { diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index 2c1d82138..1f89d3745 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -15,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; @@ -48,7 +49,7 @@ public override async Task RegisterCodeFixesAsync(CodeFixContext context) || containingMethod is null || target.Ancestors().TakeWhile(node => node != containingMethod) .Any(node => node is LockStatementSyntax or CatchFilterClauseSyntax) - || !CanConvertToAsync(semanticModel, containingMethod, context.CancellationToken) + || !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)) { @@ -85,7 +86,11 @@ public override async Task RegisterCodeFixesAsync(CodeFixContext context) /// public override FixAllProvider GetFixAllProvider() => WellKnownFixAllProviders.BatchFixer; - private static bool CanConvertToAsync(SemanticModel semanticModel, MethodDeclarationSyntax method, CancellationToken cancellationToken) + private static async Task CanConvertToAsyncAsync( + Document document, + SemanticModel semanticModel, + MethodDeclarationSyntax method, + CancellationToken cancellationToken) { if (method.Modifiers.Any(SyntaxKind.AsyncKeyword)) { @@ -105,7 +110,8 @@ private static bool CanConvertToAsync(SemanticModel semanticModel, MethodDeclara || method.DescendantNodes( node => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax) .OfType() - .Any(variable => semanticModel.GetDeclaredSymbol(variable, cancellationToken) is ILocalSymbol { RefKind: not RefKind.None })) + .Any(variable => semanticModel.GetDeclaredSymbol(variable, cancellationToken) is ILocalSymbol local + && (local.RefKind != RefKind.None || local.Type.IsRefLikeType))) { return false; } @@ -115,7 +121,50 @@ private static bool CanConvertToAsync(SemanticModel semanticModel, MethodDeclara || (!method.Modifiers.Any(SyntaxKind.PartialKeyword) && !methodSymbol.IsVirtual && !methodSymbol.IsOverride - && !methodSymbol.FindInterfacesImplemented().Any()); + && !methodSymbol.FindInterfacesImplemented().Any() + && !await HasMethodGroupReferenceAsync(document.Project.Solution, methodSymbol, cancellationToken).ConfigureAwait(false)); + } + + 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) diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index b57bc0572..12dc15207 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -1422,6 +1422,7 @@ int Value { public async Task CodeFixIsNotOfferedForMethodsThatCannotBeAsync() { var test = @" +using System; using System.Collections.Generic; using System.Threading.Tasks; @@ -1457,11 +1458,24 @@ int RefLocal(Task task) { return task.[|Result|]; } + int RefLikeLocal(Task task) { + RefLike value = default; + 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; + } } "; From d8cc9725f317096277aecb2868d4670e9448ef49 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 21:05:36 -0600 Subject: [PATCH 20/26] Reject by-ref parameter completion proofs Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 5 ++++ .../VSTHRD103UseAsyncOptionAnalyzerTests.cs | 26 +++++++++++++++++++ 2 files changed, 31 insertions(+) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index 7d25577cc..8b53e5ee2 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -510,6 +510,11 @@ private static bool HasTaskCompletedCore( (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; diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs index f88bf60d0..41666852b 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs @@ -1397,6 +1397,32 @@ 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 TaskGetAwaiterGetResultInTaskReturningMethodGeneratesWarning() { From 2a1a1bc98af52d84bce19adeeae413cd97f06719 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 21:27:25 -0600 Subject: [PATCH 21/26] Harden VSTHRD002 flow and code fixes Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 181 ++++++++++--- .../VSTHRD002UseJtfRunAnalyzer.cs | 39 ++- .../VSTHRD103UseAsyncOptionAnalyzer.cs | 30 ++- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 139 ++++++++-- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 149 ++++++++++ .../VSTHRD103UseAsyncOptionAnalyzerTests.cs | 254 ++++++++++++++++++ 6 files changed, 727 insertions(+), 65 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index 8b53e5ee2..ee421ad10 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -116,31 +116,44 @@ bool DefinitelyPrecedesNode(SyntaxNode candidate) } } } + 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 - && context.SemanticModel.GetSymbolInfo(UnwrapParentheses(refAssignment.Expression), context.CancellationToken).Symbol is ISymbol assignedFrom) + && context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol is ILocalSymbol { RefKind: not RefKind.None } reboundLocal) { - HashSet assignedTargets = GetRefTargets(assignedFrom); - if ((!includeAllCandidates && DefinitelyPrecedesNode(candidate)) - || !refTargets.TryGetValue(reboundLocal, out HashSet? existingTargets)) + 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) { - refTargets[reboundLocal] = assignedTargets; - if (potentialOnlyRefLocals.Contains(assignedFrom)) + if ((!includeAllCandidates && DefinitelyPrecedesNode(candidate)) + || !refTargets.TryGetValue(reboundLocal, out HashSet? existingTargets)) { - potentialOnlyRefLocals.Add(reboundLocal); + refTargets[reboundLocal] = assignedTargets; + if (assignedFrom is null || potentialOnlyRefLocals.Contains(assignedFrom)) + { + potentialOnlyRefLocals.Add(reboundLocal); + } + else + { + potentialOnlyRefLocals.Remove(reboundLocal); + } } else { - potentialOnlyRefLocals.Remove(reboundLocal); - } - } - else - { - existingTargets.UnionWith(assignedTargets); - if (potentialOnlyRefLocals.Contains(assignedFrom)) - { - potentialOnlyRefLocals.Add(reboundLocal); + existingTargets.UnionWith(assignedTargets); + if (assignedFrom is null || potentialOnlyRefLocals.Contains(assignedFrom)) + { + potentialOnlyRefLocals.Add(reboundLocal); + } } } } @@ -382,7 +395,7 @@ private static bool IsTaskReassignedInContinuation( || nestedFunction?.SpanStart < accessSyntax.SpanStart; if ((argument.SpanStart < beforePosition || isDeferredWrite) && (argument.RefKindKeyword.IsKind(SyntaxKind.RefKeyword) || argument.RefKindKeyword.IsKind(SyntaxKind.OutKeyword)) - && IsOneOfSymbols(context, argument.Expression, taskSymbols)) + && MayAliasTaskStorage(context, argument.Expression, taskSymbols)) { return true; } @@ -520,6 +533,11 @@ private static bool HasTaskCompletedCore( return false; } + if (ContainsPotentialControlFlowBypass(accessSyntax)) + { + return false; + } + if (IsWithinCompletedTaskBranch(context, accessSyntax, taskSymbols, potentialTaskSymbols)) { return Utils.IsTask(taskType) || !MayHaveAwaitedTaskBefore(context, accessSyntax, taskSymbols); @@ -534,7 +552,10 @@ private static bool HasTaskCompletedCore( StatementSyntax? containingStatement = accessSyntax.FirstAncestorOrSelf(); if (containingStatement is null) { - return false; + 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) @@ -612,10 +633,65 @@ or BaseMethodDeclarationSyntax bool DescendIntoChildren(SyntaxNode child) => child == searchRoot || child is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; + var taskAndCopySymbols = new HashSet(taskSymbols, SymbolEqualityComparer.Default); + bool IsTaskOrCopy(ExpressionSyntax expression) + { + ISymbol? expressionSymbol = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(expression), context.CancellationToken).Symbol; + return expressionSymbol is object && taskAndCopySymbols.Contains(expressionSymbol); + } + + bool addedCopy; + do + { + addedCopy = false; + foreach (VariableDeclaratorSyntax variable in searchRoot.DescendantNodes(DescendIntoChildren) + .OfType() + .Where(variable => variable.SpanStart < accessSyntax.SpanStart && variable.Initializer is object)) + { + if (context.SemanticModel.GetDeclaredSymbol(variable, context.CancellationToken) is ILocalSymbol local + && IsTaskOrCopy(variable.Initializer!.Value) + && taskAndCopySymbols.Add(local)) + { + addedCopy = true; + } + } + + foreach (AssignmentExpressionSyntax assignment in searchRoot.DescendantNodes(DescendIntoChildren) + .OfType() + .Where(assignment => assignment.SpanStart < accessSyntax.SpanStart)) + { + if (context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol is ILocalSymbol local + && IsTaskOrCopy(assignment.Right) + && taskAndCopySymbols.Add(local)) + { + addedCopy = true; + } + } + } + while (addedCopy); + + IImmutableSet taskAndCopySymbolSet = taskAndCopySymbols.ToImmutableHashSet(SymbolEqualityComparer.Default); return searchRoot.DescendantNodes(DescendIntoChildren) .OfType() .Any(awaitExpression => awaitExpression.SpanStart < accessSyntax.SpanStart - && AwaitCompletesTask(context, awaitExpression, taskSymbols)); + && AwaitCompletesTask(context, awaitExpression, taskAndCopySymbolSet)); + } + + 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( @@ -653,7 +729,7 @@ or ForEachStatementSyntax || (binary.IsKind(SyntaxKind.LogicalOrExpression) && ConditionProvesCompletion(context, binary.Left, taskSymbols, conditionValue: false))); if (leftProvesCompletion - && !MayReassignTask(context, binary, potentialTaskSymbols, binary.Left.Span.End, accessSyntax.SpanStart)) + && !MayReassignTask(context, binary, potentialTaskSymbols, binary.Left.SpanStart - 1, accessSyntax.SpanStart)) { return true; } @@ -850,36 +926,38 @@ private static bool TryGetAwaitExpression( private static bool TryGetAwaitExpression( SyntaxNodeAnalysisContext context, - StatementSyntax statement, + SyntaxNode node, IImmutableSet taskSymbols, int beforePosition, [NotNullWhen(true)] out AwaitExpressionSyntax? awaitExpression) { static bool DescendIntoChildren(SyntaxNode node) => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; - if (statement is WhileStatementSyntax or DoStatementSyntax or ForStatementSyntax or ForEachStatementSyntax or ForEachVariableStatementSyntax) + if (node is WhileStatementSyntax or DoStatementSyntax or ForStatementSyntax or ForEachStatementSyntax or ForEachVariableStatementSyntax) { awaitExpression = null; return false; } - foreach (AwaitExpressionSyntax candidate in statement.DescendantNodes(DescendIntoChildren).OfType().Reverse()) + foreach (AwaitExpressionSyntax candidate in node.DescendantNodes(DescendIntoChildren).OfType().Reverse()) { if (candidate.Span.End >= beforePosition) { continue; } - IEnumerable ancestorsWithinStatement = candidate.Ancestors().TakeWhile(node => node != statement); + IEnumerable ancestorsWithinStatement = candidate.Ancestors().TakeWhile(ancestor => ancestor != node); bool isConditionallyExecuted = ancestorsWithinStatement.Any( - node => node is StatementSyntax - or SwitchExpressionSyntax - or ConditionalAccessExpressionSyntax + ancestor => ancestor is StatementSyntax or WhenClauseSyntax or CatchFilterClauseSyntax - || (node is ConditionalExpressionSyntax conditional + || (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)) - || (node is BinaryExpressionSyntax binary + || (ancestor is BinaryExpressionSyntax binary && (binary.IsKind(SyntaxKind.LogicalAndExpression) || binary.IsKind(SyntaxKind.LogicalOrExpression) || binary.IsKind(SyntaxKind.CoalesceExpression)) @@ -963,7 +1041,7 @@ bool DescendIntoChildren(SyntaxNode child) => if (argument.SpanStart > afterPosition && argument.SpanStart < beforePosition && (argument.RefKindKeyword.IsKind(SyntaxKind.RefKeyword) || argument.RefKindKeyword.IsKind(SyntaxKind.OutKeyword)) - && IsOneOfSymbols(context, argument.Expression, taskSymbols)) + && MayAliasTaskStorage(context, argument.Expression, taskSymbols)) { return true; } @@ -1024,7 +1102,46 @@ private static bool IsAssignmentToTask(SyntaxNodeAnalysisContext context, Expres return tuple.Arguments.Any(argument => IsAssignmentToTask(context, argument.Expression, taskSymbols)); } - return IsOneOfSymbols(context, 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.GetSymbolInfo(invocation, context.CancellationToken).Symbol is IMethodSymbol method + && (method.ReturnsByRef || method.ReturnsByRefReadonly)) + { + for (int i = 0; i < invocation.ArgumentList.Arguments.Count && i < method.Parameters.Length; i++) + { + if (method.Parameters[i].RefKind != RefKind.None + && MayAliasTaskStorage(context, invocation.ArgumentList.Arguments[i].Expression, taskSymbols)) + { + return true; + } + } + } + + return false; } private static bool IsOneOfSymbols(SyntaxNodeAnalysisContext context, ExpressionSyntax expression, IImmutableSet symbols) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index fdd14399f..8b2704192 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -70,7 +70,7 @@ public override void Initialize(AnalysisContext context) compilationContext.Options, CommonInterest.FileNamePatternForSyncMethodsToExcludeFromVSTHRD103, compilationContext.CancellationToken).ToImmutableArray(); - if (taskSymbol is object) + if (taskSymbol is object || !configuredSyncBlockingMethods.IsEmpty) { compilationContext.RegisterCodeBlockStartAction(codeBlockContext => { @@ -80,9 +80,12 @@ public override void Initialize(AnalysisContext context) { bool analyzeWholeCodeBlock = propertySymbol is object || !methodSymbol!.HasAsyncCompatibleReturnType(); codeBlockContext.RegisterSyntaxNodeAction( - Utils.DebuggableWrapper(c => AnalyzeInvocation(c, configuredSyncBlockingMethods, methodsExcludedFromVSTHRD103, analyzeWholeCodeBlock)), + Utils.DebuggableWrapper(c => AnalyzeInvocation(c, configuredSyncBlockingMethods, methodsExcludedFromVSTHRD103, analyzeWholeCodeBlock, taskSymbol is object)), SyntaxKind.InvocationExpression); - codeBlockContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(c => AnalyzeMemberAccess(c, analyzeWholeCodeBlock)), SyntaxKind.SimpleMemberAccessExpression); + if (taskSymbol is object) + { + codeBlockContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(c => AnalyzeMemberAccess(c, analyzeWholeCodeBlock)), SyntaxKind.SimpleMemberAccessExpression); + } } }); } @@ -126,10 +129,11 @@ private static void AnalyzeInvocation( SyntaxNodeAnalysisContext context, ImmutableArray configuredSyncBlockingMethods, ImmutableArray methodsExcludedFromVSTHRD103, - bool analyzeWholeCodeBlock) + bool analyzeWholeCodeBlock, + bool analyzeBuiltInBlockingMethods) { var invocationExpressionSyntax = (InvocationExpressionSyntax)context.Node; - if (ShouldAnalyze(context, analyzeWholeCodeBlock)) + if (analyzeBuiltInBlockingMethods && ShouldAnalyze(context, analyzeWholeCodeBlock)) { InspectMemberAccess( context, @@ -215,10 +219,27 @@ private static bool HasAsyncAlternative( } private static bool HasSupersetOfParameterTypes(IMethodSymbol candidateMethod, IMethodSymbol baselineMethod) - => baselineMethod.Parameters.Length <= candidateMethod.Parameters.Length - && baselineMethod.Parameters.All( - baselineParameter => candidateMethod.Parameters.Any( - candidateParameter => SymbolEqualityComparer.Default.Equals(baselineParameter.Type, candidateParameter.Type))); + { + if (baselineMethod.Parameters.Length > candidateMethod.Parameters.Length) + { + return false; + } + + var remainingCandidateTypes = candidateMethod.Parameters.Select(parameter => parameter.Type).ToList(); + foreach (IParameterSymbol baselineParameter in baselineMethod.Parameters) + { + int match = remainingCandidateTypes.FindIndex( + candidateType => SymbolEqualityComparer.Default.Equals(baselineParameter.Type, candidateType)); + if (match < 0) + { + return false; + } + + remainingCandidateTypes.RemoveAt(match); + } + + return true; + } private static bool IsInTaskReturningMethodOrDelegate(SyntaxNodeAnalysisContext context) { diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs index 6439b536f..27aa43202 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs @@ -124,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; @@ -202,7 +213,20 @@ private static bool HasSupersetOfParameterTypes(IMethodSymbol candidateMethod, I return false; } - return baselineMethod.Parameters.All(baselineParameter => candidateMethod.Parameters.Any(candidateParameter => baselineParameter.Type?.Equals(candidateParameter.Type, SymbolEqualityComparer.Default) ?? false)); + var remainingCandidateTypes = candidateMethod.Parameters.Select(parameter => parameter.Type).ToList(); + foreach (IParameterSymbol baselineParameter in baselineMethod.Parameters) + { + int match = remainingCandidateTypes.FindIndex( + candidateType => SymbolEqualityComparer.Default.Equals(baselineParameter.Type, candidateType)); + if (match < 0) + { + return false; + } + + remainingCandidateTypes.RemoveAt(match); + } + + return true; } private static bool IsInTaskReturningMethodOrDelegate(SyntaxNodeAnalysisContext context) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index 1f89d3745..b8e8595d0 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -47,8 +47,7 @@ public override async Task RegisterCodeFixesAsync(CodeFixContext context) MethodDeclarationSyntax? containingMethod = target.FirstAncestorOrSelf(); if (semanticModel is null || containingMethod is null - || target.Ancestors().TakeWhile(node => node != containingMethod) - .Any(node => node is LockStatementSyntax or CatchFilterClauseSyntax) + || 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)) @@ -92,37 +91,135 @@ private static async Task CanConvertToAsyncAsync( MethodDeclarationSyntax method, CancellationToken cancellationToken) { + if (!IsMethodLocallyConvertible(semanticModel, method, cancellationToken, out IMethodSymbol? methodSymbol)) + { + return false; + } + if (method.Modifiers.Any(SyntaxKind.AsyncKeyword)) { return true; } - IMethodSymbol? methodSymbol = semanticModel.GetDeclaredSymbol(method, cancellationToken); - if (methodSymbol is null - || methodSymbol.Parameters.Any(parameter => parameter.RefKind != RefKind.None || parameter.Type.IsRefLikeType) - || methodSymbol.ReturnsByRef - || methodSymbol.ReturnsByRefReadonly - || methodSymbol.ReturnType.IsRefLikeType - || method.DescendantNodes( + 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( + && !method.DescendantNodes( node => node is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax) - .OfType() - .Any(variable => semanticModel.GetDeclaredSymbol(variable, cancellationToken) is ILocalSymbol local - && (local.RefKind != RefKind.None || local.Type.IsRefLikeType))) + .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.IsVirtual + && !methodSymbol.IsOverride + && !methodSymbol.FindInterfacesImplemented().Any(); + + private static async Task CanConvertCallerChainAsync( + Solution solution, + IMethodSymbol method, + HashSet visitedMethods, + CancellationToken cancellationToken) + { + if (!visitedMethods.Add(method.OriginalDefinition)) { - return false; + return true; } - bool changesContract = !methodSymbol.HasAsyncCompatibleReturnType(); - return !changesContract - || (!method.Modifiers.Any(SyntaxKind.PartialKeyword) - && !methodSymbol.IsVirtual - && !methodSymbol.IsOverride - && !methodSymbol.FindInterfacesImplemented().Any() - && !await HasMethodGroupReferenceAsync(document.Project.Solution, methodSymbol, cancellationToken).ConfigureAwait(false)); + 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)) + { + 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.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( diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 12dc15207..15319ffba 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -1202,6 +1202,66 @@ Task FAsync(Task task) { await verifyTest.RunAsync(); } + [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 ConfiguredSyncBlockingMethodRequiresOneToOneAsyncParameterMatches() + { + string test = """ + using System.Threading.Tasks; + + class CustomWaiter + { + internal void Join(int first, int second) { } + + internal Task JoinAsync(int first, string second) => Task.CompletedTask; + } + + class Test + { + Task FAsync(CustomWaiter waiter) + { + waiter.[|Join|](1, 2); + return Task.CompletedTask; + } + } + """; + + var verifyTest = new CSVerify.Test + { + TestCode = test, + FixedCode = test, + }; + verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncBlockingMethods.txt", "[CustomWaiter]::Join")); + await verifyTest.RunAsync(); + } + [Fact] public async Task ConfiguredGenericExtensionReceiverDoesNotThrow() { @@ -1463,6 +1523,11 @@ int RefLikeLocal(Task task) { 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) { @@ -1476,6 +1541,10 @@ void MethodGroup(Task task) { void UseMethodGroup() { Action action = MethodGroup; } + + static void Create(out RefLike value) { + value = default; + } } "; @@ -1505,6 +1574,86 @@ void F(object gate) { 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 CodeFixIsNotOfferedWhenChangingMethodContract() { diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs index 41666852b..f2f4e878c 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs @@ -1423,6 +1423,260 @@ await CSVerify.VerifyAnalyzerAsync( CSVerify.Diagnostic(DescriptorNoAlternativeMethod).WithLocation(0).WithArguments("Result")); } + [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 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 AsyncAlternativeRequiresOneToOneParameterMatches() + { + string test = """ + using System.Threading.Tasks; + + class CustomWaiter + { + internal void Join(int first, int second) { } + + internal Task JoinAsync(int first, string second) => Task.CompletedTask; + } + + class Test + { + Task FAsync(CustomWaiter waiter) + { + waiter.Join(1, 2); + return Task.CompletedTask; + } + } + """; + + await CSVerify.VerifyAnalyzerAsync(test); + } + [Fact] public async Task TaskGetAwaiterGetResultInTaskReturningMethodGeneratesWarning() { From 9c60d57ae56b844e9e67a765f3bbb725c838ae04 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 21:32:24 -0600 Subject: [PATCH 22/26] Validate async alternative applicability Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 49 +++++++++++++++++++ .../VSTHRD002UseJtfRunAnalyzer.cs | 25 +--------- .../VSTHRD103UseAsyncOptionAnalyzer.cs | 33 +------------ .../VSTHRD002UseJtfRunAnalyzerTests.cs | 33 ++++++++++--- .../VSTHRD103UseAsyncOptionAnalyzerTests.cs | 46 +++++++++++++++-- 5 files changed, 119 insertions(+), 67 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index ee421ad10..8b2206e56 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -287,6 +287,55 @@ internal static void InspectMemberAccess( return context.SemanticModel.GetSymbolInfo(receiver, context.CancellationToken).Symbol; } + /// + /// Determines whether an async alternative is applicable to the arguments of a synchronous invocation. + /// + internal static bool IsApplicableAsyncAlternative( + SyntaxNodeAnalysisContext context, + InvocationExpressionSyntax invocation, + IMethodSymbol candidateMethod) + { + 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; + } + + 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. /// diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index 8b2704192..54b19ec58 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -215,30 +215,7 @@ private static bool HasAsyncAlternative( .Any(candidate => !candidate.IsObsolete() && candidate.Name != declaringMethodName && candidate.HasAsyncCompatibleReturnType() - && HasSupersetOfParameterTypes(candidate, invokedMethod)); - } - - private static bool HasSupersetOfParameterTypes(IMethodSymbol candidateMethod, IMethodSymbol baselineMethod) - { - if (baselineMethod.Parameters.Length > candidateMethod.Parameters.Length) - { - return false; - } - - var remainingCandidateTypes = candidateMethod.Parameters.Select(parameter => parameter.Type).ToList(); - foreach (IParameterSymbol baselineParameter in baselineMethod.Parameters) - { - int match = remainingCandidateTypes.FindIndex( - candidateType => SymbolEqualityComparer.Default.Equals(baselineParameter.Type, candidateType)); - if (match < 0) - { - return false; - } - - remainingCandidateTypes.RemoveAt(match); - } - - return true; + && CSharpCommonInterest.IsApplicableAsyncAlternative(context, invocation, candidate)); } private static bool IsInTaskReturningMethodOrDelegate(SyntaxNodeAnalysisContext context) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs index 27aa43202..3d76461be 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD103UseAsyncOptionAnalyzer.cs @@ -169,7 +169,7 @@ MemberBindingExpressionSyntax memberBinding when invocationExpressionSyntax.Firs foreach (IMethodSymbol m in symbols.OfType()) { if (!m.IsObsolete() - && HasSupersetOfParameterTypes(m, methodSymbol) + && CSharpCommonInterest.IsApplicableAsyncAlternative(context, invocationExpressionSyntax, m) && m.Name != invocationDeclaringMethod?.Identifier.Text && m.HasAsyncCompatibleReturnType()) { @@ -198,37 +198,6 @@ MemberBindingExpressionSyntax memberBinding when invocationExpressionSyntax.Firs } } - /// - /// 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; - } - - var remainingCandidateTypes = candidateMethod.Parameters.Select(parameter => parameter.Type).ToList(); - foreach (IParameterSymbol baselineParameter in baselineMethod.Parameters) - { - int match = remainingCandidateTypes.FindIndex( - candidateType => SymbolEqualityComparer.Default.Equals(baselineParameter.Type, candidateType)); - if (match < 0) - { - return false; - } - - remainingCandidateTypes.RemoveAt(match); - } - - return true; - } - private static bool IsInTaskReturningMethodOrDelegate(SyntaxNodeAnalysisContext context) { // We want to scan invocations that occur inside Task and Task-returning delegates or methods. diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 15319ffba..150659916 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -1231,23 +1231,39 @@ void F(CustomWaiter waiter) } [Fact] - public async Task ConfiguredSyncBlockingMethodRequiresOneToOneAsyncParameterMatches() + public async Task ConfiguredSyncBlockingMethodRequiresApplicableAsyncAlternative() { string test = """ using System.Threading.Tasks; class CustomWaiter { - internal void Join(int first, int second) { } + internal void Join(int value) { } - internal Task JoinAsync(int first, string second) => Task.CompletedTask; + 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) + Task FAsync(CustomWaiter waiter, ReorderedWaiter reordered, OptionalWaiter optional) { - waiter.[|Join|](1, 2); + waiter.[|Join|](1); + reordered.[|Join|](1, ""); + optional.Join(1); return Task.CompletedTask; } } @@ -1258,7 +1274,12 @@ Task FAsync(CustomWaiter waiter) TestCode = test, FixedCode = test, }; - verifyTest.TestState.AdditionalFiles.Add(("vs-threading.SyncBlockingMethods.txt", "[CustomWaiter]::Join")); + verifyTest.TestState.AdditionalFiles.Add( + ("vs-threading.SyncBlockingMethods.txt", """ + [CustomWaiter]::Join + [ReorderedWaiter]::Join + [OptionalWaiter]::Join + """)); await verifyTest.RunAsync(); } diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs index f2f4e878c..f17aaf187 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs @@ -1652,23 +1652,31 @@ await CSVerify.VerifyAnalyzerAsync( } [Fact] - public async Task AsyncAlternativeRequiresOneToOneParameterMatches() + public async Task AsyncAlternativeMustBeApplicableToInvocation() { string test = """ using System.Threading.Tasks; class CustomWaiter { - internal void Join(int first, int second) { } + internal void Join(int value) { } - internal Task JoinAsync(int first, string second) => Task.CompletedTask; + 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) + Task FAsync(CustomWaiter waiter, ReorderedWaiter reordered) { - waiter.Join(1, 2); + waiter.Join(1); + reordered.Join(1, ""); return Task.CompletedTask; } } @@ -1677,6 +1685,34 @@ Task FAsync(CustomWaiter waiter) 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() { From 87226fedfd1ff22225c64f545bbef95a5a128748 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 21:54:47 -0600 Subject: [PATCH 23/26] Harden VSTHRD002 contract conversion Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 87 ++++++++++++------- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 61 ++++++++++++- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 79 +++++++++++++++++ .../VSTHRD103UseAsyncOptionAnalyzerTests.cs | 33 +++++++ 4 files changed, 230 insertions(+), 30 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index 8b2206e56..f74d40434 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -682,48 +682,77 @@ or BaseMethodDeclarationSyntax bool DescendIntoChildren(SyntaxNode child) => child == searchRoot || child is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; - var taskAndCopySymbols = new HashSet(taskSymbols, SymbolEqualityComparer.Default); - bool IsTaskOrCopy(ExpressionSyntax expression) + bool IsDefinitelyExecutedBefore(SyntaxNode candidate, AwaitExpressionSyntax awaitExpression) { - ISymbol? expressionSymbol = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(expression), context.CancellationToken).Symbol; - return expressionSymbol is object && taskAndCopySymbols.Contains(expressionSymbol); - } + StatementSyntax? awaitStatement = awaitExpression.FirstAncestorOrSelf(); + if (awaitStatement?.Parent is not BlockSyntax block) + { + return false; + } - bool addedCopy; - do - { - addedCopy = false; - foreach (VariableDeclaratorSyntax variable in searchRoot.DescendantNodes(DescendIntoChildren) - .OfType() - .Where(variable => variable.SpanStart < accessSyntax.SpanStart && variable.Initializer is object)) + StatementSyntax? candidateStatement = candidate.FirstAncestorOrSelf(); + while (candidateStatement is object && candidateStatement.Parent != block) { - if (context.SemanticModel.GetDeclaredSymbol(variable, context.CancellationToken) is ILocalSymbol local - && IsTaskOrCopy(variable.Initializer!.Value) - && taskAndCopySymbols.Add(local)) + if (candidateStatement.Parent is not BlockSyntax containingBlock) { - addedCopy = true; + return false; } + + candidateStatement = containingBlock; + } + + return candidateStatement is object + && block.Statements.IndexOf(candidateStatement) < block.Statements.IndexOf(awaitStatement); + } + + foreach (AwaitExpressionSyntax awaitExpression in searchRoot.DescendantNodes(DescendIntoChildren) + .OfType() + .Where(awaitExpression => awaitExpression.SpanStart < accessSyntax.SpanStart) + .OrderBy(awaitExpression => awaitExpression.SpanStart)) + { + var taskAndCopySymbols = new HashSet(taskSymbols, SymbolEqualityComparer.Default); + bool IsTaskOrCopy(ExpressionSyntax expression) + { + ISymbol? expressionSymbol = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(expression), context.CancellationToken).Symbol; + return expressionSymbol is object && taskAndCopySymbols.Contains(expressionSymbol); } - foreach (AssignmentExpressionSyntax assignment in searchRoot.DescendantNodes(DescendIntoChildren) - .OfType() - .Where(assignment => assignment.SpanStart < accessSyntax.SpanStart)) + IEnumerable copyOperations = searchRoot.DescendantNodes(DescendIntoChildren) + .Where(node => node.SpanStart < awaitExpression.SpanStart + && node is VariableDeclaratorSyntax or AssignmentExpressionSyntax) + .OrderBy(node => node.SpanStart); + foreach (SyntaxNode copyOperation in copyOperations) { - if (context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol is ILocalSymbol local - && IsTaskOrCopy(assignment.Right) - && taskAndCopySymbols.Add(local)) + 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 ILocalSymbol assignedLocal) { - addedCopy = true; + if (IsTaskOrCopy(assignment.Right)) + { + taskAndCopySymbols.Add(assignedLocal); + } + else if (IsDefinitelyExecutedBefore(assignment, awaitExpression)) + { + taskAndCopySymbols.Remove(assignedLocal); + } } } + + if (AwaitCompletesTask( + context, + awaitExpression, + taskAndCopySymbols.ToImmutableHashSet(SymbolEqualityComparer.Default))) + { + return true; + } } - while (addedCopy); - IImmutableSet taskAndCopySymbolSet = taskAndCopySymbols.ToImmutableHashSet(SymbolEqualityComparer.Default); - return searchRoot.DescendantNodes(DescendIntoChildren) - .OfType() - .Any(awaitExpression => awaitExpression.SpanStart < accessSyntax.SpanStart - && AwaitCompletesTask(context, awaitExpression, taskAndCopySymbolSet)); + return false; } private static bool ContainsPotentialControlFlowBypass(SyntaxNode accessSyntax) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index b8e8595d0..9186eac8f 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -166,9 +166,60 @@ private static bool IsUnsupportedLocal(ISymbol? symbol) 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(); + && !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.Zip(method.Parameters, ParametersHaveEquivalentSignatures).All(match => match)); + } + + 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, @@ -209,6 +260,14 @@ private static async Task CanConvertCallerChainAsync( return false; } + if (callingMethodSymbol.ReturnType is INamedTypeSymbol { Arity: 0 } nonGenericReturnType + && nonGenericReturnType.IsAsyncCompatibleReturnType() + && invocation.FirstAncestorOrSelf() is { Expression: { } returnExpression } + && returnExpression.FullSpan.Contains(invocation.Span)) + { + return false; + } + if (!callingMethodSymbol.HasAsyncCompatibleReturnType() && (!CanChangeMethodContract(callingMethod, callingMethodSymbol) || await HasMethodGroupReferenceAsync(solution, callingMethodSymbol, cancellationToken).ConfigureAwait(false) diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 150659916..1fdee3b7f 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -1675,6 +1675,62 @@ static int CallerWithRefLikeLocal(Task task) 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 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 CodeFixIsNotOfferedWhenChangingMethodContract() { @@ -1707,6 +1763,29 @@ public int InterfaceMethod(Task task) { 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() { diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs index f17aaf187..78fcaa4b7 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs @@ -1569,6 +1569,39 @@ async Task GetResultAsync(ValueTask task) 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 RefConditionalAliasMutationAfterAwait_GeneratesWarning() { From 394ae0a437be1250c75127c6ff8d460f7bda1b5d Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 22:09:33 -0600 Subject: [PATCH 24/26] Cover additional VSTHRD002 flow cases Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 126 +++++++++++++++++- .../VSTHRD002UseJtfRunAnalyzer.cs | 19 ++- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 80 +++++++++++ .../VSTHRD103UseAsyncOptionAnalyzerTests.cs | 83 ++++++++++++ 4 files changed, 297 insertions(+), 11 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index f74d40434..f148f8f14 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -278,6 +278,50 @@ internal static void InspectMemberAccess( } } + internal static void InspectMemberBinding( + SyntaxNodeAnalysisContext context, + MemberBindingExpressionSyntax memberBinding, + ExpressionSyntax receiver, + SyntaxNode accessSyntax, + DiagnosticDescriptor descriptor, + IEnumerable problematicMethods) + { + 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 (memberBinding.Name.Identifier.ValueText == item.Method.Name + && receiverType.Name == item.Method.ContainingType.Name + && receiverType.BelongsToNamespace(item.Method.ContainingType.Namespace) + && IsBuiltInBlockingMember(context, accessedSymbol, item.Method)) + { + if (HasTaskCompleted(context, receiver, accessSyntax)) + { + return; + } + + context.ReportDiagnostic(Diagnostic.Create(descriptor, memberBinding.Name.GetLocation())); + } + } + } + /// /// Gets the symbol represented by the normalized task-like receiver of a blocking member access. /// @@ -713,8 +757,17 @@ bool IsDefinitelyExecutedBefore(SyntaxNode candidate, AwaitExpressionSyntax awai var taskAndCopySymbols = new HashSet(taskSymbols, SymbolEqualityComparer.Default); bool IsTaskOrCopy(ExpressionSyntax expression) { - ISymbol? expressionSymbol = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(expression), context.CancellationToken).Symbol; - return expressionSymbol is object && taskAndCopySymbols.Contains(expressionSymbol); + 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 + && IsSupportedConfigureAwaitInvocation(context, configureAwaitInvocation) + && IsTaskOrCopy(configureAwaitAccess.Expression); } IEnumerable copyOperations = searchRoot.DescendantNodes(DescendIntoChildren) @@ -730,15 +783,16 @@ bool IsTaskOrCopy(ExpressionSyntax expression) taskAndCopySymbols.Add(declaredLocal); } else if (copyOperation is AssignmentExpressionSyntax assignment - && context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol is ILocalSymbol assignedLocal) + && context.SemanticModel.GetSymbolInfo(assignment.Left, context.CancellationToken).Symbol is ISymbol assignedSymbol + && assignedSymbol is ILocalSymbol or IParameterSymbol) { if (IsTaskOrCopy(assignment.Right)) { - taskAndCopySymbols.Add(assignedLocal); + taskAndCopySymbols.Add(assignedSymbol); } else if (IsDefinitelyExecutedBefore(assignment, awaitExpression)) { - taskAndCopySymbols.Remove(assignedLocal); + taskAndCopySymbols.Remove(assignedSymbol); } } } @@ -1039,7 +1093,10 @@ or CatchFilterClauseSyntax && (binary.IsKind(SyntaxKind.LogicalAndExpression) || binary.IsKind(SyntaxKind.LogicalOrExpression) || binary.IsKind(SyntaxKind.CoalesceExpression)) - && binary.Right.FullSpan.Contains(candidate.Span))); + && 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; @@ -1085,12 +1142,67 @@ private static bool AwaitCompletesTask( || (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)))); + && 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)) + { + 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); diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs index 54b19ec58..09c179b8e 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/VSTHRD002UseJtfRunAnalyzer.cs @@ -135,10 +135,21 @@ private static void AnalyzeInvocation( var invocationExpressionSyntax = (InvocationExpressionSyntax)context.Node; if (analyzeBuiltInBlockingMethods && ShouldAnalyze(context, analyzeWholeCodeBlock)) { - InspectMemberAccess( - context, - invocationExpressionSyntax.Expression as MemberAccessExpressionSyntax, - CommonInterest.ProblematicSyncBlockingMethods); + 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 diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 1fdee3b7f..9039262c6 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -220,6 +220,68 @@ 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_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() { @@ -1202,6 +1264,24 @@ Task FAsync(Task task) { 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() { diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs index 78fcaa4b7..f5071028e 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs @@ -1602,6 +1602,89 @@ async Task GetResultAsync(ValueTask task, ValueTask other) }.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 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() { From bdb808ff98877ee5622ce0138e1973a92916bc49 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 22:53:58 -0600 Subject: [PATCH 25/26] Harden VSTHRD002 flow analysis Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../CSharpCommonInterest.cs | 233 +++++++++++++++--- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 27 +- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 162 ++++++++++++ .../VSTHRD103UseAsyncOptionAnalyzerTests.cs | 153 ++++++++++++ 4 files changed, 534 insertions(+), 41 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs index f148f8f14..99c2f6b1f 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CSharp/CSharpCommonInterest.cs @@ -12,6 +12,7 @@ using Microsoft.CodeAnalysis.CSharp; using Microsoft.CodeAnalysis.CSharp.Syntax; using Microsoft.CodeAnalysis.Diagnostics; +using Microsoft.CodeAnalysis.Operations; namespace Microsoft.VisualStudio.Threading.Analyzers; @@ -278,6 +279,15 @@ internal static void InspectMemberAccess( } } + /// + /// 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, @@ -403,16 +413,23 @@ private static bool HasTaskCompletedInContinuation( { foreach (AnonymousFunctionExpressionSyntax anonymousFunction in accessSyntax.Ancestors().OfType()) { - if (anonymousFunction.Parent is not ArgumentSyntax anonymousFunctionArgument + 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 - || continuationInvocation.ArgumentList.Arguments.FirstOrDefault() != anonymousFunctionArgument) + || 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; } - if (context.SemanticModel.GetSymbolInfo(continuationInvocation, context.CancellationToken).Symbol is not IMethodSymbol invokedMethod - || invokedMethod.Name != nameof(Task.ContinueWith) - || !Utils.IsTask(invokedMethod.ContainingType)) + if (continuationOperation.TargetMethod.Name != nameof(Task.ContinueWith) + || !Utils.IsTask(continuationOperation.TargetMethod.ContainingType)) { continue; } @@ -435,17 +452,19 @@ private static bool HasTaskCompletedInContinuation( if (accessSyntax.Ancestors().TakeWhile(node => node != anonymousFunction) .Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax)) { - potentialTaskSymbols = potentialTaskSymbols.Union(GetSymbolAndRefAliases( + (ImmutableHashSet outerDefiniteAliases, ImmutableHashSet outerPotentialAliases) = GetSymbolAndRefAliases( context, accessSyntax, completedTask, anonymousFunction, - includeAllCandidates: true).Potential); + includeAllCandidates: true); + taskSymbols = taskSymbols.Union(outerDefiniteAliases); + potentialTaskSymbols = potentialTaskSymbols.Union(outerPotentialAliases); } ISymbol? receiverSymbol = context.SemanticModel.GetSymbolInfo(UnwrapParentheses(taskReceiver), context.CancellationToken).Symbol; if (receiverSymbol is object - && taskSymbols.Contains(receiverSymbol) + && (SymbolEqualityComparer.Default.Equals(receiverSymbol, completedTask) || taskSymbols.Contains(receiverSymbol)) && !IsTaskReassignedInContinuation(context, anonymousFunction, accessSyntax, potentialTaskSymbols)) { return true; @@ -633,7 +652,7 @@ private static bool HasTaskCompletedCore( if (IsWithinCompletedTaskBranch(context, accessSyntax, taskSymbols, potentialTaskSymbols)) { - return Utils.IsTask(taskType) || !MayHaveAwaitedTaskBefore(context, accessSyntax, taskSymbols); + return Utils.IsTask(taskType) || !MayHaveConsumedValueTaskBefore(context, accessSyntax, taskSymbols); } // Awaiting an IValueTaskSource-backed ValueTask consumes it, so a later Result access is not safe. @@ -711,7 +730,7 @@ or ForEachVariableStatementSyntax } } - private static bool MayHaveAwaitedTaskBefore( + private static bool MayHaveConsumedValueTaskBefore( SyntaxNodeAnalysisContext context, SyntaxNode accessSyntax, IImmutableSet taskSymbols) @@ -726,10 +745,10 @@ or BaseMethodDeclarationSyntax bool DescendIntoChildren(SyntaxNode child) => child == searchRoot || child is not AnonymousFunctionExpressionSyntax and not LocalFunctionStatementSyntax; - bool IsDefinitelyExecutedBefore(SyntaxNode candidate, AwaitExpressionSyntax awaitExpression) + bool IsDefinitelyExecutedBefore(SyntaxNode candidate, SyntaxNode consumption) { - StatementSyntax? awaitStatement = awaitExpression.FirstAncestorOrSelf(); - if (awaitStatement?.Parent is not BlockSyntax block) + StatementSyntax? consumptionStatement = consumption.FirstAncestorOrSelf(); + if (consumptionStatement?.Parent is not BlockSyntax block) { return false; } @@ -746,13 +765,10 @@ bool IsDefinitelyExecutedBefore(SyntaxNode candidate, AwaitExpressionSyntax awai } return candidateStatement is object - && block.Statements.IndexOf(candidateStatement) < block.Statements.IndexOf(awaitStatement); + && block.Statements.IndexOf(candidateStatement) < block.Statements.IndexOf(consumptionStatement); } - foreach (AwaitExpressionSyntax awaitExpression in searchRoot.DescendantNodes(DescendIntoChildren) - .OfType() - .Where(awaitExpression => awaitExpression.SpanStart < accessSyntax.SpanStart) - .OrderBy(awaitExpression => awaitExpression.SpanStart)) + ImmutableHashSet GetTaskAndCopySymbolsBefore(SyntaxNode consumption) { var taskAndCopySymbols = new HashSet(taskSymbols, SymbolEqualityComparer.Default); bool IsTaskOrCopy(ExpressionSyntax expression) @@ -766,12 +782,11 @@ bool IsTaskOrCopy(ExpressionSyntax expression) return expression is InvocationExpressionSyntax configureAwaitInvocation && configureAwaitInvocation.Expression is MemberAccessExpressionSyntax { Name.Identifier.ValueText: nameof(Task.ConfigureAwait) } configureAwaitAccess - && IsSupportedConfigureAwaitInvocation(context, configureAwaitInvocation) && IsTaskOrCopy(configureAwaitAccess.Expression); } IEnumerable copyOperations = searchRoot.DescendantNodes(DescendIntoChildren) - .Where(node => node.SpanStart < awaitExpression.SpanStart + .Where(node => node.SpanStart < consumption.SpanStart && node is VariableDeclaratorSyntax or AssignmentExpressionSyntax) .OrderBy(node => node.SpanStart); foreach (SyntaxNode copyOperation in copyOperations) @@ -790,23 +805,163 @@ bool IsTaskOrCopy(ExpressionSyntax expression) { taskAndCopySymbols.Add(assignedSymbol); } - else if (IsDefinitelyExecutedBefore(assignment, awaitExpression)) + else if (IsDefinitelyExecutedBefore(assignment, consumption)) { taskAndCopySymbols.Remove(assignedSymbol); } } } - if (AwaitCompletesTask( - context, - awaitExpression, - taskAndCopySymbols.ToImmutableHashSet(SymbolEqualityComparer.Default))) + return taskAndCopySymbols.ToImmutableHashSet(SymbolEqualityComparer.Default); + } + + 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; + } + } + + 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) @@ -1187,7 +1342,9 @@ or BaseMethodDeclarationSyntax ?? 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)) + || MayReassignTask(context, searchRoot, collectionSymbol, initializer.Span.End, awaitExpression.SpanStart) + || NestedFunctionMayReassignTask(context, awaitExpression, taskSymbols) + || NestedFunctionMayReassignTask(context, awaitExpression, collectionSymbol)) { return false; } @@ -1230,7 +1387,7 @@ bool DescendIntoChildren(SyntaxNode child) => { if (argument.SpanStart > afterPosition && argument.SpanStart < beforePosition - && (argument.RefKindKeyword.IsKind(SyntaxKind.RefKeyword) || argument.RefKindKeyword.IsKind(SyntaxKind.OutKeyword)) + && context.SemanticModel.GetOperation(argument, context.CancellationToken) is IArgumentOperation { Parameter.RefKind: RefKind.Ref or RefKind.Out } && MayAliasTaskStorage(context, argument.Expression, taskSymbols)) { return true; @@ -1318,16 +1475,22 @@ private static bool MayAliasTaskStorage( } if (expression is InvocationExpressionSyntax invocation - && context.SemanticModel.GetSymbolInfo(invocation, context.CancellationToken).Symbol is IMethodSymbol method - && (method.ReturnsByRef || method.ReturnsByRefReadonly)) + && context.SemanticModel.GetOperation(invocation, context.CancellationToken) is IInvocationOperation invocationOperation + && (invocationOperation.TargetMethod.ReturnsByRef || invocationOperation.TargetMethod.ReturnsByRefReadonly)) { - for (int i = 0; i < invocation.ArgumentList.Arguments.Count && i < method.Parameters.Length; i++) + IMethodSymbol method = invocationOperation.TargetMethod; + if ((method.ReducedFrom ?? method).Parameters is [{ RefKind: not RefKind.None }, ..] + && invocation.Expression is MemberAccessExpressionSyntax memberAccess + && MayAliasTaskStorage(context, memberAccess.Expression, taskSymbols)) { - if (method.Parameters[i].RefKind != RefKind.None - && MayAliasTaskStorage(context, invocation.ArgumentList.Arguments[i].Expression, taskSymbols)) - { - return true; - } + return true; + } + + if (invocationOperation.Arguments.Any(argument => argument.Parameter?.RefKind != RefKind.None + && argument.Syntax is ArgumentSyntax argumentSyntax + && MayAliasTaskStorage(context, argumentSyntax.Expression, taskSymbols))) + { + return true; } } diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index 9186eac8f..dcbd46763 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -66,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) @@ -85,6 +85,15 @@ 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, @@ -248,7 +257,9 @@ private static async Task CanConvertCallerChainAsync( || callingMethod is null || IsAwaitForbiddenAt(invocation, callingMethod) || invocation.Ancestors().TakeWhile(node => node != callingMethod) - .Any(node => node is AnonymousFunctionExpressionSyntax or LocalFunctionStatementSyntax)) + .Any(node => node is AnonymousFunctionExpressionSyntax + or LocalFunctionStatementSyntax + or ConditionalAccessExpressionSyntax)) { return false; } @@ -260,10 +271,14 @@ private static async Task CanConvertCallerChainAsync( return false; } - if (callingMethodSymbol.ReturnType is INamedTypeSymbol { Arity: 0 } nonGenericReturnType - && nonGenericReturnType.IsAsyncCompatibleReturnType() - && invocation.FirstAncestorOrSelf() is { Expression: { } returnExpression } - && returnExpression.FullSpan.Contains(invocation.Span)) + 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)) + && (callerReturnType.Arity == 0 + || !Utils.IsTask(callerReturnType) + || !semanticModel.Compilation.ClassifyConversion(method.ReturnType, callerReturnType.TypeArguments[0]).IsImplicit)) { return false; } diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index 9039262c6..bcda68ee9 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -240,6 +240,29 @@ async void GetResultAsync(Task task) 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() { @@ -632,6 +655,41 @@ void Guarded(ValueTask task) { 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() { @@ -694,6 +752,28 @@ void ContinueWith(Func, int> del) { } await CSVerify.VerifyCodeFixAsync(test, expected, test); } + [Fact] + 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() { @@ -1788,6 +1868,88 @@ static Task Caller(Task 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 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() { diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs index f5071028e..f34335c4f 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD103UseAsyncOptionAnalyzerTests.cs @@ -1423,6 +1423,30 @@ await CSVerify.VerifyAnalyzerAsync( 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() { @@ -1664,6 +1688,135 @@ async Task GetResultAsync(ValueTask task, ValueTask copy) 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() { From 2e2d14e647eed2fafabf27979c0e5d4cb8aeb667 Mon Sep 17 00:00:00 2001 From: Andrew Arnott Date: Fri, 21 Aug 2026 23:16:18 -0600 Subject: [PATCH 26/26] Harden VSTHRD002 caller conversion Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../VSTHRD002UseJtfRunCodeFixWithAwait.cs | 10 ++- .../VSTHRD002UseJtfRunAnalyzerTests.cs | 72 +++++++++++++++++++ 2 files changed, 76 insertions(+), 6 deletions(-) diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs index dcbd46763..d2f109ec2 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.CodeFixes/VSTHRD002UseJtfRunCodeFixWithAwait.cs @@ -192,8 +192,9 @@ private static bool HasAsyncNameCollision(IMethodSymbol method) return method.ContainingType.GetMembers(asyncName) .OfType() .Any(candidate => candidate.Arity == method.Arity - && candidate.Parameters.Length == method.Parameters.Length - && candidate.Parameters.Zip(method.Parameters, ParametersHaveEquivalentSignatures).All(match => match)); + && 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) @@ -275,10 +276,7 @@ or LocalFunctionStatementSyntax && callerReturnType.IsAsyncCompatibleReturnType() && ((invocation.FirstAncestorOrSelf() is { Expression: { } returnExpression } && returnExpression.FullSpan.Contains(invocation.Span)) - || (callingMethod.ExpressionBody?.Expression.FullSpan.Contains(invocation.Span) is true)) - && (callerReturnType.Arity == 0 - || !Utils.IsTask(callerReturnType) - || !semanticModel.Compilation.ClassifyConversion(method.ReturnType, callerReturnType.TypeArguments[0]).IsImplicit)) + || (callingMethod.ExpressionBody?.Expression.FullSpan.Contains(invocation.Span) is true))) { return false; } diff --git a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs index bcda68ee9..3f43f8548 100644 --- a/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs +++ b/test/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD002UseJtfRunAnalyzerTests.cs @@ -1926,6 +1926,55 @@ static DerivedTask GetValue(Task 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() { @@ -1973,6 +2022,29 @@ Task GetValueAsync(Task 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() {