diff --git a/src/TUnit.Mocks.SourceGenerator/Discovery/MockDiscoveryCache.cs b/src/TUnit.Mocks.SourceGenerator/Discovery/MockDiscoveryCache.cs new file mode 100644 index 00000000000..c2485e329e9 --- /dev/null +++ b/src/TUnit.Mocks.SourceGenerator/Discovery/MockDiscoveryCache.cs @@ -0,0 +1,213 @@ +using System.Collections.Concurrent; +using System.Collections.Immutable; +using System.Runtime.CompilerServices; +using Microsoft.CodeAnalysis; +using TUnit.Mocks.SourceGenerator.Models; + +namespace TUnit.Mocks.SourceGenerator.Discovery; + +/// +/// Per-compilation memo for the discovery transforms. +/// +/// Every Mock.Of<T>() / T.Mock() call site is transformed on its own, and each +/// transform used to walk the full member surface and the transitive interface closure of its +/// target. A type mocked at a hundred call sites was modelled a hundred times per compilation, and +/// the transforms run again for every site on every edit. The models only depend on the target +/// symbol, the mocking mode and the consuming compilation (member accessibility and +/// InternalsVisibleTo via , namespace conflicts via the +/// compilation's own declarations), so within one compilation the first call site's result is +/// reused by all others. +/// +/// +/// The cache is keyed on the instance and never outlives it, so there is +/// no cross-compilation state to invalidate: a new compilation (any edit) starts empty. +/// +/// +/// Values are computed before GetOrAdd runs, so two transforms that miss on the same key +/// at the same time can both build a model and one of them is discarded. The computation is pure, +/// so the outcome is the same either way; downstream consumers rely on value equality, not on +/// reference identity. A cancelled computation throws before storing, so it is never memoized. +/// +/// +internal sealed class MockDiscoveryCache +{ + private static readonly ConditionalWeakTable Caches = new(); + private static readonly ConditionalWeakTable.CreateValueCallback Create = static _ => new MockDiscoveryCache(); + + private int _mayReferenceGeneratedStaticExtensions = -1; + + private MockDiscoveryCache() + { + } + + public static MockDiscoveryCache For(Compilation compilation) => Caches.GetValue(compilation, Create); + + /// Results of BuildSingleTypeModel; means "not mockable". + public ConcurrentDictionary SingleTypeModels { get; } = new(); + + /// A single-type model followed by its transitive auto-mock interface models. + public ConcurrentDictionary> ModelsWithTransitiveDependencies { get; } = new(); + + /// All models produced by one Mock.Of<T1, T2, ...>() combination. + public ConcurrentDictionary> MultiTypeModels { get; } = new(); + + /// + /// Whether the compilation can see any type named *_MockStaticExtension, in any + /// namespace. A generator never sees its own output, so such a type can only be declared in + /// source or come from a referenced assembly; when there is none, a T.Mock() call + /// cannot already bind to one and the per-site binding check can be skipped. + /// + public bool MayReferenceGeneratedStaticExtensions(Compilation compilation) + { + var value = Volatile.Read(ref _mayReferenceGeneratedStaticExtensions); + if (value < 0) + { + value = ScanForStaticExtensions(compilation) ? 1 : 0; + Volatile.Write(ref _mayReferenceGeneratedStaticExtensions, value); + } + + return value == 1; + } + + private static bool ScanForStaticExtensions(Compilation compilation) + { + // Source declarations: answered from the declaration table, no symbols are created. + if (compilation.ContainsSymbolsWithName(IsStaticExtensionName, SymbolFilter.Type)) + { + return true; + } + + // Referenced assemblies: an extension producing TUnit.Mocks mocks has to reference + // TUnit.Mocks, so only those assemblies (a handful) are walked, across every namespace. + foreach (var assembly in compilation.SourceModule.ReferencedAssemblySymbols) + { + if (ReferencesTUnitMocks(assembly) && ContainsStaticExtension(assembly.GlobalNamespace)) + { + return true; + } + } + + return false; + } + + private static bool IsStaticExtensionName(string name) + => name.EndsWith("_MockStaticExtension", System.StringComparison.Ordinal); + + private static bool ReferencesTUnitMocks(IAssemblySymbol assembly) + { + if (assembly.Name == "TUnit.Mocks") + { + return true; + } + + foreach (var module in assembly.Modules) + { + foreach (var reference in module.ReferencedAssemblies) + { + if (reference.Name == "TUnit.Mocks") + { + return true; + } + } + } + + return false; + } + + private static bool ContainsStaticExtension(INamespaceSymbol ns) + { + // Extension classes are top-level static classes, so nested types need no visit. + foreach (var type in ns.GetTypeMembers()) + { + if (IsStaticExtensionName(type.Name)) + { + return true; + } + } + + foreach (var child in ns.GetNamespaceMembers()) + { + if (ContainsStaticExtension(child)) + { + return true; + } + } + + return false; + } +} + +/// Memo key for a single mocked type in one mocking mode. +internal readonly struct SingleTypeKey : System.IEquatable +{ + public SingleTypeKey(INamedTypeSymbol type, bool isPartialMock, bool isWrapMock) + { + Type = type; + IsPartialMock = isPartialMock; + IsWrapMock = isWrapMock; + } + + public INamedTypeSymbol Type { get; } + public bool IsPartialMock { get; } + public bool IsWrapMock { get; } + + public bool Equals(SingleTypeKey other) + => IsPartialMock == other.IsPartialMock + && IsWrapMock == other.IsWrapMock + && SymbolEqualityComparer.IncludeNullability.Equals(Type, other.Type); + + public override bool Equals(object? obj) => obj is SingleTypeKey other && Equals(other); + + public override int GetHashCode() + { + unchecked + { + var hash = SymbolEqualityComparer.IncludeNullability.GetHashCode(Type); + hash = hash * 31 + (IsPartialMock ? 1 : 0); + hash = hash * 31 + (IsWrapMock ? 1 : 0); + return hash; + } + } +} + +/// Memo key for an ordered list of type arguments. +internal readonly struct TypeListKey : System.IEquatable +{ + public TypeListKey(ImmutableArray types) => Types = types; + + public ImmutableArray Types { get; } + + public bool Equals(TypeListKey other) + { + if (Types.Length != other.Types.Length) + { + return false; + } + + for (var i = 0; i < Types.Length; i++) + { + if (!SymbolEqualityComparer.IncludeNullability.Equals(Types[i], other.Types[i])) + { + return false; + } + } + + return true; + } + + public override bool Equals(object? obj) => obj is TypeListKey other && Equals(other); + + public override int GetHashCode() + { + unchecked + { + var hash = 17; + foreach (var type in Types) + { + hash = hash * 31 + SymbolEqualityComparer.IncludeNullability.GetHashCode(type); + } + + return hash; + } + } +} diff --git a/src/TUnit.Mocks.SourceGenerator/Discovery/MockTypeDiscovery.cs b/src/TUnit.Mocks.SourceGenerator/Discovery/MockTypeDiscovery.cs index 574224a09fa..e979ef2f26f 100644 --- a/src/TUnit.Mocks.SourceGenerator/Discovery/MockTypeDiscovery.cs +++ b/src/TUnit.Mocks.SourceGenerator/Discovery/MockTypeDiscovery.cs @@ -66,7 +66,7 @@ public static ImmutableArray TransformToModels(GeneratorSyntaxCon // Verify this is TUnit.Mocks.Mock.Of() or TUnit.Mocks.MockRepository.Of() var containingTypeName = method.ContainingType?.Name; if ((containingTypeName != "Mock" && containingTypeName != "MockRepository") || - method.ContainingNamespace?.ToDisplayString() != "TUnit.Mocks") + !IsTUnitMocksNamespace(method.ContainingNamespace)) return ImmutableArray.Empty; var isDelegateMock = method.Name == "OfDelegate"; @@ -139,11 +139,29 @@ public static ImmutableArray TransformToModels(GeneratorSyntaxCon ct); } + // Every Mock.Of() site naming the same combination produces the same models. + var cache = MockDiscoveryCache.For(compilation).MultiTypeModels; + var key = new TypeListKey(method.TypeArguments); + if (cache.TryGetValue(key, out var cached)) + return cached; + + return cache.GetOrAdd(key, BuildMultiTypeModels( + namedType, method.TypeArguments, isPartialMock, compilationAssembly, compilation, ct)); + } + + private static ImmutableArray BuildMultiTypeModels( + INamedTypeSymbol namedType, + ImmutableArray typeArguments, + bool isPartialMock, + IAssemblySymbol? compilationAssembly, + Compilation compilation, + CancellationToken ct) + { // Multi-type mock: validate additional type args are all interfaces var additionalTypes = new List(); - for (int i = 1; i < method.TypeArguments.Length; i++) + for (int i = 1; i < typeArguments.Length; i++) { - if (method.TypeArguments[i] is not INamedTypeSymbol additionalType) + if (typeArguments[i] is not INamedTypeSymbol additionalType) return ImmutableArray.Empty; if (additionalType.TypeKind != TypeKind.Interface) return ImmutableArray.Empty; @@ -325,19 +343,21 @@ private static void CollectTransitiveInterfaceTypes( // Skip BCL/system interfaces — they have members (indexers, explicit implementations) // that the mock generator cannot handle, and auto-mocking them is rarely useful. - var ns = namedReturn.ContainingNamespace?.ToDisplayString() ?? ""; - if (IsFrameworkNamespace(ns)) + if (IsInFrameworkNamespace(namedReturn)) continue; + // Add returns false if already discovered/visited — skip without re-walking. Checked + // before the static-abstract scan below: the same return type typically recurs across + // many members, and a type that fails that scan fails it every time, so recording it + // as visited first never changes which types are generated. + if (!visited.Add(namedReturn.GetFullyQualifiedName())) continue; + // Skip interfaces that have static abstract members — using them as type arguments // in Mock/MockEngine triggers CS8920 because the static abstract members // don't have a most specific implementation in the interface. if (HasStaticAbstractMembers(namedReturn)) continue; - // Add returns false if already discovered/visited — skip without re-walking. - if (!visited.Add(namedReturn.GetFullyQualifiedName())) continue; - var model = BuildSingleTypeModel( namedReturn, isPartialMock: false, @@ -360,22 +380,53 @@ private static void CollectTransitiveInterfaceTypes( /// private static ITypeSymbol UnwrapAsyncType(ITypeSymbol type) { - if (type is INamedTypeSymbol { IsGenericType: true } named) + // Matches the definitions System.Threading.Tasks.Task and ValueTask by + // name, without formatting the symbol to a display string for every member. + if (type is INamedTypeSymbol { IsGenericType: true, Arity: 1, ContainingType: null } named + && named.Name is ("Task" or "ValueTask") + && named.ConstructedFrom.TypeParameters[0].Name == "TResult" + && IsNamespace(named.ContainingNamespace, "Tasks", "Threading", "System")) { - var constructedName = named.ConstructedFrom.ToDisplayString(); - if (constructedName is "System.Threading.Tasks.Task" - or "System.Threading.Tasks.ValueTask") - { - return named.TypeArguments[0]; - } + return named.TypeArguments[0]; } return type; } - private static bool IsFrameworkNamespace(string ns) => - ns == "System" || ns.StartsWith("System.") || - ns == "Microsoft" || ns.StartsWith("Microsoft.") || - ns == "Windows" || ns.StartsWith("Windows."); + /// + /// True for types under the System, Microsoft or Windows root namespaces. + /// + private static bool IsInFrameworkNamespace(INamedTypeSymbol type) + { + var ns = type.ContainingNamespace; + if (ns is null || ns.IsGlobalNamespace) + return false; + + while (ns.ContainingNamespace is { IsGlobalNamespace: false } parent) + { + ns = parent; + } + + return ns.Name is "System" or "Microsoft" or "Windows"; + } + + private static bool IsTUnitMocksNamespace(INamespaceSymbol? ns) + => IsNamespace(ns, "Mocks", "TUnit"); + + /// + /// True when is exactly the namespace whose segments, innermost first, + /// are . + /// + private static bool IsNamespace(INamespaceSymbol? ns, params string[] segmentsInnermostFirst) + { + foreach (var segment in segmentsInnermostFirst) + { + if (ns is null || ns.IsGlobalNamespace || ns.Name != segment) + return false; + ns = ns.ContainingNamespace; + } + + return ns is { IsGlobalNamespace: true }; + } /// /// Returns true if the interface (or any of its base interfaces) has static abstract members @@ -439,6 +490,22 @@ private static ImmutableArray BuildModelWithTransitiveDependencie IAssemblySymbol? compilationAssembly, Compilation compilation, CancellationToken cancellationToken) + { + var cache = MockDiscoveryCache.For(compilation).ModelsWithTransitiveDependencies; + var key = new SingleTypeKey(namedType, isPartialMock, isWrapMock: false); + if (cache.TryGetValue(key, out var cached)) + return cached; + + return cache.GetOrAdd(key, CreateModelWithTransitiveDependencies( + namedType, isPartialMock, compilationAssembly, compilation, cancellationToken)); + } + + private static ImmutableArray CreateModelWithTransitiveDependencies( + INamedTypeSymbol namedType, + bool isPartialMock, + IAssemblySymbol? compilationAssembly, + Compilation compilation, + CancellationToken cancellationToken) { var model = BuildSingleTypeModel( namedType, @@ -470,6 +537,26 @@ private static ImmutableArray BuildModelWithTransitiveDependencie Compilation compilation, bool isWrapMock, CancellationToken cancellationToken) + { + // The model depends only on the symbol, the mode flags and the compilation + // (compilationAssembly is always compilation.Assembly), so it is shared by every call + // site and every transitive walk that reaches the same type. + var cache = MockDiscoveryCache.For(compilation).SingleTypeModels; + var key = new SingleTypeKey(namedType, isPartialMock, isWrapMock); + if (cache.TryGetValue(key, out var cached)) + return cached; + + return cache.GetOrAdd(key, CreateSingleTypeModel( + namedType, isPartialMock, compilationAssembly, compilation, isWrapMock, cancellationToken)); + } + + private static MockTypeModel? CreateSingleTypeModel( + INamedTypeSymbol namedType, + bool isPartialMock, + IAssemblySymbol? compilationAssembly, + Compilation compilation, + bool isWrapMock, + CancellationToken cancellationToken) { // An interface with abstract members this compilation can't access (e.g. `internal` // members declared in another assembly) cannot be implemented by any type we could emit, @@ -625,18 +712,25 @@ public static ImmutableArray TransformMockExtensionInvocation( if (namedType.TypeKind is not (TypeKind.Interface or TypeKind.Class)) return ImmutableArray.Empty; - // Skip if .Mock() already resolves to a generated specialization (2nd incremental pass). - // The generated per-type extension lives in a class named *_MockStaticExtension. + var compilation = context.SemanticModel.Compilation; + + // Skip if .Mock() already resolves to a generated specialization. The generated per-type + // extension lives in a class named *_MockStaticExtension in the TUnit.Mocks namespace. // This covers both interfaces (wrapper return type in TUnit.Mocks.Generated) and - // classes (Mock return type in TUnit.Mocks). - var invocationSymbol = context.SemanticModel.GetSymbolInfo(invocation, ct); - if (invocationSymbol.Symbol is IMethodSymbol resolved - && resolved.ContainingType?.Name is { } containingName - && containingName.EndsWith("_MockStaticExtension")) - return ImmutableArray.Empty; + // classes (Mock return type in TUnit.Mocks). A generator never sees its own output, + // so such an extension can only come from a referenced assembly or hand-written source + // (in any namespace); binding the whole invocation is only worth it when the compilation + // can see one at all. + if (MockDiscoveryCache.For(compilation).MayReferenceGeneratedStaticExtensions(compilation)) + { + var invocationSymbol = context.SemanticModel.GetSymbolInfo(invocation, ct); + if (invocationSymbol.Symbol is IMethodSymbol resolved + && resolved.ContainingType?.Name is { } containingName + && containingName.EndsWith("_MockStaticExtension")) + return ImmutableArray.Empty; + } var isPartialMock = namedType.TypeKind == TypeKind.Class; - var compilation = context.SemanticModel.Compilation; var compilationAssembly = compilation.Assembly; return BuildModelWithTransitiveDependencies( NormalizeSingleMockType(namedType), @@ -652,7 +746,7 @@ public static ImmutableArray TransformMockExtensionInvocation( /// Semantic transform for [assembly: GenerateMock(typeof(T))]. /// Extracts the type argument and pairs each model with its attribute location. /// - public static ImmutableArray TransformGenerateMockAttribute( + public static EquatableArray TransformGenerateMockAttribute( GeneratorAttributeSyntaxContext context, CancellationToken ct) { // The target symbol for an assembly attribute is the assembly itself @@ -664,7 +758,7 @@ public static ImmutableArray TransformGenerateMockAttribu { if (attr.AttributeClass?.Name is not ("GenerateMockAttribute" or "GenerateMock")) continue; - if (attr.AttributeClass?.ContainingNamespace?.ToDisplayString() != "TUnit.Mocks") + if (!IsTUnitMocksNamespace(attr.AttributeClass?.ContainingNamespace)) continue; if (attr.ConstructorArguments.Length != 1) @@ -696,6 +790,6 @@ public static ImmutableArray TransformGenerateMockAttribu } } - return requests.ToImmutable(); + return new EquatableArray(requests.ToImmutable()); } } diff --git a/src/TUnit.Mocks.SourceGenerator/MockGenerator.cs b/src/TUnit.Mocks.SourceGenerator/MockGenerator.cs index 50f9116983b..088c96ed3a9 100644 --- a/src/TUnit.Mocks.SourceGenerator/MockGenerator.cs +++ b/src/TUnit.Mocks.SourceGenerator/MockGenerator.cs @@ -9,13 +9,13 @@ namespace TUnit.Mocks.SourceGenerator; [Generator(LanguageNames.CSharp)] public class MockGenerator : IIncrementalGenerator { - private readonly Action _emitSources; + private readonly Action _emitSources; public MockGenerator() : this(EmitSources) { } - internal MockGenerator(Action emitSources) + internal MockGenerator(Action emitSources) { _emitSources = emitSources; } @@ -41,7 +41,8 @@ namespace TUnit.Mocks.Generated; transform: static (ctx, ct) => CreateRequests( MockTypeDiscovery.TransformToModels(ctx, ct), ctx.Node.GetLocation())) - .SelectMany((requests, _) => requests); + .WithTrackingName(MockTrackingNames.MockOfInvocations) + .SelectMany((requests, _) => requests.AsImmutableArray()); // Step 1b: Find all [assembly: GenerateMock(typeof(T))] attributes var attributeTypes = context.SyntaxProvider @@ -49,7 +50,8 @@ namespace TUnit.Mocks.Generated; "TUnit.Mocks.GenerateMockAttribute", predicate: static (node, _) => true, transform: MockTypeDiscovery.TransformGenerateMockAttribute) - .SelectMany((requests, _) => requests); + .WithTrackingName(MockTrackingNames.GenerateMockAttributes) + .SelectMany((requests, _) => requests.AsImmutableArray()); // Attribute-only requests have no invocation for the TM006 analyzer to inspect. context.RegisterSourceOutput(attributeTypes, static (spc, request) => @@ -72,10 +74,11 @@ namespace TUnit.Mocks.Generated; transform: static (ctx, ct) => CreateRequests( MockTypeDiscovery.TransformMockExtensionInvocation(ctx, ct), ctx.Node.GetLocation())) - .SelectMany((requests, _) => requests); + .WithTrackingName(MockTrackingNames.MockExtensionInvocations) + .SelectMany((requests, _) => requests.AsImmutableArray()); // Step 2: Merge all sources and deduplicate - var distinctTypes = mockTypes + var distinctRequests = mockTypes .Collect() .Combine(attributeTypes.Collect()) .Combine(extensionTypes.Collect()) @@ -96,36 +99,98 @@ namespace TUnit.Mocks.Generated; // collision — it just has to agree on who emits the shared member surface (#6834). return SharedMemberSurfaceResolver.Resolve( GeneratedNameCollisionDetector.Annotate(requests)); - }); + }) + .WithTrackingName(MockTrackingNames.DistinctRequests); + + // Step 3: Generate source for each unique type. The request's location is dropped first, + // so a call site that merely moves (an edit above it) leaves the model — and with it the + // generated source — cached. + var emitResults = distinctRequests + .Select(static (request, _) => request.Model) + .WithTrackingName(MockTrackingNames.DistinctModels) + .Select((model, _) => Emit(model)) + .WithTrackingName(MockTrackingNames.EmitResults); + + context.RegisterSourceOutput(emitResults, AddEmittedSources); + + // Step 4: TM009 needs the location of the request that produced a failed model. Pairing + // happens here rather than in the emit step so locations never invalidate emitted source. + var failures = emitResults + .Where(static result => result.Failed) + .Collect() + .Combine(distinctRequests.Collect()); - // Step 3: Generate source for each unique type - context.RegisterSourceOutput(distinctTypes, GenerateMockSafely); + context.RegisterSourceOutput(failures, ReportGenerationFailures); } - private void GenerateMockSafely(SourceProductionContext spc, MockGenerationRequest request) + private MockEmitResult Emit(MockTypeModel model) { + var sink = new MockSourceSink(); + try { - _emitSources(spc, request.Model); + _emitSources(sink, model); + return new MockEmitResult(model, sink.ToEquatableArray(), null, null); } catch (Exception exception) when (exception is not OperationCanceledException) { + // Sources added before the failure are still emitted, as they were when generation + // wrote straight to the SourceProductionContext. + return new MockEmitResult(model, sink.ToEquatableArray(), exception.GetType().Name, exception.Message); + } + } + + private static void AddEmittedSources(SourceProductionContext spc, MockEmitResult result) + { + var model = result.Model; + if (model.CollidesWith is not null) + { + spc.ReportDiagnostic(Diagnostic.Create( + Diagnostics.TM008_GeneratedNameCollision, + Location.None, + model.FullyQualifiedName, + MockImplBuilder.GetCompositeSafeName(model), + model.CollidesWith)); + } + + foreach (var source in result.Sources) + { + spc.AddSource(source.HintName, source.Source); + } + } + + private static void ReportGenerationFailures( + SourceProductionContext spc, + (ImmutableArray Failures, ImmutableArray Requests) input) + { + foreach (var failure in input.Failures) + { + var location = Location.None; + foreach (var request in input.Requests) + { + if (request.Model.Equals(failure.Model)) + { + location = request.SourceLocation.ToLocation(); + break; + } + } + spc.ReportDiagnostic(Diagnostic.Create( Diagnostics.TM009_GenerationFailed, - request.SourceLocation.ToLocation(), - request.Model.FullyQualifiedName, - exception.GetType().Name, - exception.Message)); + location, + failure.Model.FullyQualifiedName, + failure.FailureExceptionType, + failure.FailureMessage)); } } - private static ImmutableArray CreateRequests( + private static EquatableArray CreateRequests( ImmutableArray models, Location location) { if (models.IsDefaultOrEmpty) { - return ImmutableArray.Empty; + return EquatableArray.Empty; } var sourceLocation = MockSourceLocation.From(location); @@ -135,7 +200,7 @@ private static ImmutableArray CreateRequests( requests.Add(new MockGenerationRequest(model, sourceLocation)); } - return requests.MoveToImmutable(); + return new EquatableArray(requests.MoveToImmutable()); } private static void AddDistinctRequests( @@ -154,16 +219,11 @@ private static void AddDistinctRequests( } } - internal static void EmitSources(SourceProductionContext spc, MockTypeModel model) + internal static void EmitSources(MockSourceSink spc, MockTypeModel model) { if (model.CollidesWith is not null) { - spc.ReportDiagnostic(Diagnostic.Create( - Diagnostics.TM008_GeneratedNameCollision, - Location.None, - model.FullyQualifiedName, - MockImplBuilder.GetCompositeSafeName(model), - model.CollidesWith)); + // TM008 is reported when the (empty) result is added to the compilation. return; } @@ -205,7 +265,7 @@ internal static void EmitSources(SourceProductionContext spc, MockTypeModel mode } } - private static void GenerateSingleTypeMock(SourceProductionContext spc, MockTypeModel model) + private static void GenerateSingleTypeMock(MockSourceSink spc, MockTypeModel model) { var fileName = GetSafeFileName(model); @@ -244,7 +304,7 @@ private static void GenerateSingleTypeMock(SourceProductionContext spc, MockType } } - private static void GenerateUnconstructableClassStub(SourceProductionContext spc, MockTypeModel model) + private static void GenerateUnconstructableClassStub(MockSourceSink spc, MockTypeModel model) { var extensionSource = MockStaticExtensionBuilder.BuildForPartialMock(model); if (!string.IsNullOrEmpty(extensionSource)) @@ -253,7 +313,7 @@ private static void GenerateUnconstructableClassStub(SourceProductionContext spc } } - private static void GenerateDelegateMock(SourceProductionContext spc, MockTypeModel model) + private static void GenerateDelegateMock(MockSourceSink spc, MockTypeModel model) { var fileName = GetSafeFileName(model); @@ -264,7 +324,7 @@ private static void GenerateDelegateMock(SourceProductionContext spc, MockTypeMo spc.AddSource($"{fileName}_MockDelegateFactory.g.cs", factorySource); } - private static void GenerateWrapMock(SourceProductionContext spc, MockTypeModel model) + private static void GenerateWrapMock(MockSourceSink spc, MockTypeModel model) { var fileName = GetSafeFileName(model); @@ -281,14 +341,14 @@ private static void GenerateWrapMock(SourceProductionContext spc, MockTypeModel } } - private static void GenerateMultiInterfaceMock(SourceProductionContext spc, MockTypeModel model) + private static void GenerateMultiInterfaceMock(MockSourceSink spc, MockTypeModel model) { var fileName = GetSafeFileName(model); var implFactorySource = BuildCombinedImplAndFactory(model); spc.AddSource($"{fileName}_MockImplFactory.g.cs", implFactorySource); } - private static void GenerateImplFactoryMembersAndEvents(SourceProductionContext spc, MockTypeModel model, string fileName) + private static void GenerateImplFactoryMembersAndEvents(MockSourceSink spc, MockTypeModel model, string fileName) { var implFactorySource = BuildCombinedImplAndFactory(model); spc.AddSource($"{fileName}_MockImplFactory.g.cs", implFactorySource); @@ -296,7 +356,7 @@ private static void GenerateImplFactoryMembersAndEvents(SourceProductionContext GenerateMembersAndEvents(spc, model, fileName); } - private static void GenerateMembersAndEvents(SourceProductionContext spc, MockTypeModel model, string fileName) + private static void GenerateMembersAndEvents(MockSourceSink spc, MockTypeModel model, string fileName) { var membersSource = MockMembersBuilder.Build(model); spc.AddSource($"{fileName}_MockMembers.g.cs", membersSource); diff --git a/src/TUnit.Mocks.SourceGenerator/MockTrackingNames.cs b/src/TUnit.Mocks.SourceGenerator/MockTrackingNames.cs new file mode 100644 index 00000000000..61e91bcb6dc --- /dev/null +++ b/src/TUnit.Mocks.SourceGenerator/MockTrackingNames.cs @@ -0,0 +1,14 @@ +namespace TUnit.Mocks.SourceGenerator; + +/// +/// Names of the pipeline steps, for incrementality tests. +/// +internal static class MockTrackingNames +{ + public const string MockOfInvocations = nameof(MockOfInvocations); + public const string GenerateMockAttributes = nameof(GenerateMockAttributes); + public const string MockExtensionInvocations = nameof(MockExtensionInvocations); + public const string DistinctRequests = nameof(DistinctRequests); + public const string DistinctModels = nameof(DistinctModels); + public const string EmitResults = nameof(EmitResults); +} diff --git a/src/TUnit.Mocks.SourceGenerator/Models/EquatableArray.cs b/src/TUnit.Mocks.SourceGenerator/Models/EquatableArray.cs index 2578488e11a..f2e03a8468f 100644 --- a/src/TUnit.Mocks.SourceGenerator/Models/EquatableArray.cs +++ b/src/TUnit.Mocks.SourceGenerator/Models/EquatableArray.cs @@ -27,6 +27,8 @@ public bool Equals(EquatableArray other) if (_array.IsDefault && other._array.IsDefault) return true; if (_array.IsDefault || other._array.IsDefault) return false; if (_array.Length != other._array.Length) return false; + // Shared backing array (e.g. a memoized model reused by several call sites). + if (_array == other._array) return true; for (int i = 0; i < _array.Length; i++) { diff --git a/src/TUnit.Mocks.SourceGenerator/Models/MockEmitResult.cs b/src/TUnit.Mocks.SourceGenerator/Models/MockEmitResult.cs new file mode 100644 index 00000000000..befa325c0e3 --- /dev/null +++ b/src/TUnit.Mocks.SourceGenerator/Models/MockEmitResult.cs @@ -0,0 +1,35 @@ +using System; +using System.Collections.Generic; +using System.Collections.Immutable; + +namespace TUnit.Mocks.SourceGenerator.Models; + +/// A generated file: its hint name and text. +internal readonly record struct GeneratedMockSource(string HintName, string Source); + +/// +/// The sources generated for one model, plus the failure that interrupted generation, if any. +/// Carries no source location, so moving a call site never invalidates generated output; TM009 +/// pairs a failure with its request location in a separate step. +/// +internal sealed record MockEmitResult( + MockTypeModel Model, + EquatableArray Sources, + string? FailureExceptionType, + string? FailureMessage) +{ + public bool Failed => FailureExceptionType is not null; +} + +/// Collects generated files for a model; the incremental pipeline adds them later. +internal sealed class MockSourceSink +{ + private readonly List _sources = new(); + + public void AddSource(string hintName, string source) => _sources.Add(new GeneratedMockSource(hintName, source)); + + public EquatableArray ToEquatableArray() + => _sources.Count == 0 + ? EquatableArray.Empty + : new EquatableArray(_sources.ToImmutableArray()); +} diff --git a/src/TUnit.Mocks.SourceGenerator/Models/MockTypeModel.cs b/src/TUnit.Mocks.SourceGenerator/Models/MockTypeModel.cs index 9ff8d5322f9..db49ea16be2 100644 --- a/src/TUnit.Mocks.SourceGenerator/Models/MockTypeModel.cs +++ b/src/TUnit.Mocks.SourceGenerator/Models/MockTypeModel.cs @@ -98,6 +98,9 @@ public bool LacksAccessibleConstructor public bool Equals(MockTypeModel? other) { if (other is null) return false; + // Discovery memoizes models per compilation, so duplicates reaching dedup are usually the + // very same instance. + if (ReferenceEquals(this, other)) return true; return FullyQualifiedName == other.FullyQualifiedName && OpenGenericTypeOfExpression == other.OpenGenericTypeOfExpression && Name == other.Name @@ -123,6 +126,13 @@ public bool Equals(MockTypeModel? other) && SecondaryMemberIdMaps.Equals(other.SecondaryMemberIdMaps); } + /// + /// Shallow hash over the identity and shape of the model. The deep member arrays are left to + /// : hashing them walked every member, parameter and + /// nested array of every model on each dedup pass, while the identity fields alone already + /// separate distinct models. Everything hashed here is also compared by Equals, so equal + /// models still hash equally. + /// public override int GetHashCode() { unchecked @@ -135,15 +145,15 @@ public override int GetHashCode() hash = hash * 31 + IsWrapMock.GetHashCode(); hash = hash * 31 + IsPublic.GetHashCode(); hash = hash * 31 + UseFallbackNamespace.GetHashCode(); - hash = hash * 31 + TypeParameters.GetHashCode(); - hash = hash * 31 + Methods.GetHashCode(); - hash = hash * 31 + Properties.GetHashCode(); - hash = hash * 31 + Events.GetHashCode(); + hash = hash * 31 + TypeParameters.Length; + hash = hash * 31 + Methods.Length; + hash = hash * 31 + Properties.Length; + hash = hash * 31 + Events.Length; hash = hash * 31 + AdditionalInterfaceNames.GetHashCode(); hash = hash * 31 + HasStaticAbstractMembers.GetHashCode(); hash = hash * 31 + IsSecondaryMemberSurface.GetHashCode(); hash = hash * 31 + (CollidesWith?.GetHashCode() ?? 0); - hash = hash * 31 + SecondaryMemberIdMaps.GetHashCode(); + hash = hash * 31 + SecondaryMemberIdMaps.Length; return hash; } } diff --git a/tests/TUnit.Mocks.SourceGenerator.Tests/MockDiscoveryCacheTests.cs b/tests/TUnit.Mocks.SourceGenerator.Tests/MockDiscoveryCacheTests.cs new file mode 100644 index 00000000000..6234485a019 --- /dev/null +++ b/tests/TUnit.Mocks.SourceGenerator.Tests/MockDiscoveryCacheTests.cs @@ -0,0 +1,94 @@ +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using TUnit.Mocks.SourceGenerator.Discovery; + +namespace TUnit.Mocks.SourceGenerator.Tests; + +/// +/// Locks in the gate that decides whether T.Mock() call sites need a binding check for an +/// existing *_MockStaticExtension. The test-pinned Roslyn cannot parse C# 14 +/// extension(...) blocks, so the gate is exercised directly rather than end to end. +/// +public class MockDiscoveryCacheTests : SnapshotTestBase +{ + private const string Consumer = """ + public class TestUsage + { + } + """; + + [Test] + public async Task No_Static_Extension_Anywhere_Skips_Binding_Check() + { + var compilation = CreateCompilation(Consumer); + + await Assert.That(MockDiscoveryCache.For(compilation).MayReferenceGeneratedStaticExtensions(compilation)).IsFalse(); + } + + [Test] + public async Task Source_Declared_Static_Extension_In_Other_Namespace_Is_Detected() + { + var compilation = CreateCompilation(""" + namespace Custom.Mocking + { + public static class INotifier_MockStaticExtension + { + } + } + """); + + await Assert.That(MockDiscoveryCache.For(compilation).MayReferenceGeneratedStaticExtensions(compilation)).IsTrue(); + } + + [Test] + public async Task Referenced_Static_Extension_In_Other_Namespace_Is_Detected() + { + var reference = CreateExternalAssemblyReference(""" + namespace ExternalLib.Mocking + { + public interface IExternalNotifier + { + } + + public static class IExternalNotifier_MockStaticExtension + { + // Uses a TUnit.Mocks type so the assembly references TUnit.Mocks. + public static global::TUnit.Mocks.Mock? Mock() => null; + } + } + """); + var compilation = CreateCompilation(Consumer, reference); + + await Assert.That(MockDiscoveryCache.For(compilation).MayReferenceGeneratedStaticExtensions(compilation)).IsTrue(); + } + + [Test] + public async Task Referenced_Assembly_Without_TUnit_Mocks_Reference_Is_Not_Walked() + { + var reference = CreateExternalAssemblyReference(""" + namespace ExternalLib + { + public static class Unrelated_MockStaticExtension + { + } + } + """); + var compilation = CreateCompilation(Consumer, reference); + + await Assert.That(MockDiscoveryCache.For(compilation).MayReferenceGeneratedStaticExtensions(compilation)).IsFalse(); + } + + private static Compilation CreateCompilation(string source, MetadataReference? reference = null) + { + var parseOptions = CSharpParseOptions.Default.WithLanguageVersion(LanguageVersion.Preview); + IEnumerable references = reference is null + ? GetCachedReferences() + : GetCachedReferences().Append(reference); + + return CSharpCompilation.Create( + "TestAssembly", + [CSharpSyntaxTree.ParseText(source, parseOptions)], + references, + new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary)); + } +} diff --git a/tests/TUnit.Mocks.SourceGenerator.Tests/MockGeneratorDiagnosticTests.cs b/tests/TUnit.Mocks.SourceGenerator.Tests/MockGeneratorDiagnosticTests.cs index 7c553e4b50a..0fc84404d97 100644 --- a/tests/TUnit.Mocks.SourceGenerator.Tests/MockGeneratorDiagnosticTests.cs +++ b/tests/TUnit.Mocks.SourceGenerator.Tests/MockGeneratorDiagnosticTests.cs @@ -86,7 +86,7 @@ public interface IBroken { void Break(); } } private static void EmitWithInjectedFailure( - SourceProductionContext context, + MockSourceSink context, MockTypeModel model) { if (model.FullyQualifiedName == "global::IBroken") diff --git a/tests/TUnit.Mocks.SourceGenerator.Tests/MockGeneratorIncrementalityTests.cs b/tests/TUnit.Mocks.SourceGenerator.Tests/MockGeneratorIncrementalityTests.cs new file mode 100644 index 00000000000..1cb744eb727 --- /dev/null +++ b/tests/TUnit.Mocks.SourceGenerator.Tests/MockGeneratorIncrementalityTests.cs @@ -0,0 +1,164 @@ +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; + +namespace TUnit.Mocks.SourceGenerator.Tests; + +/// +/// Checks which pipeline steps re-run between generator passes. Emitting a mock is by far the most +/// expensive step, so edits that do not change a mocked type must leave its emit step cached. +/// +public class MockGeneratorIncrementalityTests : SnapshotTestBase +{ + private static readonly CSharpParseOptions ParseOptions = + CSharpParseOptions.Default.WithLanguageVersion(LanguageVersion.Preview); + + private const string Contracts = """ + namespace TestNamespace; + + public interface IGreeter + { + string Greet(string name); + IClock Clock { get; } + } + + public interface IClock + { + int Now(); + } + + public abstract class Repository + { + public abstract int Count(); + } + """; + + private const string Usage = """ + using TUnit.Mocks; + + namespace TestNamespace; + + public class Usage + { + void M() + { + _ = Mock.Of(); + _ = IGreeter.Mock(); + _ = Mock.Of(); + } + } + """; + + [Test] + public async Task Unrelated_Edit_Leaves_Emitted_Source_Cached() + { + var compilation = CreateCompilation(Contracts, Usage); + var driver = RunTracked(CreateDriver(), compilation); + + var edited = compilation.AddSyntaxTrees( + CSharpSyntaxTree.ParseText("namespace TestNamespace; public class Unrelated { }", ParseOptions)); + driver = RunTracked(driver, edited); + + var result = driver.GetRunResult().Results.Single(); + await AssertAllOutputs(result, MockTrackingNames.DistinctModels, IncrementalStepRunReason.Cached, IncrementalStepRunReason.Unchanged); + await AssertAllOutputs(result, MockTrackingNames.EmitResults, IncrementalStepRunReason.Cached); + } + + [Test] + public async Task Moving_A_Call_Site_Leaves_Emitted_Source_Cached() + { + var compilation = CreateCompilation(Contracts, Usage); + var driver = RunTracked(CreateDriver(), compilation); + var firstSources = GetGeneratedSources(driver); + + // A blank line above the call sites shifts every request location. + var usageTree = compilation.SyntaxTrees.Single(t => t.ToString().Contains("class Usage")); + var edited = compilation.ReplaceSyntaxTree( + usageTree, + CSharpSyntaxTree.ParseText(Usage.Replace("public class Usage", "\npublic class Usage"), ParseOptions)); + driver = RunTracked(driver, edited); + + var result = driver.GetRunResult().Results.Single(); + + // The requests did change (their locations moved), so the test is not vacuous... + await Assert.That(result.TrackedSteps[MockTrackingNames.DistinctRequests] + .SelectMany(step => step.Outputs) + .Any(output => output.Reason == IncrementalStepRunReason.Modified)) + .IsTrue(); + + // ...but the models, and therefore the generated source, did not. + await AssertAllOutputs(result, MockTrackingNames.DistinctModels, IncrementalStepRunReason.Cached, IncrementalStepRunReason.Unchanged); + await AssertAllOutputs(result, MockTrackingNames.EmitResults, IncrementalStepRunReason.Cached); + await Assert.That(GetGeneratedSources(driver)).IsEquivalentTo(firstSources); + } + + [Test] + public async Task Changing_A_Mocked_Interface_Regenerates_Its_Source() + { + var compilation = CreateCompilation(Contracts, Usage); + var driver = RunTracked(CreateDriver(), compilation); + + var contractsTree = compilation.SyntaxTrees.Single(t => t.ToString().Contains("interface IGreeter")); + var edited = compilation.ReplaceSyntaxTree( + contractsTree, + CSharpSyntaxTree.ParseText( + Contracts.Replace("string Greet(string name);", "string Greet(string name);\n void Wave();"), + ParseOptions)); + driver = RunTracked(driver, edited); + + var result = driver.GetRunResult().Results.Single(); + var emitOutputs = result.TrackedSteps[MockTrackingNames.EmitResults] + .SelectMany(step => step.Outputs) + .ToList(); + + await Assert.That(emitOutputs.Any(output => output.Reason == IncrementalStepRunReason.Modified)).IsTrue(); + await Assert.That(GetGeneratedSources(driver).Any(source => source.Contains("Wave", StringComparison.Ordinal))).IsTrue(); + } + + private static async Task AssertAllOutputs( + GeneratorRunResult result, + string stepName, + params IncrementalStepRunReason[] allowedReasons) + { + var reasons = result.TrackedSteps[stepName] + .SelectMany(step => step.Outputs) + .Select(output => output.Reason) + .ToList(); + + await Assert.That(reasons).IsNotEmpty(); + foreach (var reason in reasons) + { + await Assert.That(allowedReasons).Contains(reason); + } + } + + private static CSharpCompilation CreateCompilation(params string[] sources) + => CSharpCompilation.Create( + "TestAssembly", + sources.Select(source => CSharpSyntaxTree.ParseText(source, ParseOptions)), + GetCachedReferences(), + new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary)); + + private static GeneratorDriver CreateDriver() + => CSharpGeneratorDriver.Create( + [new MockGenerator().AsSourceGenerator()], + parseOptions: ParseOptions, + driverOptions: new GeneratorDriverOptions(IncrementalGeneratorOutputKind.None, trackIncrementalGeneratorSteps: true)); + + private static GeneratorDriver RunTracked(GeneratorDriver driver, Compilation compilation) + { + driver = driver.RunGenerators(compilation); + + var errors = driver.GetRunResult().Diagnostics.Where(d => d.Severity == DiagnosticSeverity.Error).ToList(); + if (errors.Count > 0) + { + throw new InvalidOperationException(string.Join(Environment.NewLine, errors)); + } + + return driver; + } + + private static string[] GetGeneratedSources(GeneratorDriver driver) + => driver.GetRunResult().GeneratedTrees + .Select(tree => tree.GetText().ToString()) + .ToArray(); +}