diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD010MainThreadUsageAnalyzerTests.cs b/src/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD010MainThreadUsageAnalyzerTests.cs index 65d247a1b..895eb9201 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD010MainThreadUsageAnalyzerTests.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers.Tests/VSTHRD010MainThreadUsageAnalyzerTests.cs @@ -267,6 +267,121 @@ void VerifyOnUIThread() { this.VerifyCSharpDiagnostic(test, this.expect); } + [Fact] + public void RequiresUIThreadTransitive() + { + var test = @" +using System; +using Microsoft.VisualStudio.Shell.Interop; + +class Test { + void F() { + VerifyOnUIThread(); + IVsSolution sln = null; + sln.SetProperty(1000, null); + } + + void G() { + F(); + } + + void H() { + G(); + } + + int MainThreadGetter { + get { + H(); + return 0; + } + + set { + } + } + + int MainThreadSetter { + get => 0; + set => H(); + } + + int CallMainThreadGetter_get() => MainThreadGetter; // Flagged + int CallMainThreadGetter_get2() => this.MainThreadGetter; // Flagged + int CallMainThreadGetter_get3() => ((Test)this).MainThreadGetter; // Flagged + int CallMainThreadGetter_set() => MainThreadGetter = 1; + int CallMainThreadGetter_set2() => this.MainThreadGetter = 1; + + int CallMainThreadSetter_get() => MainThreadSetter; + int CallMainThreadSetter_get2() => this.MainThreadSetter; + int CallMainThreadSetter_set() => MainThreadSetter = 1; // Flagged + int CallMainThreadSetter_set2() => this.MainThreadSetter = 1; // Flagged + int CallMainThreadSetter_set3() => ((Test)this).MainThreadSetter = 1; // Flagged + + // None of these should produce diagnostics since we're not invoking the members. + string NameOfFoo1() => nameof(MainThreadGetter); + string NameOfFoo2() => nameof(MainThreadSetter); + string NameOfThisFoo1() => nameof(this.MainThreadGetter); + string NameOfThisFoo2() => nameof(this.MainThreadSetter); + string NameOfH() => nameof(H); + Action GAsDelegate() => this.G; + + void VerifyOnUIThread() { + } +} +"; + DiagnosticResult CreateDiagnostic(int line, int column, int endLine, int endColumn) => + new DiagnosticResult + { + Id = this.expect.Id, + Message = this.expect.Message, + SkipVerifyMessage = this.expect.SkipVerifyMessage, + Severity = this.expect.Severity, + Locations = new[] { new DiagnosticResultLocation("Test0.cs", line, column, endLine, endColumn) }, + }; + var expect = new DiagnosticResult[] + { + CreateDiagnostic(12, 10, 12, 11), + CreateDiagnostic(16, 10, 16, 11), + CreateDiagnostic(21, 9, 21, 12), + CreateDiagnostic(32, 9, 32, 12), + CreateDiagnostic(35, 9, 35, 33), + CreateDiagnostic(36, 9, 36, 34), + CreateDiagnostic(37, 9, 37, 34), + CreateDiagnostic(43, 9, 43, 33), + CreateDiagnostic(44, 9, 44, 34), + CreateDiagnostic(45, 9, 45, 34), + }; + this.VerifyCSharpDiagnostic(test, expect); + } + + [Fact] + public void RequiresUIThreadNotTransitiveIfNotExplicit() + { + var test = @" +using System; +using Microsoft.VisualStudio.Shell.Interop; + +class Test { + void F() { + IVsSolution sln = null; + sln.SetProperty(1000, null); + } + + void G() { + F(); + } + + void H() { + G(); + } + + void VerifyOnUIThread() { + } +} +"; + this.expect.Locations = new[] { new DiagnosticResultLocation("Test0.cs", 8, 13, 8, 24) }; + this.VerifyCSharpDiagnostic(test, this.expect); + } + [Fact] public void InvokeVsSolutionAfterSwitchedToMainThreadAsync() { diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.cs.xlf b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.cs.xlf index 6978353df..cb0a1f639 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.cs.xlf +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.cs.xlf @@ -192,6 +192,11 @@ Await JoinableTaskFactory.SwitchToMainThreadAsync() first. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.de.xlf b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.de.xlf index 84e993f15..d4d778617 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.de.xlf +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.de.xlf @@ -192,6 +192,11 @@ Await JoinableTaskFactory.SwitchToMainThreadAsync() first. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.es.xlf b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.es.xlf index d1d09ffa8..e322b41dc 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.es.xlf +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.es.xlf @@ -192,6 +192,11 @@ Await JoinableTaskFactory.SwitchToMainThreadAsync() first. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.fr.xlf b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.fr.xlf index 65a16a595..0a1e272d0 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.fr.xlf +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.fr.xlf @@ -192,6 +192,11 @@ Await JoinableTaskFactory.SwitchToMainThreadAsync() first. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.it.xlf b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.it.xlf index 1dde34ccc..ce821efdc 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.it.xlf +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.it.xlf @@ -192,6 +192,11 @@ Await JoinableTaskFactory.SwitchToMainThreadAsync() first. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.ja.xlf b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.ja.xlf index 5a79c4e0a..10d52b00d 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.ja.xlf +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.ja.xlf @@ -192,6 +192,11 @@ Await JoinableTaskFactory.SwitchToMainThreadAsync() first. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.ko.xlf b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.ko.xlf index da317e3a1..68e870677 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.ko.xlf +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.ko.xlf @@ -192,6 +192,11 @@ Await JoinableTaskFactory.SwitchToMainThreadAsync() first. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.pl.xlf b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.pl.xlf index 49cbfcafc..e2091f797 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.pl.xlf +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.pl.xlf @@ -192,6 +192,11 @@ Await JoinableTaskFactory.SwitchToMainThreadAsync() first. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.pt-BR.xlf b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.pt-BR.xlf index 0c1fe4c8d..896e93818 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.pt-BR.xlf +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.pt-BR.xlf @@ -192,6 +192,11 @@ Await JoinableTaskFactory.SwitchToMainThreadAsync() first. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.ru.xlf b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.ru.xlf index 91377a8c7..2fc72469f 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.ru.xlf +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.ru.xlf @@ -192,6 +192,11 @@ Await JoinableTaskFactory.SwitchToMainThreadAsync() first. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.tr.xlf b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.tr.xlf index 7f4b91f1e..5d7248d31 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.tr.xlf +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.tr.xlf @@ -192,6 +192,11 @@ Await JoinableTaskFactory.SwitchToMainThreadAsync() first. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.zh-Hans.xlf b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.zh-Hans.xlf index 469739371..5b428fc05 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.zh-Hans.xlf +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.zh-Hans.xlf @@ -192,6 +192,11 @@ Await JoinableTaskFactory.SwitchToMainThreadAsync() first. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.zh-Hant.xlf b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.zh-Hant.xlf index ab13b9857..508dffdba 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.zh-Hant.xlf +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/MultilingualResources/Microsoft.VisualStudio.Threading.Analyzers.zh-Hant.xlf @@ -192,6 +192,11 @@ Await JoinableTaskFactory.SwitchToMainThreadAsync() first. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/Strings.Designer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers/Strings.Designer.cs index 750cf9960..5c2fbb48b 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/Strings.Designer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/Strings.Designer.cs @@ -163,6 +163,15 @@ internal static string VSTHRD010_MessageFormat_NoAssertingMethod { } } + /// + /// Looks up a localized string similar to Add call to {0}() at start of member body because this member invokes other members that require the main thread.. + /// + internal static string VSTHRD010_MessageFormat_TransitiveMainThreadUser { + get { + return ResourceManager.GetString("VSTHRD010_MessageFormat_TransitiveMainThreadUser", resourceCulture); + } + } + /// /// Looks up a localized string similar to Invoke single-threaded types on Main thread. /// diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/Strings.resx b/src/Microsoft.VisualStudio.Threading.Analyzers/Strings.resx index b2f98a4d9..6c8235849 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/Strings.resx +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/Strings.resx @@ -257,4 +257,8 @@ Use AsyncLazy<T> instead. Await JoinableTaskFactory.SwitchToMainThreadAsync() first. {0} is a type name and {1} is the name of a method that throws if not called from the main thread. + + Add call to {0}() at start of member body because this member invokes other members that require the main thread. + {0} is a method name. + \ No newline at end of file diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/Utils.cs b/src/Microsoft.VisualStudio.Threading.Analyzers/Utils.cs index ae4a96cb5..c1ffb054b 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/Utils.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/Utils.cs @@ -276,6 +276,22 @@ internal static bool IsObsolete(this ISymbol symbol) return symbol.GetAttributes().Any(a => a.AttributeClass.Name == nameof(ObsoleteAttribute) && a.AttributeClass.BelongsToNamespace(Namespaces.System)); } + internal static bool IsOnLeftHandOfAssignment(SyntaxNode syntaxNode) + { + SyntaxNode parent = null; + while ((parent = syntaxNode.Parent) != null) + { + if (parent is AssignmentExpressionSyntax assignment) + { + return assignment.Left == syntaxNode; + } + + syntaxNode = parent; + } + + return false; + } + internal static IEnumerable FindInterfacesImplemented(this ISymbol symbol) { if (symbol == null) @@ -607,9 +623,9 @@ internal static NameSyntax QualifyName(IReadOnlyList qualifiers, SimpleN /// /// Determines whether an expression appears inside a C# "nameof" pseudo-method. /// - internal static bool IsWithinNameOf(ExpressionSyntax memberAccess) + internal static bool IsWithinNameOf(SyntaxNode syntaxNode) { - var invocation = memberAccess?.FirstAncestorOrSelf(); + var invocation = syntaxNode?.FirstAncestorOrSelf(); return (invocation?.Expression as IdentifierNameSyntax)?.Identifier.Text == "nameof" && invocation.ArgumentList.Arguments.Count == 1; } diff --git a/src/Microsoft.VisualStudio.Threading.Analyzers/VSTHRD010MainThreadUsageAnalyzer.cs b/src/Microsoft.VisualStudio.Threading.Analyzers/VSTHRD010MainThreadUsageAnalyzer.cs index a7c15673d..d9640513f 100644 --- a/src/Microsoft.VisualStudio.Threading.Analyzers/VSTHRD010MainThreadUsageAnalyzer.cs +++ b/src/Microsoft.VisualStudio.Threading.Analyzers/VSTHRD010MainThreadUsageAnalyzer.cs @@ -9,6 +9,7 @@ using Microsoft.CodeAnalysis.CSharp; using Microsoft.CodeAnalysis.CSharp.Syntax; using Microsoft.CodeAnalysis.Diagnostics; + using Microsoft.CodeAnalysis.Semantics; using Microsoft.CodeAnalysis.Text; /// @@ -62,6 +63,23 @@ public class VSTHRD010MainThreadUsageAnalyzer : DiagnosticAnalyzer defaultSeverity: DiagnosticSeverity.Warning, isEnabledByDefault: true); + internal static readonly DiagnosticDescriptor DescriptorTransitiveMainThreadUser = new DiagnosticDescriptor( + id: Id, + title: Strings.VSTHRD010_Title, + messageFormat: Strings.VSTHRD010_MessageFormat_TransitiveMainThreadUser, + helpLinkUri: Utils.GetHelpLink(Id), + category: "Usage", + defaultSeverity: DiagnosticSeverity.Warning, + isEnabledByDefault: true); + + /// + /// A reusable value to return from . + /// + private static readonly ImmutableArray ReusableSupportedDescriptors = ImmutableArray.Create( + Descriptor, + DescriptorNoAssertingMethod, + DescriptorTransitiveMainThreadUser); + private enum ThreadingContext { /// @@ -82,13 +100,7 @@ private enum ThreadingContext } /// - public override ImmutableArray SupportedDiagnostics - { - get - { - return ImmutableArray.Create(Descriptor); - } - } + public override ImmutableArray SupportedDiagnostics => ReusableSupportedDescriptors; /// public override void Initialize(AnalysisContext context) @@ -96,29 +108,154 @@ public override void Initialize(AnalysisContext context) context.EnableConcurrentExecution(); context.ConfigureGeneratedCodeAnalysis(GeneratedCodeAnalysisFlags.Analyze); - context.RegisterCompilationStartAction(ctxt => + context.RegisterCompilationStartAction(compilationStartContext => { - var mainThreadAssertingMethods = CommonInterest.ReadMethods(ctxt, CommonInterest.FileNamePatternForMethodsThatAssertMainThread).ToImmutableArray(); - var mainThreadSwitchingMethods = CommonInterest.ReadMethods(ctxt, CommonInterest.FileNamePatternForMethodsThatSwitchToMainThread).ToImmutableArray(); - var typesRequiringMainThread = CommonInterest.ReadTypes(ctxt, CommonInterest.FileNamePatternForTypesRequiringMainThread).ToImmutableArray(); + var mainThreadAssertingMethods = CommonInterest.ReadMethods(compilationStartContext, CommonInterest.FileNamePatternForMethodsThatAssertMainThread).ToImmutableArray(); + var mainThreadSwitchingMethods = CommonInterest.ReadMethods(compilationStartContext, CommonInterest.FileNamePatternForMethodsThatSwitchToMainThread).ToImmutableArray(); + var typesRequiringMainThread = CommonInterest.ReadTypes(compilationStartContext, CommonInterest.FileNamePatternForTypesRequiringMainThread).ToImmutableArray(); + + var methodsDeclaringUIThreadRequirement = new HashSet(); + var callerToCalleeMap = new Dictionary>(); - ctxt.RegisterCodeBlockStartAction(ctxt2 => + compilationStartContext.RegisterCodeBlockStartAction(codeBlockStartContext => { var methodAnalyzer = new MethodAnalyzer { MainThreadAssertingMethods = mainThreadAssertingMethods, MainThreadSwitchingMethods = mainThreadSwitchingMethods, TypesRequiringMainThread = typesRequiringMainThread, + MethodsDeclaringUIThreadRequirement = methodsDeclaringUIThreadRequirement, }; - ctxt2.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(methodAnalyzer.AnalyzeInvocation), SyntaxKind.InvocationExpression); - ctxt2.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(methodAnalyzer.AnalyzeMemberAccess), SyntaxKind.SimpleMemberAccessExpression); - ctxt2.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(methodAnalyzer.AnalyzeCast), SyntaxKind.CastExpression); - ctxt2.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(methodAnalyzer.AnalyzeAs), SyntaxKind.AsExpression); - ctxt2.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(methodAnalyzer.AnalyzeAs), SyntaxKind.IsExpression); + codeBlockStartContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(methodAnalyzer.AnalyzeInvocation), SyntaxKind.InvocationExpression); + codeBlockStartContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(methodAnalyzer.AnalyzeMemberAccess), SyntaxKind.SimpleMemberAccessExpression); + codeBlockStartContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(methodAnalyzer.AnalyzeCast), SyntaxKind.CastExpression); + codeBlockStartContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(methodAnalyzer.AnalyzeAs), SyntaxKind.AsExpression); + codeBlockStartContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(methodAnalyzer.AnalyzeAs), SyntaxKind.IsExpression); + }); + + compilationStartContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(c => AddToCallerCalleeMap(c, callerToCalleeMap)), SyntaxKind.InvocationExpression); + compilationStartContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(c => AddToCallerCalleeMap(c, callerToCalleeMap)), SyntaxKind.SimpleMemberAccessExpression); + compilationStartContext.RegisterSyntaxNodeAction(Utils.DebuggableWrapper(c => AddToCallerCalleeMap(c, callerToCalleeMap)), SyntaxKind.IdentifierName); + + compilationStartContext.RegisterCompilationEndAction(compilationEndContext => + { + var calleeToCallerMap = CreateCalleeToCallerMap(callerToCalleeMap); + var transitiveClosureOfMainThreadRequiringMethods = GetTransitiveClosureOfMainThreadRequiringMethods(methodsDeclaringUIThreadRequirement, calleeToCallerMap); + foreach (var implicitUserMethod in transitiveClosureOfMainThreadRequiringMethods.Except(methodsDeclaringUIThreadRequirement)) + { + var declarationSyntax = implicitUserMethod.DeclaringSyntaxReferences.FirstOrDefault()?.GetSyntax(compilationEndContext.CancellationToken); + SyntaxToken memberNameSyntax = default(SyntaxToken); + switch (declarationSyntax) + { + case MethodDeclarationSyntax methodDeclarationSyntax: + memberNameSyntax = methodDeclarationSyntax.Identifier; + break; + case AccessorDeclarationSyntax accessorDeclarationSyntax: + memberNameSyntax = accessorDeclarationSyntax.Keyword; + break; + } + var location = memberNameSyntax.GetLocation(); + if (location != null) + { + var exampleAssertingMethod = mainThreadAssertingMethods.FirstOrDefault(); + compilationEndContext.ReportDiagnostic(Diagnostic.Create(DescriptorTransitiveMainThreadUser, location, exampleAssertingMethod)); + } + } }); }); } + private static HashSet GetTransitiveClosureOfMainThreadRequiringMethods(HashSet methodsRequiringUIThread, Dictionary> calleeToCallerMap) + { + var result = new HashSet(); + + void MarkMethod(IMethodSymbol method) + { + if (result.Add(method) && calleeToCallerMap.TryGetValue(method, out var callers)) + { + foreach (var caller in callers) + { + MarkMethod(caller); + } + } + } + + foreach (var method in methodsRequiringUIThread) + { + MarkMethod(method); + } + + return result; + } + + private static void AddToCallerCalleeMap(SyntaxNodeAnalysisContext context, Dictionary> callerToCalleeMap) + { + if (Utils.IsWithinNameOf(context.Node)) + { + return; + } + + IMethodSymbol GetPropertyAccessor(IPropertySymbol propertySymbol) + { + if (propertySymbol != null) + { + return Utils.IsOnLeftHandOfAssignment(context.Node) + ? propertySymbol.SetMethod + : propertySymbol.GetMethod; + } + + return null; + } + + ISymbol targetMethod = null; + switch (context.Node) + { + case InvocationExpressionSyntax invocationExpressionSyntax: + targetMethod = context.SemanticModel.GetSymbolInfo(invocationExpressionSyntax.Expression).Symbol; + break; + case MemberAccessExpressionSyntax memberAccessExpressionSyntax: + targetMethod = GetPropertyAccessor(context.SemanticModel.GetSymbolInfo(memberAccessExpressionSyntax.Name).Symbol as IPropertySymbol); + break; + case IdentifierNameSyntax identifierNameSyntax: + targetMethod = GetPropertyAccessor(context.SemanticModel.GetSymbolInfo(identifierNameSyntax).Symbol as IPropertySymbol); + break; + } + + if (context.ContainingSymbol is IMethodSymbol caller && targetMethod is IMethodSymbol callee) + { + lock (callerToCalleeMap) + { + if (!callerToCalleeMap.TryGetValue(caller, out HashSet callees)) + { + callerToCalleeMap[caller] = callees = new HashSet(); + } + + callees.Add(callee); + } + } + } + + private static Dictionary> CreateCalleeToCallerMap(Dictionary> callerToCalleeMap) + { + var result = new Dictionary>(); + + foreach (var item in callerToCalleeMap) + { + var caller = item.Key; + foreach (var callee in item.Value) + { + if (!result.TryGetValue(callee, out var callers)) + { + result[callee] = callers = new HashSet(); + } + + callers.Add(caller); + } + } + + return result; + } + private class MethodAnalyzer { private ImmutableDictionary methodDeclarationNodes = ImmutableDictionary.Empty; @@ -129,6 +266,8 @@ private class MethodAnalyzer internal ImmutableArray TypesRequiringMainThread { get; set; } + internal HashSet MethodsDeclaringUIThreadRequirement { get; set; } + internal void AnalyzeInvocation(SyntaxNodeAnalysisContext context) { var invocationSyntax = (InvocationExpressionSyntax)context.Node; @@ -140,6 +279,14 @@ internal void AnalyzeInvocation(SyntaxNodeAnalysisContext context) { if (this.MainThreadAssertingMethods.Contains(invokedMethod) || this.MainThreadSwitchingMethods.Contains(invokedMethod)) { + if (context.ContainingSymbol is IMethodSymbol callingMethod) + { + lock (this.MethodsDeclaringUIThreadRequirement) + { + this.MethodsDeclaringUIThreadRequirement.Add(callingMethod); + } + } + this.methodDeclarationNodes = this.methodDeclarationNodes.SetItem(methodDeclaration, ThreadingContext.MainThread); return; } @@ -219,7 +366,7 @@ private bool AnalyzeTypeWithinContext(ITypeSymbol type, ISymbol symbol, SyntaxNo Location location = focusDiagnosticOn ?? context.Node.GetLocation(); var exampleAssertingMethod = this.MainThreadAssertingMethods.FirstOrDefault(); var descriptor = exampleAssertingMethod.Name != null ? Descriptor : DescriptorNoAssertingMethod; - context.ReportDiagnostic(Diagnostic.Create(descriptor, location, type.Name, exampleAssertingMethod.ToString())); + context.ReportDiagnostic(Diagnostic.Create(descriptor, location, type.Name, exampleAssertingMethod)); return true; } }