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();
+}