From abd649c4df53315cf5d214eae63bd806d62e5ae0 Mon Sep 17 00:00:00 2001 From: caran Date: Fri, 10 Apr 2026 17:01:29 +0200 Subject: [PATCH 01/21] Extensions and bug fixes --- Benchmarking/Benchmarks/Benchmarks.csproj | 4 +- Benchmarking/Benchmarks/Program.cs | 40 +- ...toryGenerator.Extensions.AspNetCore.csproj | 2 +- FactoryGenerator/FactoryGenerator.cs | 390 ++++++++++++++++-- ...nerator.Extensions.AspNetCore.Tests.csproj | 2 +- .../InjectionDetectionTests.cs | 45 ++ Tests/TestData/Inherited/Types.cs | 58 +++ Tests/TestData/Inheritor/Inheritor.csproj | 1 + Tests/TestData/Inheritor/Types.cs | 18 +- Tests/TestWebApp/TestWebApp.csproj | 2 +- 10 files changed, 519 insertions(+), 43 deletions(-) diff --git a/Benchmarking/Benchmarks/Benchmarks.csproj b/Benchmarking/Benchmarks/Benchmarks.csproj index 7da39a4..ef7a918 100644 --- a/Benchmarking/Benchmarks/Benchmarks.csproj +++ b/Benchmarking/Benchmarks/Benchmarks.csproj @@ -2,10 +2,10 @@ Exe - net9.0 + net10.0 + preview enable enable - true diff --git a/Benchmarking/Benchmarks/Program.cs b/Benchmarking/Benchmarks/Program.cs index ed35ed6..4101508 100644 --- a/Benchmarking/Benchmarks/Program.cs +++ b/Benchmarking/Benchmarks/Program.cs @@ -7,6 +7,8 @@ namespace Benchmarks; +// ── Dictionary-based resolution (existing path) ────────────────────────────── + [MemoryDiagnoser] [JsonExporterAttribute.Full] [JsonExporterAttribute.FullCompressed] @@ -36,10 +38,38 @@ public class ResolveBenchmarks public IContainer CreateFromSelf() => new DependencyInjectionContainer(m_container); } +// ── Static-extension resolution (C# 14 / .NET 10+ path) ───────────────────── +// +// Each Resolve(container?) call inlines the full construction chain directly — +// no dictionary lookup, no factory-method indirection. +// +// Null-container variants bypass the singleton cache entirely and perform a +// fresh allocation on every call, exposing the raw construction cost. + +[MemoryDiagnoser] +[JsonExporterAttribute.Full] +[JsonExporterAttribute.FullCompressed] +public class StaticExtensionBenchmarks +{ + private readonly DependencyInjectionContainer m_container = new(default, default, new NonInjectedClass()); + [Benchmark] + public ISingleton ExtensionResolveSingleton() => ISingleton.Resolve(m_container); + [Benchmark] + public ISingleton ExtensionResolveSingletonNullContainer() => ISingleton.Resolve(null); + [Benchmark] + public IOverridable ExtensionResolveTransient() => IOverridable.Resolve(m_container); + [Benchmark] + public ChainA ExtensionResolveChain() => ChainA.Resolve(m_container); + [Benchmark] + public ChainA ExtensionResolveChainNullContainer() => ChainA.Resolve(null); + [Benchmark] + public ArrayConsumer ExtensionResolveWithCollection() => ArrayConsumer.Resolve(m_container); + [Benchmark] + public ArrayConsumer ExtensionResolveWithCollectionNullContainer() => ArrayConsumer.Resolve(null); +} + internal static class Program { - private static void Main(string[] args) - { - var summary = BenchmarkRunner.Run(); - } -} \ No newline at end of file + private static void Main(string[] args) => + BenchmarkSwitcher.FromAssembly(typeof(Program).Assembly).Run(args); +} diff --git a/FactoryGenerator.Extensions.AspNetCore/FactoryGenerator.Extensions.AspNetCore.csproj b/FactoryGenerator.Extensions.AspNetCore/FactoryGenerator.Extensions.AspNetCore.csproj index 79f3779..95b0faa 100644 --- a/FactoryGenerator.Extensions.AspNetCore/FactoryGenerator.Extensions.AspNetCore.csproj +++ b/FactoryGenerator.Extensions.AspNetCore/FactoryGenerator.Extensions.AspNetCore.csproj @@ -1,7 +1,7 @@  - net8.0 + net10.0 FactoryGenerator.Extensions.AspNetCore latest diff --git a/FactoryGenerator/FactoryGenerator.cs b/FactoryGenerator/FactoryGenerator.cs index be14e7f..50834d0 100644 --- a/FactoryGenerator/FactoryGenerator.cs +++ b/FactoryGenerator/FactoryGenerator.cs @@ -5,6 +5,7 @@ using System.Text; using System.Threading; using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; using Microsoft.CodeAnalysis.CSharp.Syntax; using Microsoft.CodeAnalysis.Diagnostics; @@ -33,6 +34,10 @@ public void Initialize(IncrementalGeneratorInitializationContext context) .Collect(); var combined = attributes.Combine(compilation).Combine(syntaxUsages).Combine(logProvider); context.RegisterSourceOutput(combined, MakeAutofacModule); + + var supportsStaticExtensions = context.ParseOptionsProvider.Select(IsAtLeastCSharp14); + var extensionData = attributes.Combine(compilation).Combine(supportsStaticExtensions); + context.RegisterSourceOutput(extensionData, MakeStaticExtensions); } private IncrementalValueProvider SetupLog(IncrementalGeneratorInitializationContext context) @@ -189,12 +194,12 @@ private IContainer GetTop() }} public IContainer? Base {{ get; }} public IContainer? Inheritor {{ get; set; }} - private readonly object m_lock = new(); + internal readonly object m_lock = new(); private Dictionary> m_lookup; private Dictionary m_booleans; private List>? resolvedInstances; - private List> GetResolvedInstances() + internal List> GetResolvedInstances() {{ if (resolvedInstances is null) lock (m_lock) @@ -382,7 +387,7 @@ public bool GetBoolean(string key) { if (!parameter.IsCollection) continue; if (parameter.CollectionElementFullName is null) continue; - var name = parameter.Name; + var name = "coll_" + parameter.CollectionElementMemberName!; log.Log(LogLevel.Debug, $"Creating Collection: {name} of element type {parameter.CollectionElementFullName}"); MakeArray(arrayDeclarations, name, parameter.CollectionElementFullName, parameter.CollectionElementMemberName!, interfaceInjectors); constructorParameters.Remove(parameter); @@ -423,7 +428,7 @@ public bool GetBoolean(string key) var lifetimeParameters = string.Join(", ", allParameters); log.Log(LogLevel.Debug, $"Resulting Constructor: {constructor}"); - var constructorFields = string.Join("\n\t", allArguments.Select(arg => arg + ";")); + var constructorFields = string.Join("\n\t", allArguments.Select(arg => "internal " + arg + ";")); var constructorAssignments = string.Join("\n\t\t", allArguments.Select(arg => arg.Split(' ').Last()).Select(arg => $"this.{arg} = {arg};")); var resolvedConstructorAssignments = string.Join("\n\t\t", @@ -434,7 +439,7 @@ public bool GetBoolean(string key) // ReadOnlySpan is a ref struct and cannot be placed in the lookup dictionary var localizedForDict = localizedParameters.Where(p => p.CollectionKind != CollectionKind.ReadOnlySpan).ToList(); var localizedPairs = localizedForDict - .Select(p => (TypeName: p.TypeFullName, Expression: CollectionDictExpression(p.CollectionKind, p.Name))) + .Select(p => (TypeName: p.TypeFullName, Expression: CollectionDictExpression(p.CollectionKind, "coll_" + p.CollectionElementMemberName!))) .ToList(); var requestedPairs = requestedUsages.Select(u => (TypeName: u.FullName, MemberName: u.MemberName)).ToList(); var constructorPairs = constructorParameters.Select(p => (TypeName: p.TypeFullName, Expression: p.Name)).ToList(); @@ -481,13 +486,13 @@ public ILifetimeScope BeginLifetimeScope() GetResolvedInstances().Add(new WeakReference(scope)); return scope; }} - private readonly object m_lock = new(); + internal readonly object m_lock = new(); private {ClassName} m_fallback; private Dictionary> m_lookup; private Dictionary m_booleans; private List>? resolvedInstances; - private List> GetResolvedInstances() + internal List> GetResolvedInstances() {{ if (resolvedInstances is null) lock (m_lock) @@ -627,17 +632,16 @@ internal static void Register() private static void CheckForCycles(ImmutableArray dataInjections) { - var tree = new Dictionary>(); + // Build adjacency list: interface/type name → set of dependency names + var graph = new Dictionary>(); + // Map each interface name back to its concrete type for error messages + var nodeOwner = new Dictionary(); + foreach (var injection in dataInjections) { if (injection.Lambda != null) continue; - var node = new List(); - foreach (var ifaceName in injection.InterfaceFullNames) - { - if (!tree.ContainsKey(ifaceName)) - tree[ifaceName] = node; - } + var deps = new HashSet(); foreach (var ctor in injection.Constructors) { foreach (var parameter in ctor.Parameters) @@ -656,19 +660,70 @@ private static void CheckForCycles(ImmutableArray dataInjections) : parameter.TypeFullName; } - node.Add(depName); - if (tree.TryGetValue(depName, out var list)) - { - foreach (var ifaceName in injection.InterfaceFullNames) - { - if (list.Contains(ifaceName)) - throw new InvalidOperationException( - $"Cyclic Dependency Detected between {injection.TypeFullName} and {ifaceName}"); - } - } + deps.Add(depName); + } + } + + foreach (var ifaceName in injection.InterfaceFullNames) + { + if (!graph.ContainsKey(ifaceName)) + { + graph[ifaceName] = deps; + nodeOwner[ifaceName] = injection.TypeFullName; + } + else + { + // Multiple implementations of the same interface — merge edges + foreach (var d in deps) + graph[ifaceName].Add(d); } } } + + // DFS-based cycle detection + // 0 = unvisited, 1 = in-progress (on current path), 2 = done + var state = new Dictionary(); + var path = new List(); + + foreach (var node in graph.Keys) + { + if (!state.TryGetValue(node, out var s) || s == 0) + DfsCycleCheck(node, graph, state, path, nodeOwner); + } + } + + private static void DfsCycleCheck(string node, Dictionary> graph, + Dictionary state, List path, + Dictionary nodeOwner) + { + state[node] = 1; // in-progress + path.Add(node); + + if (graph.TryGetValue(node, out var deps)) + { + foreach (var dep in deps) + { + state.TryGetValue(dep, out var depState); + if (depState == 1) + { + // Back-edge found — extract the cycle + var cycleStart = path.IndexOf(dep); + var cycle = path.GetRange(cycleStart, path.Count - cycleStart); + cycle.Add(dep); + string owner; + if (!nodeOwner.TryGetValue(node, out owner)) + owner = node; + throw new InvalidOperationException( + $"Cyclic Dependency Detected: {string.Join(" \u2192 ", cycle)} (via {owner})"); + } + + if (depState == 0 && graph.ContainsKey(dep)) + DfsCycleCheck(dep, graph, state, path, nodeOwner); + } + } + + path.RemoveAt(path.Count - 1); + state[node] = 2; // done } private static string Constructor(string usingStatements, string constructorFields, string constructor, string constructorAssignments, int dictSize, @@ -993,15 +1048,20 @@ private static string MakeMethodCall(ImmutableArray parameters, H /// in a generated constructor call. Collection params are converted from the cached /// IEnumerable<T> factory to the exact type requested. /// - private static string CollectionConstructorArg(ParameterData parameter) => - parameter.CollectionKind switch + private static string CollectionConstructorArg(ParameterData parameter) + { + if (!parameter.IsCollection) + return parameter.Name; + var memberName = "coll_" + parameter.CollectionElementMemberName!; + return parameter.CollectionKind switch { - CollectionKind.Array => $"{parameter.Name}.ToArray()", - CollectionKind.List => $"{parameter.Name}.ToList()", - CollectionKind.ImmutableArray => $"ImmutableArray.CreateRange({parameter.Name})", - CollectionKind.ReadOnlySpan => $"new global::System.ReadOnlySpan<{parameter.CollectionElementFullName}>({parameter.Name}.ToArray())", - _ => parameter.Name, // Enumerable or plain missing → use name directly + CollectionKind.Array => $"{memberName}.ToArray()", + CollectionKind.List => $"{memberName}.ToList()", + CollectionKind.ImmutableArray => $"ImmutableArray.CreateRange({memberName})", + CollectionKind.ReadOnlySpan => $"new global::System.ReadOnlySpan<{parameter.CollectionElementFullName}>({memberName}.ToArray())", + _ => memberName, // Enumerable → use element member name directly }; + } /// /// Returns the expression used in the Func<object> lambda inside the lookup dictionary @@ -1015,5 +1075,273 @@ private static string CollectionDictExpression(CollectionKind kind, string facto CollectionKind.ImmutableArray => $"ImmutableArray.CreateRange({factoryName})", _ => factoryName, // Enumerable → direct }; + + // ── C# 14 static-extension generation ──────────────────────────────────── + + private static bool IsAtLeastCSharp14(ParseOptions options, CancellationToken _) + { + if (options is not CSharpParseOptions csOptions) return false; + // C# 14 = 1400 in Roslyn's LanguageVersion enum. + // LanguageVersion.Preview == int.MaxValue, which is also >= 1400. + const int CSharp14 = 1400; + return (int)csOptions.LanguageVersion >= CSharp14; + } + + private static void MakeStaticExtensions( + SourceProductionContext context, + ((ImmutableArray Injections, Compilation Compilation) Left, bool SupportsExtensions) data) + { + if (!data.SupportsExtensions) return; + var source = GenerateStaticExtensions(data.Left.Injections, data.Left.Compilation); + context.AddSource("DependencyInjectionContainer.StaticExtensions.g.cs", source); + } + + private static string GenerateStaticExtensions( + ImmutableArray dataInjections, Compilation compilation) + { + var ordered = dataInjections.Reverse().ToList(); + foreach (var injection in ordered.ToArray()) + { + if (!injection.IsTestType) continue; + ordered.Remove(injection); + ordered.Add(injection); + } + + var interfaceInjectors = new Dictionary>(); + var interfaceMemberNames = new Dictionary(); + foreach (var injection in ordered) + { + for (var i = 0; i < injection.InterfaceFullNames.Length; i++) + { + var ifaceFull = injection.InterfaceFullNames[i]; + var ifaceMember = injection.InterfaceMemberNames[i]; + if (!interfaceInjectors.ContainsKey(ifaceFull)) + { + interfaceInjectors[ifaceFull] = new List(); + interfaceMemberNames[ifaceFull] = ifaceMember; + } + interfaceInjectors[ifaceFull].Add(injection); + } + } + + var availableInterfaces = interfaceInjectors.Keys.ToImmutableArray(); + + var sb = new StringBuilder(); + sb.AppendLine($@"using System.CodeDom.Compiler; +namespace {compilation.Assembly.Name}.Generated; +#nullable enable"); + + foreach (var kvp in interfaceMemberNames) + { + var ifaceFull = kvp.Key; + var ifaceMember = kvp.Value; + var possibilities = interfaceInjectors[ifaceFull]; + var className = ifaceMember + "Extensions"; + var body = StaticExtensionBody(ifaceFull, ifaceMember, possibilities, availableInterfaces); + + sb.AppendLine($@"[GeneratedCode(""{ToolName}"", ""{Version}"")] +public static class {className} +{{ + extension({ifaceFull}) + {{ +{body} + }} +}}"); + } + + return sb.ToString(); + } + + /// + /// Produces the method declaration(s) that go inside an extension(...) block + /// for the given interface. + /// + private static string StaticExtensionBody( + string ifaceFull, + string ifaceMember, + List possibilities, + ImmutableArray availableInterfaces) + { + var hasBooleans = possibilities.Any(p => p.BooleanInjection != null); + var chosen = possibilities.Last(); + + // Boolean-switched or lambda: the container already encodes all of that logic in its + // own factory method — call it directly (no dictionary) and fall back to inline + // construction when the container is absent. + if (hasBooleans || chosen.Lambda != null) + { + string nullFallback; + if (chosen.Lambda != null) + { + nullFallback = + $"throw new global::System.InvalidOperationException(" + + $"\"Cannot resolve {ifaceFull} without a container\")"; + } + else + { + var fallbackImpl = possibilities.LastOrDefault(p => p.BooleanInjection == null); + nullFallback = fallbackImpl != null + ? InlineCreation(fallbackImpl, availableInterfaces, nullContainer: true) + : $"throw new global::System.InvalidOperationException(" + + $"\"Cannot resolve {ifaceFull} without a container\")"; + } + return $" public static {ifaceFull} Resolve({ClassName}? container) =>" + + $" container?.{ifaceMember}() ?? {nullFallback};"; + } + + var creation = InlineCreation(chosen, availableInterfaces, nullContainer: false); + var nullCreation = InlineCreation(chosen, availableInterfaces, nullContainer: true); + + // Singleton / Scoped: inline double-checked locking directly against the container's + // cache field, bypassing both the dictionary and the factory-method call. + if (chosen.Singleton || chosen.Scoped) + { + if (chosen.Disposable) + { + return $@" public static {ifaceFull} Resolve({ClassName}? container) + {{ + if (container != null) + {{ + var cached = container.{chosen.LazyFieldName}; + if (cached != null) return cached; + lock (container.m_lock) + {{ + cached = container.{chosen.LazyFieldName}; + if (cached != null) return cached; + var value = {creation}; + container.GetResolvedInstances().Add(new global::System.WeakReference(value)); + container.{chosen.LazyFieldName} = value; + return value; + }} + }} + return {nullCreation}; + }}"; + } + + return $@" public static {ifaceFull} Resolve({ClassName}? container) + {{ + if (container != null) + {{ + var cached = container.{chosen.LazyFieldName}; + if (cached != null) return cached; + lock (container.m_lock) + {{ + cached = container.{chosen.LazyFieldName}; + if (cached != null) return cached; + return container.{chosen.LazyFieldName} = {creation}; + }} + }} + return {nullCreation}; + }}"; + } + + // Transient + disposable: construct and register for disposal tracking. + if (chosen.Disposable) + { + return $@" public static {ifaceFull} Resolve({ClassName}? container) + {{ + var value = {creation}; + container?.GetResolvedInstances().Add(new global::System.WeakReference(value)); + return value; + }}"; + } + + // Transient, non-disposable: pure inline factory chain. + return $" public static {ifaceFull} Resolve({ClassName}? container) => {creation};"; + } + + /// + /// Builds a new ConcreteType(...) expression for the given injection where every + /// DI-registered dependency is resolved via its own static Resolve extension. + /// + private static string InlineCreation( + InjectionData injection, + ImmutableArray availableInterfaces, + bool nullContainer) + { + HashSet? missing = null; + HashSet? nullableDefaults = null; + var ctor = GetBestConstructor(injection, availableInterfaces, ref missing, ref nullableDefaults); + + if (ctor == null) + return $"default! /* no constructor found for {injection.TypeFullName} */"; + + var resolveArg = nullContainer ? "null" : "container"; + var args = new List(); + + foreach (var parameter in ctor.Parameters) + { + // Nullable params with no DI registration → null + if (nullableDefaults?.Contains(parameter) == true) + { + args.Add("null"); + continue; + } + + var typeLookup = parameter.IsNullable + ? parameter.TypeFullName.TrimEnd('?') + : parameter.TypeFullName; + + // DI-registered: recurse via the static extension + if (availableInterfaces.Contains(typeLookup)) + { + args.Add($"{typeLookup}.Resolve({resolveArg})"); + continue; + } + + // C# will supply the default value — omit the argument + if (parameter.HasExplicitDefault) + continue; + + // Empty params array — omit + if (parameter.IsParams) + continue; + + // Collection: convert from the container's IEnumerable lazy property to the + // exact collection type the constructor demands, mirroring CollectionConstructorArg. + if (parameter.IsCollection && parameter.CollectionElementFullName != null) + { + var elemType = parameter.CollectionElementFullName; + string collectionArg; + if (nullContainer) + { + collectionArg = parameter.CollectionKind switch + { + CollectionKind.List => $"new global::System.Collections.Generic.List<{elemType}>()", + CollectionKind.ImmutableArray => $"global::System.Collections.Immutable.ImmutableArray<{elemType}>.Empty", + CollectionKind.ReadOnlySpan => $"global::System.ReadOnlySpan<{elemType}>.Empty", + _ => $"global::System.Array.Empty<{elemType}>()", // Array or Enumerable + }; + } + else + { + var src = $"container!.coll_{parameter.CollectionElementMemberName}"; + collectionArg = parameter.CollectionKind switch + { + CollectionKind.Array => $"{src}.ToArray()", + CollectionKind.List => $"{src}.ToList()", + CollectionKind.ImmutableArray => $"global::System.Collections.Immutable.ImmutableArray.CreateRange({src})", + CollectionKind.ReadOnlySpan => $"new global::System.ReadOnlySpan<{elemType}>({src}.ToArray())", + _ => src, // Enumerable — direct + }; + } + args.Add(collectionArg); + continue; + } + + if (parameter.IsNullable) + { + args.Add("null"); + continue; + } + + // Required non-DI parameter: sourced from the container's internal field. + // When nullContainer is true we use default! — callers accept that null-container + // mode cannot satisfy non-DI required dependencies. + args.Add(nullContainer ? "default!" : $"container!.{parameter.Name}"); + } + + return $"new {injection.TypeFullName}({string.Join(", ", args)})"; + } } } \ No newline at end of file diff --git a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj index fc8a819..e4de7bb 100644 --- a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj +++ b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj @@ -1,7 +1,7 @@  - net8.0 + net10.0 false true diff --git a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs index 5cceb2e..f29ca42 100644 --- a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs +++ b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs @@ -299,6 +299,51 @@ public void ReadOnlySpanConstructorParameterIsResolved() { m_container.Resolve().Count.ShouldBe(3); } + + // ── Cross-array reentrancy tests ────────────────────────────────────────── + // Ensures that reentrancy guards are per-array-type, not global. Resolving + // IEnumerable triggers construction of CrossA3 which needs + // IEnumerable. That second resolution must not be blocked. + + [Fact] + public void CrossArrayReentrancyResolvesAllA() + { + var items = m_container.Resolve().Items.ToList(); + items.Count.ShouldBe(3); + } + + [Fact] + public void CrossArrayReentrancyResolvesBInsideCrossA3() + { + var items = m_container.Resolve().Items.ToList(); + var crossA3 = items.OfType().ShouldHaveSingleItem(); + crossA3.Deps.Count().ShouldBe(2); + } + + // ── Inheritor + Base array tests ────────────────────────────────────────── + + [Fact] + public void InheritorAndBaseContainerMergeArrays() + { + var parent = new DependencyInjectionContainer(false, false, new NonInjectedClass()); + var child = new DependencyInjectionContainer(parent); + // Inherited defines SplitBase1 + SplitBase2 (2 items per container). + // Inheritor defines SplitInheritor1..3 (3 more per container). + // Each standalone container has 5. After merging, the child sees its own 5 + // plus the parent's 5 = 10. + child.Resolve().Items.Count().ShouldBe(10); + } + + [Fact] + public void BaseContainerSeesInheritorArraysAfterLinking() + { + var parent = new DependencyInjectionContainer(false, false, new NonInjectedClass()); + var child = new DependencyInjectionContainer(parent); + // After linking, the parent's Inheritor is the child. Resolving on the + // parent should now include its own 5 plus the child's 5 = 10. + parent.Resolve().Items.Count().ShouldBe(10); + } + private class DummyContainer : IContainer { public const string DummyText = "I am a bit of text"; diff --git a/Tests/TestData/Inherited/Types.cs b/Tests/TestData/Inherited/Types.cs index 70b1483..8b81859 100644 --- a/Tests/TestData/Inherited/Types.cs +++ b/Tests/TestData/Inherited/Types.cs @@ -237,4 +237,62 @@ public class ImmutableArrayConsumer(ImmutableArray arrays) public class ReadOnlySpanConsumer(ReadOnlySpan arrays) { public int Count { get; } = arrays.Length; +} + +// ── Cross-array reentrancy tests ───────────────────────────────────────────── +// Resolving IEnumerable should work even though CrossA3 depends on +// IEnumerable. The reentrancy flag is per-collection type, so +// resolving the B array must not be blocked by the A array's reentrancy guard. + +public interface ICrossArrayB; + +[Inject] +public class CrossB1 : ICrossArrayB; + +[Inject] +public class CrossB2 : ICrossArrayB; + +public interface ICrossArrayA; + +[Inject] +public class CrossA1 : ICrossArrayA; + +[Inject] +public class CrossA2 : ICrossArrayA; + +/// +/// Implementation of ICrossArrayA that depends on an array of ICrossArrayB. +/// When the container builds IEnumerable<ICrossArrayA> and encounters CrossA3 +/// it must resolve IEnumerable<ICrossArrayB>. This must succeed because the +/// reentrancy guard is local to each array type. +/// +[Inject] +public class CrossA3(IEnumerable deps) : ICrossArrayA +{ + public IEnumerable Deps { get; } = deps; +} + +[Inject, Self] +public class CrossArrayConsumer(IEnumerable items) +{ + public IEnumerable Items { get; } = items; +} + +// ── Inheritor + Base array tests ───────────────────────────────────────────── +// Interface whose implementations are split across the Inherited and Inheritor +// projects, so we can verify that arrays merge correctly across container +// hierarchies (both Base → child and Inheritor → child directions). + +public interface ISplitArray; + +[Inject] +public class SplitBase1 : ISplitArray; + +[Inject] +public class SplitBase2 : ISplitArray; + +[Inject, Self] +public class SplitArrayConsumer(IEnumerable items) +{ + public IEnumerable Items { get; } = items; } \ No newline at end of file diff --git a/Tests/TestData/Inheritor/Inheritor.csproj b/Tests/TestData/Inheritor/Inheritor.csproj index 1b4aa7e..4f0d091 100644 --- a/Tests/TestData/Inheritor/Inheritor.csproj +++ b/Tests/TestData/Inheritor/Inheritor.csproj @@ -14,6 +14,7 @@ enable enable true + preview \ No newline at end of file diff --git a/Tests/TestData/Inheritor/Types.cs b/Tests/TestData/Inheritor/Types.cs index cb2ad0b..917cd29 100644 --- a/Tests/TestData/Inheritor/Types.cs +++ b/Tests/TestData/Inheritor/Types.cs @@ -1,4 +1,5 @@ -using FactoryGenerator.Attributes; +using System.Runtime.InteropServices; +using FactoryGenerator.Attributes; using Inherited; using Inheritor.Generated; @@ -50,4 +51,17 @@ public static IEnumerable Method() var array = container.Resolve>(); return array; } -} \ No newline at end of file +} +// ── Inheritor + Base array tests ───────────────────────────────────────────── +// Additional ISplitArray implementations in the Inheritor project. When a child +// container is created from a parent, the merged IEnumerable should +// contain items from both Inherited (Base) and Inheritor. + +[Inject] +public class SplitInheritor1 : ISplitArray; + +[Inject] +public class SplitInheritor2 : ISplitArray; + +[Inject] +public class SplitInheritor3 : ISplitArray; \ No newline at end of file diff --git a/Tests/TestWebApp/TestWebApp.csproj b/Tests/TestWebApp/TestWebApp.csproj index be9b04a..333e941 100644 --- a/Tests/TestWebApp/TestWebApp.csproj +++ b/Tests/TestWebApp/TestWebApp.csproj @@ -1,7 +1,7 @@  - net8.0 + net10.0 enable enable From 736966388cf404a91e6949b03664ec3541daa5ec Mon Sep 17 00:00:00 2001 From: caran Date: Fri, 10 Apr 2026 17:38:20 +0200 Subject: [PATCH 02/21] TUnit Migration --- Directory.Packages.props | 15 +--- ...nerator.Extensions.AspNetCore.Tests.csproj | 7 +- .../IntegrationTests.cs | 5 +- .../FactoryGenerator.Tests.csproj | 8 +- Tests/FactoryGenerator.Tests/GlobalUsings.cs | 2 +- .../InjectionDetectionTests.cs | 83 ++++++++++--------- 6 files changed, 52 insertions(+), 68 deletions(-) diff --git a/Directory.Packages.props b/Directory.Packages.props index 70bdb68..d8cbf67 100644 --- a/Directory.Packages.props +++ b/Directory.Packages.props @@ -6,21 +6,14 @@ - - - + + + - - - - - - all - runtime; build; native; contentfiles; analyzers; buildtransitive - + all diff --git a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj index e4de7bb..c61af28 100644 --- a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj +++ b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj @@ -7,14 +7,9 @@ - + - - - runtime; build; native; contentfiles; analyzers; buildtransitive - all - diff --git a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs index a559d3c..d0b75e5 100644 --- a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs +++ b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs @@ -1,4 +1,4 @@ -using System.Net; +using System.Net; using System.Threading.Tasks; using Microsoft.AspNetCore.Builder; using Microsoft.AspNetCore.Hosting; @@ -7,7 +7,6 @@ using Microsoft.Extensions.DependencyInjection; using Microsoft.Extensions.Hosting; using Shouldly; -using Xunit; using FactoryGenerator; using FactoryGenerator.Attributes; using FactoryGenerator.Extensions.AspNetCore; @@ -37,7 +36,7 @@ public class OtherService : IOtherService public class IntegrationTests { - [Fact] + [Test] public async Task Middleware_Integrates_FactoryGenerator_With_RequestServices() { // Setup diff --git a/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj b/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj index 1bb0dd5..69c4846 100644 --- a/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj +++ b/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj @@ -9,14 +9,8 @@ - + - - - - runtime; build; native; contentfiles; analyzers; buildtransitive - all - diff --git a/Tests/FactoryGenerator.Tests/GlobalUsings.cs b/Tests/FactoryGenerator.Tests/GlobalUsings.cs index 8c927eb..7aa3922 100644 --- a/Tests/FactoryGenerator.Tests/GlobalUsings.cs +++ b/Tests/FactoryGenerator.Tests/GlobalUsings.cs @@ -1 +1 @@ -global using Xunit; \ No newline at end of file +global using TUnit.Core; \ No newline at end of file diff --git a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs index f29ca42..09085c4 100644 --- a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs +++ b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs @@ -11,13 +11,16 @@ public class InjectionDetectionTests() { private readonly IContainer m_container = new DependencyInjectionContainer(default, default, new NonInjectedClass()); - [Fact] + [After(Test)] + public void DisposeContainer() => m_container.Dispose(); + + [Test] public void InjectedTypesAreResolvable() { m_container.Resolve().ShouldBeOfType(); } - [Fact] + [Test] public void SingletonInjectionsResolveToTheSameInstanceEverytime() { var first = m_container.Resolve(); @@ -25,7 +28,7 @@ public void SingletonInjectionsResolveToTheSameInstanceEverytime() ReferenceEquals(first, second).ShouldBeTrue(); } - [Fact] + [Test] public void NonSingleInjectionsResolveToDifferentInstanceEverytime() { var first = m_container.Resolve(); @@ -33,7 +36,7 @@ public void NonSingleInjectionsResolveToDifferentInstanceEverytime() ReferenceEquals(first, second).ShouldBeFalse(); } - [Fact] + [Test] public void ResolveUsesArguments() { var dummy = new NonInjectedClass(); @@ -41,22 +44,22 @@ public void ResolveUsesArguments() myContainer.Resolve().NonInjectedClassArgument.ShouldBe(dummy); } - [Theory] - [InlineData(true, typeof(EnabledImplementation))] - [InlineData(false, typeof(FallbackImplementation))] + [Test] + [Arguments(true, typeof(EnabledImplementation))] + [Arguments(false, typeof(FallbackImplementation))] public void PickupSingleInjectionWithBoolean(bool value, System.Type expected) { var myContainer = new DependencyInjectionContainer(value, default, default!); myContainer.Resolve().ShouldBeOfType(expected); } - [Fact] + [Test] public void PickupSingleInjectionFromMethod() { m_container.Resolve().ShouldBeOfType(); } - [Fact] + [Test] public void DoNotPickupNonInjection() { try @@ -71,7 +74,7 @@ public void DoNotPickupNonInjection() true.ShouldBeFalse(); } - [Fact] + [Test] public void DontPickupIDisposable() { try @@ -86,7 +89,7 @@ public void DontPickupIDisposable() true.ShouldBeFalse(); } - [Fact] + [Test] public void DontPickupExcluded() { try @@ -101,26 +104,26 @@ public void DontPickupExcluded() true.ShouldBeFalse(); } - [Fact] + [Test] public void PickupTypesSpecifiedByAs() { m_container.Resolve().ShouldBeOfType(); } - [Fact] + [Test] public void PickupInheritedInterfaces() { m_container.Resolve().ShouldBeOfType(); } - [Fact] + [Test] public void InheritorsOverride() { m_container.Resolve().ShouldBeOfType(); } - [Fact] + [Test] public void DisposingContainerDisposesSingletons() { ISingletonDisposer singleton; @@ -132,7 +135,7 @@ public void DisposingContainerDisposesSingletons() ((DisposableSingleton) singleton).WasDisposed.ShouldBeTrue(); } - [Fact] + [Test] public void DisposingLifetimeContainerDoesNotDisposeSingletons() { ISingletonDisposer singleton; @@ -150,7 +153,7 @@ public void DisposingLifetimeContainerDoesNotDisposeSingletons() ((DisposableSingleton) singleton).WasDisposed.ShouldBeTrue(); } - [Fact] + [Test] public void DisposingLifetimeContainerDisposesScoped() { IScoped singleton; @@ -164,7 +167,7 @@ public void DisposingLifetimeContainerDisposesScoped() singleton.WasDisposed.ShouldBeTrue(); } - [Fact] + [Test] public void DisposingContainerDoesNotDisposeUntrackedInstances() { IDisposer singleton; @@ -176,50 +179,50 @@ public void DisposingContainerDoesNotDisposeUntrackedInstances() ((DisposableNonSingleton) singleton).WasDisposed.ShouldBeTrue(); } - [Fact] + [Test] public void DisposingContainerDoesNotDisposesUnreferencedSingletons() { using var myContainer = new DependencyInjectionContainer(false, default, default!); } - [Fact] + [Test] public void ArrayExpressionsCollect() { m_container.Resolve().Arrays.Count().ShouldBe(3); } - [Fact] + [Test] public void RequestedArraysArePresent() { Program.Method().Count().ShouldBe(3); } - [Fact] + [Test] public void BooleanFallbackIsOverriden() { m_container.Resolve().ShouldBeOfType(); } - [Fact] + [Test] public void TryResolveWithTypeArgumentsWorks() { m_container.TryResolve(out var type).ShouldBeTrue(); type.ShouldBeOfType(); } - [Fact] + [Test] public void TryResolveWithTypeParameterWorks() { m_container.TryResolve(typeof(IType), out var type).ShouldBeTrue(); type.ShouldBeOfType(); } - [Fact] + [Test] public void ClassesInsideOtherClassesCanBeInjected() { m_container.Resolve(); } - [Fact] + [Test] public void ContainerMayCreateItself() { var newContainer = new DependencyInjectionContainer(m_container); @@ -227,20 +230,20 @@ public void ContainerMayCreateItself() resolved.Count().ShouldBe(6); var nonInjected = m_container.Resolve(); } - [Fact] + [Test] public void HierarchicalContainersResolveArraysProperly() { var newContainer = new DependencyInjectionContainer(m_container); newContainer.Resolve().Arrays.Count().ShouldBe(6); } - [Fact] + [Test] public void HierarchicalContainersResolveUsesFallBackIfItCannotFindImplementation() { var newContainer = new DependencyInjectionContainer(new DummyContainer()); newContainer.Resolve().ShouldBe(DummyContainer.DummyText); } - [Fact] + [Test] public void ContainerPropgatesRelevantBooleansCreateItself() { var baseContainer = new DependencyInjectionContainer(true, false, new()); @@ -252,7 +255,7 @@ public void ContainerPropgatesRelevantBooleansCreateItself() newContainer.GetBoolean("A").ShouldBeFalse(); newContainer.GetBoolean("TestBool").ShouldBeTrue(); } - [Fact] + [Test] public void HierarchicalContainersPropgatesBooleansUnknownToIt() { var newContainer = new DependencyInjectionContainer(new DummyContainer()); @@ -262,13 +265,13 @@ public void HierarchicalContainersPropgatesBooleansUnknownToIt() // ── Nullable parameter tests ────────────────────────────────────────────── - [Fact] + [Test] public void NullableUnregisteredParameterDefaultsToNull() { m_container.Resolve().Optional.ShouldBeNull(); } - [Fact] + [Test] public void NullableRegisteredParameterIsResolved() { m_container.Resolve().Optional.ShouldBeOfType(); @@ -276,25 +279,25 @@ public void NullableRegisteredParameterIsResolved() // ── Collection constructor parameter tests ──────────────────────────────── - [Fact] + [Test] public void ArrayConstructorParameterIsResolved() { m_container.Resolve().Arrays.Length.ShouldBe(3); } - [Fact] + [Test] public void ListConstructorParameterIsResolved() { m_container.Resolve().Arrays.Count.ShouldBe(3); } - [Fact] + [Test] public void ImmutableArrayConstructorParameterIsResolved() { m_container.Resolve().Arrays.Length.ShouldBe(3); } - [Fact] + [Test] public void ReadOnlySpanConstructorParameterIsResolved() { m_container.Resolve().Count.ShouldBe(3); @@ -305,14 +308,14 @@ public void ReadOnlySpanConstructorParameterIsResolved() // IEnumerable triggers construction of CrossA3 which needs // IEnumerable. That second resolution must not be blocked. - [Fact] + [Test] public void CrossArrayReentrancyResolvesAllA() { var items = m_container.Resolve().Items.ToList(); items.Count.ShouldBe(3); } - [Fact] + [Test] public void CrossArrayReentrancyResolvesBInsideCrossA3() { var items = m_container.Resolve().Items.ToList(); @@ -322,7 +325,7 @@ public void CrossArrayReentrancyResolvesBInsideCrossA3() // ── Inheritor + Base array tests ────────────────────────────────────────── - [Fact] + [Test] public void InheritorAndBaseContainerMergeArrays() { var parent = new DependencyInjectionContainer(false, false, new NonInjectedClass()); @@ -334,7 +337,7 @@ public void InheritorAndBaseContainerMergeArrays() child.Resolve().Items.Count().ShouldBe(10); } - [Fact] + [Test] public void BaseContainerSeesInheritorArraysAfterLinking() { var parent = new DependencyInjectionContainer(false, false, new NonInjectedClass()); From 912aec84bd609d96fd44fdb6b60a303b27476c02 Mon Sep 17 00:00:00 2001 From: caran Date: Mon, 13 Apr 2026 09:05:23 +0200 Subject: [PATCH 03/21] Fix: Dotnet 10 in CICD and removed unused using. --- .github/workflows/benchmark.yml | 2 +- .github/workflows/build.yml | 10 +++--- Benchmarking/Benchmarks/Program.cs | 32 +++++++++---------- FactoryGenerator.Attributes/IContainer.cs | 1 - .../AppBuilderExtensions.cs | 4 +-- .../FactoryGeneratorMiddleware.cs | 2 -- .../FactoryGeneratorServiceProvider.cs | 1 - .../ServiceProviderExtensions.cs | 1 - FactoryGenerator/FactoryGenerator.cs | 1 + FactoryGenerator/InjectionData.cs | 1 - FactoryGenerator/Logger.cs | 1 - FactoryGenerator/SymbolUtility.cs | 1 - .../IntegrationTests.cs | 2 -- .../InjectionDetectionTests.cs | 1 - Tests/TestData/Inherited/Types.cs | 1 - Tests/TestData/Inheritor/Types.cs | 3 +- global.json | 5 +++ 17 files changed, 29 insertions(+), 40 deletions(-) create mode 100644 global.json diff --git a/.github/workflows/benchmark.yml b/.github/workflows/benchmark.yml index a48ddd7..c0a86dc 100644 --- a/.github/workflows/benchmark.yml +++ b/.github/workflows/benchmark.yml @@ -23,7 +23,7 @@ jobs: - name: Setup dotnet ${{ matrix.dotnet-version }} uses: actions/setup-dotnet@v4 with: - dotnet-version: "9.0.x" + dotnet-version: "10.0.x" - name: Run benchmark run: cd Benchmarking/Benchmarks && dotnet run -c Release --exporters json --filter '*' diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index 02c182e..29dddb6 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -28,7 +28,7 @@ jobs: - name: Setup dotnet ${{ matrix.dotnet-version }} uses: actions/setup-dotnet@v4 with: - dotnet-version: "9.0.x" + dotnet-version: "10.0.x" - name: Restore run: dotnet restore FactoryGenerator.sln - name: Build @@ -45,7 +45,7 @@ jobs: - name: Setup dotnet ${{ matrix.dotnet-version }} uses: actions/setup-dotnet@v4 with: - dotnet-version: "9.0.x" + dotnet-version: "10.0.x" - name: Pack Generator run: dotnet pack FactoryGenerator/FactoryGenerator.csproj -o "${{ env.NuGetDirectory }}" --property:RepositoryCommit="${{ env.COMMIT_SHA }}" --property:InformationalVersion="UNRELEASED" --property:AssemblyVersion="0.0.0" --property:FileVersion="0.0.0" --property:Version="0.0.0" - name: Pack Attributes @@ -68,7 +68,7 @@ jobs: - name: Setup dotnet ${{ matrix.dotnet-version }} uses: actions/setup-dotnet@v4 with: - dotnet-version: "9.0.x" + dotnet-version: "10.0.x" - name: Pack Generator run: dotnet pack FactoryGenerator/FactoryGenerator.csproj -o "${{ env.NuGetDirectory }}" --property:RepositoryCommit="${{ env.COMMIT_SHA }}" --property:InformationalVersion="${{ github.ref_name }}" --property:AssemblyVersion="${{ github.ref_name }}" --property:FileVersion="${{ github.ref_name }}" --property:Version="${{ github.ref_name }}" - name: Pack Attributes @@ -94,7 +94,7 @@ jobs: - name: Setup dotnet ${{ matrix.dotnet-version }} uses: actions/setup-dotnet@v4 with: - dotnet-version: "9.0.x" + dotnet-version: "10.0.x" - name: Publish Nuget packages run: | for file in $(find "${{ env.NuGetDirectory }}" -type f -name "*.nupkg"); do @@ -110,7 +110,7 @@ jobs: - name: Setup dotnet ${{ matrix.dotnet-version }} uses: actions/setup-dotnet@v4 with: - dotnet-version: "9.0.x" + dotnet-version: "10.0.x" - name: Run benchmark run: cd Benchmarking/Benchmarks && dotnet run -c Release --exporters json --filter '*' diff --git a/Benchmarking/Benchmarks/Program.cs b/Benchmarking/Benchmarks/Program.cs index 4101508..dc59dcb 100644 --- a/Benchmarking/Benchmarks/Program.cs +++ b/Benchmarking/Benchmarks/Program.cs @@ -36,34 +36,32 @@ public class ResolveBenchmarks [Benchmark] public IContainer CreateFromSelf() => new DependencyInjectionContainer(m_container); -} - -// ── Static-extension resolution (C# 14 / .NET 10+ path) ───────────────────── -// -// Each Resolve(container?) call inlines the full construction chain directly — -// no dictionary lookup, no factory-method indirection. -// -// Null-container variants bypass the singleton cache entirely and perform a -// fresh allocation on every call, exposing the raw construction cost. -[MemoryDiagnoser] -[JsonExporterAttribute.Full] -[JsonExporterAttribute.FullCompressed] -public class StaticExtensionBenchmarks -{ - private readonly DependencyInjectionContainer m_container = new(default, default, new NonInjectedClass()); + // ── Static-extension resolution (C# 14 / .NET 10+ path) ───────────────────── + // + // Each Resolve(container?) call inlines the full construction chain directly — + // no dictionary lookup, no factory-method indirection. + // + // Null-container variants bypass the singleton cache entirely and perform a + // fresh allocation on every call, exposing the raw construction cost. [Benchmark] public ISingleton ExtensionResolveSingleton() => ISingleton.Resolve(m_container); + [Benchmark] public ISingleton ExtensionResolveSingletonNullContainer() => ISingleton.Resolve(null); + [Benchmark] public IOverridable ExtensionResolveTransient() => IOverridable.Resolve(m_container); + [Benchmark] public ChainA ExtensionResolveChain() => ChainA.Resolve(m_container); + [Benchmark] public ChainA ExtensionResolveChainNullContainer() => ChainA.Resolve(null); + [Benchmark] public ArrayConsumer ExtensionResolveWithCollection() => ArrayConsumer.Resolve(m_container); + [Benchmark] public ArrayConsumer ExtensionResolveWithCollectionNullContainer() => ArrayConsumer.Resolve(null); } @@ -71,5 +69,5 @@ public class StaticExtensionBenchmarks internal static class Program { private static void Main(string[] args) => - BenchmarkSwitcher.FromAssembly(typeof(Program).Assembly).Run(args); -} + BenchmarkRunner.Run(); +} \ No newline at end of file diff --git a/FactoryGenerator.Attributes/IContainer.cs b/FactoryGenerator.Attributes/IContainer.cs index b367ce9..68a4891 100644 --- a/FactoryGenerator.Attributes/IContainer.cs +++ b/FactoryGenerator.Attributes/IContainer.cs @@ -1,5 +1,4 @@ using System; -using System.Collections; using System.Collections.Generic; namespace FactoryGenerator; diff --git a/FactoryGenerator.Extensions.AspNetCore/AppBuilderExtensions.cs b/FactoryGenerator.Extensions.AspNetCore/AppBuilderExtensions.cs index 59260ef..5009be6 100644 --- a/FactoryGenerator.Extensions.AspNetCore/AppBuilderExtensions.cs +++ b/FactoryGenerator.Extensions.AspNetCore/AppBuilderExtensions.cs @@ -1,6 +1,4 @@ -using System; -using FactoryGenerator; -using Microsoft.AspNetCore.Builder; +using Microsoft.AspNetCore.Builder; namespace FactoryGenerator.Extensions.AspNetCore; diff --git a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs index 94a7e24..9c77c88 100644 --- a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs +++ b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs @@ -1,7 +1,5 @@ using System.Threading.Tasks; -using FactoryGenerator; using Microsoft.AspNetCore.Http; -using Microsoft.Extensions.DependencyInjection; namespace FactoryGenerator.Extensions.AspNetCore; diff --git a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorServiceProvider.cs b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorServiceProvider.cs index fdfdcad..826282b 100644 --- a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorServiceProvider.cs +++ b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorServiceProvider.cs @@ -1,5 +1,4 @@ using System; -using FactoryGenerator; using Microsoft.Extensions.DependencyInjection; namespace FactoryGenerator.Extensions.AspNetCore; diff --git a/FactoryGenerator.Extensions.AspNetCore/ServiceProviderExtensions.cs b/FactoryGenerator.Extensions.AspNetCore/ServiceProviderExtensions.cs index 1bf54f0..d045e3e 100644 --- a/FactoryGenerator.Extensions.AspNetCore/ServiceProviderExtensions.cs +++ b/FactoryGenerator.Extensions.AspNetCore/ServiceProviderExtensions.cs @@ -1,5 +1,4 @@ using System; -using FactoryGenerator; namespace FactoryGenerator.Extensions.AspNetCore; diff --git a/FactoryGenerator/FactoryGenerator.cs b/FactoryGenerator/FactoryGenerator.cs index 50834d0..3ce0112 100644 --- a/FactoryGenerator/FactoryGenerator.cs +++ b/FactoryGenerator/FactoryGenerator.cs @@ -1128,6 +1128,7 @@ private static string GenerateStaticExtensions( var sb = new StringBuilder(); sb.AppendLine($@"using System.CodeDom.Compiler; +using System.Linq; namespace {compilation.Assembly.Name}.Generated; #nullable enable"); diff --git a/FactoryGenerator/InjectionData.cs b/FactoryGenerator/InjectionData.cs index 2ad2cc1..249937f 100644 --- a/FactoryGenerator/InjectionData.cs +++ b/FactoryGenerator/InjectionData.cs @@ -1,5 +1,4 @@ using System; -using System.Collections.Generic; using System.Collections.Immutable; using System.Linq; diff --git a/FactoryGenerator/Logger.cs b/FactoryGenerator/Logger.cs index 7c95cce..618c82f 100644 --- a/FactoryGenerator/Logger.cs +++ b/FactoryGenerator/Logger.cs @@ -1,5 +1,4 @@ using System; -using System.Diagnostics; using System.IO; namespace FactoryGenerator diff --git a/FactoryGenerator/SymbolUtility.cs b/FactoryGenerator/SymbolUtility.cs index f28762d..3028e43 100644 --- a/FactoryGenerator/SymbolUtility.cs +++ b/FactoryGenerator/SymbolUtility.cs @@ -1,4 +1,3 @@ -using System; using System.Collections.Generic; using System.Text; using Microsoft.CodeAnalysis; diff --git a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs index d0b75e5..e93bf3e 100644 --- a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs +++ b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs @@ -7,9 +7,7 @@ using Microsoft.Extensions.DependencyInjection; using Microsoft.Extensions.Hosting; using Shouldly; -using FactoryGenerator; using FactoryGenerator.Attributes; -using FactoryGenerator.Extensions.AspNetCore; namespace FactoryGenerator.Extensions.AspNetCore.Tests; diff --git a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs index 09085c4..89b986b 100644 --- a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs +++ b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs @@ -1,4 +1,3 @@ -using System.ComponentModel; using Inherited; using Inheritor; using Inheritor.Generated; diff --git a/Tests/TestData/Inherited/Types.cs b/Tests/TestData/Inherited/Types.cs index 8b81859..98d5125 100644 --- a/Tests/TestData/Inherited/Types.cs +++ b/Tests/TestData/Inherited/Types.cs @@ -1,5 +1,4 @@ using FactoryGenerator.Attributes; -using System.Collections.Generic; using System.Collections.Immutable; namespace Inherited; diff --git a/Tests/TestData/Inheritor/Types.cs b/Tests/TestData/Inheritor/Types.cs index 917cd29..25dacab 100644 --- a/Tests/TestData/Inheritor/Types.cs +++ b/Tests/TestData/Inheritor/Types.cs @@ -1,5 +1,4 @@ -using System.Runtime.InteropServices; -using FactoryGenerator.Attributes; +using FactoryGenerator.Attributes; using Inherited; using Inheritor.Generated; diff --git a/global.json b/global.json new file mode 100644 index 0000000..e163e86 --- /dev/null +++ b/global.json @@ -0,0 +1,5 @@ +{ + "test": { + "runner": "Microsoft.Testing.Platform" + } +} \ No newline at end of file From 05aee5bb7e5dcfc47c5c1e5ac950de71cd4f9ed1 Mon Sep 17 00:00:00 2001 From: caran Date: Mon, 13 Apr 2026 09:08:45 +0200 Subject: [PATCH 04/21] Fix syntax for MTP --- .github/workflows/build.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index 29dddb6..c07a597 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -34,7 +34,7 @@ jobs: - name: Build run: dotnet build FactoryGenerator.sln --no-restore - name: Test - run: dotnet test FactoryGenerator.sln --no-build --no-restore + run: dotnet test --no-build --no-restore --solution FactoryGenerator.sln pack: runs-on: ubuntu-latest From 32ef5267cf5a5c9410348e4d8dfc7dcfe2919600 Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 09:44:16 +0200 Subject: [PATCH 05/21] Test fixes --- Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs b/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs index 5e9b480..5456778 100644 --- a/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs +++ b/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs @@ -7,7 +7,7 @@ namespace FactoryGenerator.Tests; public class ContainerRegistryTests { - [Fact] + [Test] public void ContainerEntryPointRegistersOnModuleLoad() { // The Inheritor assembly's ModuleInitializer should have already registered @@ -15,7 +15,7 @@ public void ContainerEntryPointRegistersOnModuleLoad() ContainerRegistry.RegisteredAssemblies.ShouldContain("Inheritor"); } - [Fact] + [Test] public void ContainerEntryPointCreateBuildsWorkingContainer() { // Create a base container @@ -28,7 +28,7 @@ public void ContainerEntryPointCreateBuildsWorkingContainer() chained.ShouldBeAssignableTo(); } - [Fact] + [Test] public void BuildChainCreatesWorkingContainerPipeline() { // Create a base container @@ -42,7 +42,7 @@ public void BuildChainCreatesWorkingContainerPipeline() final.Resolve().ShouldNotBeNull(); } - [Fact] + [Test] public void ContainerEntryPointAssemblyNameIsCorrect() { ContainerEntryPoint.AssemblyName.ShouldBe("Inheritor"); From db96eeb5b7610e3f6f7984e5354e30f9d0e9c44e Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 09:56:03 +0200 Subject: [PATCH 06/21] Fixed a small bug --- FactoryGenerator/FactoryGenerator.cs | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/FactoryGenerator/FactoryGenerator.cs b/FactoryGenerator/FactoryGenerator.cs index 3ce0112..e14f2ae 100644 --- a/FactoryGenerator/FactoryGenerator.cs +++ b/FactoryGenerator/FactoryGenerator.cs @@ -1239,6 +1239,15 @@ private static string StaticExtensionBody( // Transient + disposable: construct and register for disposal tracking. if (chosen.Disposable) { + if (creation != nullCreation) + { + return $@" public static {ifaceFull} Resolve({ClassName}? container) + {{ + var value = container != null ? {creation} : {nullCreation}; + container?.GetResolvedInstances().Add(new global::System.WeakReference(value)); + return value; + }}"; + } return $@" public static {ifaceFull} Resolve({ClassName}? container) {{ var value = {creation}; @@ -1248,6 +1257,10 @@ private static string StaticExtensionBody( } // Transient, non-disposable: pure inline factory chain. + if (creation != nullCreation) + { + return $" public static {ifaceFull} Resolve({ClassName}? container) => container != null ? {creation} : {nullCreation};"; + } return $" public static {ifaceFull} Resolve({ClassName}? container) => {creation};"; } From 3008eb8acee85b69aa40a5169ae9f29113aa84de Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 09:59:24 +0200 Subject: [PATCH 07/21] docs: Add Static Extensions and Plugin Architecture to README Document the C# 14 static extension methods feature that provides dictionary-free inline resolution via ISomething.Resolve(container). Add Plugin Architecture to the features list in both READMEs. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- README.md | 23 +++++++++++++++++++++++ README_NUGET.md | 2 ++ 2 files changed, 25 insertions(+) diff --git a/README.md b/README.md index fe2fb0b..46ba2c2 100644 --- a/README.md +++ b/README.md @@ -11,6 +11,8 @@ with [Autofac](https://autofac.org/) beyond syntax choices. - **Attribute-based Generation:** Simply decorate your code with attributes like ```[Inject]```,```[Singleton]```,```[Self]``` and more and your IoC container will be woven together. - **Test-Overridability:** Need to swap out one injection for another to test something? Simply ```[Inject]``` a replacement inside your test project for a new container. +- **Static Extensions (C# 14+):** On .NET 10 and later, every registered interface gains a static ```Resolve``` method that inlines the full construction chain — no dictionary, no virtual dispatch. +- **Plugin Architecture:** Load AOT-compiled plugin assemblies at runtime and chain their containers together without reflection. ## Documentation @@ -167,6 +169,27 @@ public class Provider : IProvider ``` With this code, it is now possible to do `container.Resolve()`, which will effectively return the result of `new Provider().Method()`, although, since `Method` is `[Inject]`ed as a `[Singleton]`, the result will be cached and the same instance will be returned at every call to `Resolve` as well as shared between all Injected implementations that require a `IResultType`. +### Static Extensions (C# 14 / .NET 10+) + +When targeting C# 14 or later, FactoryGenerator automatically emits [static extension methods](https://learn.microsoft.com/en-us/dotnet/csharp/whats-new/csharp-14#extension-members) for every registered interface. This provides a dictionary-free, inline resolution path that the JIT can aggressively optimize. + +Instead of `container.Resolve()`, you can call: +```csharp +var singleton = ISingleton.Resolve(container); +var transient = IOverridable.Resolve(container); +var chain = ChainA.Resolve(container); +``` + +Each generated `Resolve` method inlines the full construction chain directly — no dictionary lookup, no factory-method indirection. Singletons use double-checked locking against the container's cache field, while transients emit a pure `new` expression. + +**Null-container mode:** Passing `null` instead of a container instance bypasses the singleton/scoped cache entirely and performs a fresh allocation on every call. This is useful for one-off instances where you want pure construction cost without any shared state: +```csharp +// Fresh allocation every time — no singleton cache +var fresh = ISingleton.Resolve(null); +``` + +The static extensions are generated alongside the standard dictionary-based container and require no configuration. If the consuming project's language version is below C# 14, the extensions are simply not emitted. + ### ASP.NET Core Integration For web applications, you can integrate FactoryGenerator with the standard `IServiceProvider`. diff --git a/README_NUGET.md b/README_NUGET.md index 8dd5b04..08b748b 100644 --- a/README_NUGET.md +++ b/README_NUGET.md @@ -5,3 +5,5 @@ with [Autofac](https://autofac.org/) beyond syntax choices. - **Attribute-based Generation:** Simply decorate your code with attributes like ```[Inject]```,```[Singleton]```,```[Self]``` and more and your IoC container will be woven together. - **Test-Overridability:** Need to swap out one injection for another to test something? Simply ```[Inject]``` a replacement inside your test project for a new container. - **ASP.NET Core Integration:** Seamlessly integrate your source-generated container with the standard ASP.NET Core web pipeline. +- **Static Extensions (C# 14+):** On .NET 10 and later, every registered interface gains a static ```Resolve``` method that inlines the full construction chain — no dictionary, no virtual dispatch. +- **Plugin Architecture:** Load AOT-compiled plugin assemblies at runtime and chain their containers together without reflection. From 989a03208ac873c375777e2f1867da6e2534f502 Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 10:08:12 +0200 Subject: [PATCH 08/21] dotnet 10 --- Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj | 2 +- Tests/TestData/Inherited/Inherited.csproj | 2 +- Tests/TestData/Inheritor/Inheritor.csproj | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj b/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj index 69c4846..6da67f8 100644 --- a/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj +++ b/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj @@ -1,7 +1,7 @@ - net9.0 + net10.0 enable enable false diff --git a/Tests/TestData/Inherited/Inherited.csproj b/Tests/TestData/Inherited/Inherited.csproj index 039b8de..6ffb18e 100644 --- a/Tests/TestData/Inherited/Inherited.csproj +++ b/Tests/TestData/Inherited/Inherited.csproj @@ -5,7 +5,7 @@ - net9.0 + net10.0 enable enable diff --git a/Tests/TestData/Inheritor/Inheritor.csproj b/Tests/TestData/Inheritor/Inheritor.csproj index 4f0d091..9536b82 100644 --- a/Tests/TestData/Inheritor/Inheritor.csproj +++ b/Tests/TestData/Inheritor/Inheritor.csproj @@ -9,7 +9,7 @@ - net9.0 + net10.0 true enable enable From b920fba9d19e6903128a97ca1d19b36e9df55d92 Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 10:17:18 +0200 Subject: [PATCH 09/21] feat: Add support for configurable static extensions generation --- .gitignore | 1 - FactoryGenerator/FactoryGenerator.cs | 12 +++++++++++- FactoryGenerator/FactoryGenerator.csproj | 2 ++ FactoryGenerator/build/FactoryGenerator.props | 11 +++++++++++ README.md | 9 ++++++++- ...ctoryGenerator.Extensions.AspNetCore.Tests.csproj | 1 + Tests/TestWebApp/TestWebApp.csproj | 1 + 7 files changed, 34 insertions(+), 3 deletions(-) create mode 100644 FactoryGenerator/build/FactoryGenerator.props diff --git a/.gitignore b/.gitignore index 83371ec..45b9c53 100644 --- a/.gitignore +++ b/.gitignore @@ -21,7 +21,6 @@ [Rr]elease-x86/ [Dd]ebug-x86/ x64/ -build/ bld/ [Bb]in/ [Oo]bj/ diff --git a/FactoryGenerator/FactoryGenerator.cs b/FactoryGenerator/FactoryGenerator.cs index e14f2ae..64e1b9e 100644 --- a/FactoryGenerator/FactoryGenerator.cs +++ b/FactoryGenerator/FactoryGenerator.cs @@ -36,7 +36,10 @@ public void Initialize(IncrementalGeneratorInitializationContext context) context.RegisterSourceOutput(combined, MakeAutofacModule); var supportsStaticExtensions = context.ParseOptionsProvider.Select(IsAtLeastCSharp14); - var extensionData = attributes.Combine(compilation).Combine(supportsStaticExtensions); + var emitStaticExtensions = context.AnalyzerConfigOptionsProvider.Select(GetEmitStaticExtensions); + var staticExtensionsEnabled = supportsStaticExtensions.Combine(emitStaticExtensions) + .Select(static (pair, _) => pair.Left && pair.Right); + var extensionData = attributes.Combine(compilation).Combine(staticExtensionsEnabled); context.RegisterSourceOutput(extensionData, MakeStaticExtensions); } @@ -1087,6 +1090,13 @@ private static bool IsAtLeastCSharp14(ParseOptions options, CancellationToken _) return (int)csOptions.LanguageVersion >= CSharp14; } + private static bool GetEmitStaticExtensions(AnalyzerConfigOptionsProvider provider, CancellationToken _) + { + if (!provider.GlobalOptions.TryGetValue($"build_property.{nameof(FactoryGenerator)}_EmitStaticExtensions", out var value)) + return true; + return !string.Equals(value, "false", StringComparison.OrdinalIgnoreCase); + } + private static void MakeStaticExtensions( SourceProductionContext context, ((ImmutableArray Injections, Compilation Compilation) Left, bool SupportsExtensions) data) diff --git a/FactoryGenerator/FactoryGenerator.csproj b/FactoryGenerator/FactoryGenerator.csproj index 2be0a66..a1fc6d2 100644 --- a/FactoryGenerator/FactoryGenerator.csproj +++ b/FactoryGenerator/FactoryGenerator.csproj @@ -16,5 +16,7 @@ + + diff --git a/FactoryGenerator/build/FactoryGenerator.props b/FactoryGenerator/build/FactoryGenerator.props new file mode 100644 index 0000000..5e84143 --- /dev/null +++ b/FactoryGenerator/build/FactoryGenerator.props @@ -0,0 +1,11 @@ + + + + true + + + + + + + diff --git a/README.md b/README.md index 46ba2c2..34ddf85 100644 --- a/README.md +++ b/README.md @@ -188,7 +188,14 @@ Each generated `Resolve` method inlines the full construction chain directly — var fresh = ISingleton.Resolve(null); ``` -The static extensions are generated alongside the standard dictionary-based container and require no configuration. If the consuming project's language version is below C# 14, the extensions are simply not emitted. +The static extensions are generated alongside the standard dictionary-based container and require no additional configuration. If the consuming project's language version is below C# 14, the extensions are simply not emitted. + +**Opting out:** If you are on C# 14+ but do not want the static extensions (for example, to reduce generated code size or avoid conflicts), set the following property in your `.csproj`: +```xml + + false + +``` ### ASP.NET Core Integration For web applications, you can integrate FactoryGenerator with the standard `IServiceProvider`. diff --git a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj index c61af28..45c1d4a 100644 --- a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj +++ b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj @@ -16,6 +16,7 @@ + diff --git a/Tests/TestWebApp/TestWebApp.csproj b/Tests/TestWebApp/TestWebApp.csproj index 333e941..09bfeaa 100644 --- a/Tests/TestWebApp/TestWebApp.csproj +++ b/Tests/TestWebApp/TestWebApp.csproj @@ -10,6 +10,7 @@ + From b75c784925936271280c934fd1c8cd7ef9160c64 Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 10:17:34 +0200 Subject: [PATCH 10/21] Missed add --- Tests/TestData/Inheritor/Inheritor.csproj | 1 + 1 file changed, 1 insertion(+) diff --git a/Tests/TestData/Inheritor/Inheritor.csproj b/Tests/TestData/Inheritor/Inheritor.csproj index 9536b82..7e0c344 100644 --- a/Tests/TestData/Inheritor/Inheritor.csproj +++ b/Tests/TestData/Inheritor/Inheritor.csproj @@ -6,6 +6,7 @@ + From 87bc4c6ec5a0cc04db5782c69e86b5e292344ba3 Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 10:21:50 +0200 Subject: [PATCH 11/21] Force load --- Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs | 3 +++ 1 file changed, 3 insertions(+) diff --git a/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs b/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs index 5456778..81ff654 100644 --- a/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs +++ b/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs @@ -10,6 +10,9 @@ public class ContainerRegistryTests [Test] public void ContainerEntryPointRegistersOnModuleLoad() { + // Force the Inheritor assembly to load, triggering its ModuleInitializer + _ = typeof(ContainerEntryPoint); + // The Inheritor assembly's ModuleInitializer should have already registered // its container factory in ContainerRegistry when the assembly was loaded. ContainerRegistry.RegisteredAssemblies.ShouldContain("Inheritor"); From 84901db234d723b8f4b301ec13d4236b4ebb7751 Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 10:34:15 +0200 Subject: [PATCH 12/21] Dependency updates --- Directory.Packages.props | 47 +++++++++---------- FactoryGenerator/FactoryGenerator.csproj | 1 - .../FactoryGenerator.Tests.csproj | 2 - 3 files changed, 23 insertions(+), 27 deletions(-) diff --git a/Directory.Packages.props b/Directory.Packages.props index d8cbf67..53c8a0a 100644 --- a/Directory.Packages.props +++ b/Directory.Packages.props @@ -1,26 +1,25 @@ - - true - true - - - - - - - - - - - - - - - all - runtime; build; native; contentfiles; analyzers; buildtransitive - - - - - + + true + true + + + + + + + + + + + + + + all + runtime; build; native; contentfiles; analyzers; buildtransitive + + + + + \ No newline at end of file diff --git a/FactoryGenerator/FactoryGenerator.csproj b/FactoryGenerator/FactoryGenerator.csproj index a1fc6d2..a27e13a 100644 --- a/FactoryGenerator/FactoryGenerator.csproj +++ b/FactoryGenerator/FactoryGenerator.csproj @@ -12,7 +12,6 @@ - diff --git a/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj b/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj index 6da67f8..3472792 100644 --- a/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj +++ b/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj @@ -11,8 +11,6 @@ - - From 2baaaab0486f98ab25269f202be756337078f571 Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 13:04:34 +0200 Subject: [PATCH 13/21] More tests and fixes --- .../ContainerRegistry.cs | 26 + FactoryGenerator.Attributes/IContainer.cs | 15 + .../FactoryGeneratorMiddleware.cs | 5 +- .../ServiceProviderAdapter.cs | 28 +- FactoryGenerator/FactoryGenerator.cs | 1500 +++++++++++++---- FactoryGenerator/Injection.cs | 4 +- FactoryGenerator/InjectionData.cs | 28 - README.md | 11 +- .../IntegrationTests.cs | 78 +- .../ContainerRegistryTests.cs | 24 + .../FactoryGenerator.Tests.csproj | 2 + .../InjectionDetectionTests.cs | 116 ++ Tests/TestData/Inherited/Types.cs | 70 + Tests/TestData/Inheritor/Types.cs | 10 + 14 files changed, 1501 insertions(+), 416 deletions(-) diff --git a/FactoryGenerator.Attributes/ContainerRegistry.cs b/FactoryGenerator.Attributes/ContainerRegistry.cs index 65e43c6..6031c61 100644 --- a/FactoryGenerator.Attributes/ContainerRegistry.cs +++ b/FactoryGenerator.Attributes/ContainerRegistry.cs @@ -42,9 +42,12 @@ public static IContainer BuildChain(IContainer baseContainer) snapshot = s_registrations.OrderBy(r => r.Priority).ToList(); } + var existingAssemblies = GetAssemblyNames(baseContainer); var current = baseContainer; foreach (var registration in snapshot) { + if (!existingAssemblies.Add(registration.AssemblyName)) + continue; current = registration.Factory(current); } @@ -65,9 +68,13 @@ public static IContainer BuildChain(IContainer baseContainer, IEnumerable(s_registrations); } + var existingAssemblies = GetAssemblyNames(baseContainer); var current = baseContainer; foreach (var name in assemblyNames) { + if (!existingAssemblies.Add(name)) + continue; + var registration = snapshot.Find(r => r.AssemblyName == name); if (registration == null) { @@ -82,6 +89,25 @@ public static IContainer BuildChain(IContainer baseContainer, IEnumerable GetAssemblyNames(IContainer container) + { + var names = new HashSet(StringComparer.Ordinal); + + for (var current = container; current is not null; current = current.Base) + { + if (current is IContainerRegistrationMetadata metadata) + names.Add(metadata.AssemblyName); + } + + for (var current = container.Inheritor; current is not null; current = current.Inheritor) + { + if (current is IContainerRegistrationMetadata metadata) + names.Add(metadata.AssemblyName); + } + + return names; + } + /// /// Returns the names of all currently registered container assemblies. /// diff --git a/FactoryGenerator.Attributes/IContainer.cs b/FactoryGenerator.Attributes/IContainer.cs index 68a4891..8b5e9b2 100644 --- a/FactoryGenerator.Attributes/IContainer.cs +++ b/FactoryGenerator.Attributes/IContainer.cs @@ -24,4 +24,19 @@ public interface IContainer : ILifetimeScope { IContainer? Base { get; } IContainer? Inheritor { get; set; } +} + +public interface IContainerScopeFactory +{ + ILifetimeScope BeginLifetimeScope(IContainer? baseContainer); +} + +public interface IContainerRegistrationMetadata +{ + string AssemblyName { get; } +} + +public interface IContainerCacheInvalidator +{ + void InvalidateCollectionCaches(); } \ No newline at end of file diff --git a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs index 9c77c88..9f328d5 100644 --- a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs +++ b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs @@ -16,8 +16,11 @@ public FactoryGeneratorMiddleware(RequestDelegate next, IContainer container) public async Task Invoke(HttpContext context) { - var scope = _container.BeginLifetimeScope(); var originalProvider = context.RequestServices; + var requestContainer = new ServiceProviderAdapter(originalProvider, baseContainer: _container); + var scope = _container is IContainerScopeFactory scopeFactory + ? scopeFactory.BeginLifetimeScope(requestContainer) + : _container.BeginLifetimeScope(); context.RequestServices = new FactoryGeneratorServiceProvider(originalProvider, scope); try diff --git a/FactoryGenerator.Extensions.AspNetCore/ServiceProviderAdapter.cs b/FactoryGenerator.Extensions.AspNetCore/ServiceProviderAdapter.cs index a2dfbc2..7c4aba4 100644 --- a/FactoryGenerator.Extensions.AspNetCore/ServiceProviderAdapter.cs +++ b/FactoryGenerator.Extensions.AspNetCore/ServiceProviderAdapter.cs @@ -9,14 +9,16 @@ internal sealed class ServiceProviderAdapter : IContainer, IDisposable { private readonly IServiceProvider _serviceProvider; private readonly IServiceScope? _serviceScope; + private readonly IContainer? _baseContainer; - public ServiceProviderAdapter(IServiceProvider serviceProvider, IServiceScope? serviceScope = null) + public ServiceProviderAdapter(IServiceProvider serviceProvider, IServiceScope? serviceScope = null, IContainer? baseContainer = null) { _serviceProvider = serviceProvider; _serviceScope = serviceScope; + _baseContainer = baseContainer; } - public IContainer? Base => null; + public IContainer? Base => _baseContainer; public IContainer? Inheritor { get; set; } public void Dispose() @@ -28,6 +30,7 @@ public T Resolve() { var service = _serviceProvider.GetService(); if (service != null) return service; + if (_baseContainer is not null) return _baseContainer.Resolve(); throw new KeyNotFoundException($"The type {typeof(T)} has not been registered in the IServiceProvider."); } @@ -35,40 +38,49 @@ public object Resolve(Type type) { var service = _serviceProvider.GetService(type); if (service != null) return service; + if (_baseContainer is not null) return _baseContainer.Resolve(type); throw new KeyNotFoundException($"The type {type} has not been registered in the IServiceProvider."); } public bool TryResolve(Type type, out object? resolved) { resolved = _serviceProvider.GetService(type); - return resolved != null; + if (resolved is not null) return true; + if (_baseContainer is not null) return _baseContainer.TryResolve(type, out resolved); + return false; } public bool TryResolve(out T? resolved) { resolved = _serviceProvider.GetService(); - return resolved != null; + if (resolved is not null) return true; + if (_baseContainer is not null) return _baseContainer.TryResolve(out resolved); + return false; } public bool IsRegistered(Type type) { // IServiceProvider doesn't have a reliable IsRegistered method without resolution. // We return true if it can be resolved. - return _serviceProvider.GetService(type) != null; + return _serviceProvider.GetService(type) != null || _baseContainer?.IsRegistered(type) == true; } public bool IsRegistered() => IsRegistered(typeof(T)); - public bool GetBoolean(string key) => false; + public bool GetBoolean(string key) => _baseContainer?.GetBoolean(key) == true; public IEnumerable<(string Key, bool Value)> GetBooleans() { - yield break; + if (_baseContainer is null) + yield break; + + foreach (var boolean in _baseContainer.GetBooleans()) + yield return boolean; } public ILifetimeScope BeginLifetimeScope() { var scope = _serviceProvider.CreateScope(); - return new ServiceProviderAdapter(scope.ServiceProvider, scope); + return new ServiceProviderAdapter(scope.ServiceProvider, scope, _baseContainer); } } diff --git a/FactoryGenerator/FactoryGenerator.cs b/FactoryGenerator/FactoryGenerator.cs index 64e1b9e..a7f7e2c 100644 --- a/FactoryGenerator/FactoryGenerator.cs +++ b/FactoryGenerator/FactoryGenerator.cs @@ -6,7 +6,6 @@ using System.Threading; using Microsoft.CodeAnalysis; using Microsoft.CodeAnalysis.CSharp; -using Microsoft.CodeAnalysis.CSharp.Syntax; using Microsoft.CodeAnalysis.Diagnostics; namespace FactoryGenerator @@ -30,9 +29,7 @@ public void Initialize(IncrementalGeneratorInitializationContext context) var rest = references.SelectMany(FindMethods); var attributes = rest.Collect(); var compilation = context.CompilationProvider; - var syntaxUsages = context.SyntaxProvider.CreateSyntaxProvider(ResolveSymbols, ResolveTransformations) - .Collect(); - var combined = attributes.Combine(compilation).Combine(syntaxUsages).Combine(logProvider); + var combined = attributes.Combine(compilation).Combine(logProvider); context.RegisterSourceOutput(combined, MakeAutofacModule); var supportsStaticExtensions = context.ParseOptionsProvider.Select(IsAtLeastCSharp14); @@ -61,15 +58,13 @@ public void Initialize(IncrementalGeneratorInitializationContext context) } private void MakeAutofacModule(SourceProductionContext context, - (((ImmutableArray Injections, Compilation Compilation) Left, ImmutableArray CompileTimeResolvedTypes) Left, LoggingOptions? log) - data) + ((ImmutableArray Injections, Compilation Compilation) Left, LoggingOptions? log) data) { - var injections = data.Left.Left.Injections; - var compilation = data.Left.Left.Compilation; - var usages = data.Left.CompileTimeResolvedTypes; + var injections = data.Left.Injections; + var compilation = data.Left.Compilation; var log = data.log?.FileName == null ? NullLogger.Instance : new Logger(data.log.FileName, data.log.LogLevel); - var source = GenerateCode(injections, compilation, usages, log).ToArray(); + var source = GenerateCode(injections, compilation, log).ToArray(); context.AddSource("DependencyInjectionContainer.Lookup.g.cs", source[0]); context.AddSource("DependencyInjectionContainer.Constructor.g.cs", source[1]); context.AddSource("DependencyInjectionContainer.Declarations.g.cs", source[2]); @@ -81,30 +76,6 @@ private void MakeAutofacModule(SourceProductionContext context, context.AddSource("ContainerEntryPoint.g.cs", source[8]); } - private UsageData? ResolveTransformations(GeneratorSyntaxContext context, CancellationToken token) - { - var typeArguments = context.Node.DescendantNodes().OfType().FirstOrDefault(); - if (typeArguments is null) return null; - var identifier = typeArguments.DescendantNodes().FirstOrDefault(); - if (identifier is null) return null; - var info = context.SemanticModel.GetSymbolInfo(identifier, token); - if (info.Symbol is not INamedTypeSymbol symbol) return null; - if (!SymbolUtility.IsEnumerable(symbol)) return null; - if (symbol.TypeArguments.Length != 1) return null; - var elemType = symbol.TypeArguments[0]; - return new UsageData( - fullName: symbol.ToString()!, - memberName: SymbolUtility.MemberName(symbol).Replace("()", ""), - elementTypeFullName: elemType.ToString()!, - elementTypeMemberName: SymbolUtility.MemberName(elemType).Replace("()", "")); - } - - private bool ResolveSymbols(SyntaxNode node, CancellationToken token) - { - if (node is not MemberAccessExpressionSyntax invocation) return false; - return invocation.ToString().Contains("Resolve"); - } - private static IEnumerable FindMethods(INamespaceSymbol namespaceSymbol, CancellationToken token) { foreach (var type in SymbolUtility.GetAllTypes(namespaceSymbol)) @@ -155,7 +126,7 @@ private static INamespaceSymbol GetGlobalNamespace(Compilation compilation, Canc private const string LifetimeName = "LifetimeScope"; private static IEnumerable GenerateCode(ImmutableArray dataInjections, - Compilation compilation, ImmutableArray usages, ILogger log) + Compilation compilation, ILogger log) { CheckForCycles(dataInjections); log.Log(LogLevel.Debug, "Starting Code Generation"); @@ -173,7 +144,7 @@ namespace {compilation.Assembly.Name}.Generated; [GeneratedCode(""{ToolName}"", ""{Version}"")] #nullable enable #pragma warning disable CS0169, CS0414 -public sealed partial class {ClassName} : IContainer +public sealed partial class {ClassName} : IContainer, IContainerScopeFactory, IContainerRegistrationMetadata, IContainerCacheInvalidator {{ #pragma warning restore CS0169, CS0414 @@ -195,6 +166,62 @@ private IContainer GetTop() }} return top; }} + private void AttachToBase(IContainer baseContainer) + {{ + if (baseContainer.Inheritor is null) + {{ + baseContainer.Inheritor = this; + InvalidateCollectionCachesInChain(); + return; + }} + + var current = baseContainer.Inheritor; + while (current!.Inheritor is not null) + {{ + current = current.Inheritor; + }} + + current.Inheritor = this; + InvalidateCollectionCachesInChain(); + }} + private void DetachFromBase() + {{ + if (Base is null) + return; + + if (Base.Inheritor == this) + {{ + Base.Inheritor = Inheritor; + Inheritor = null; + InvalidateCollectionCachesInChain(); + return; + }} + + var current = Base.Inheritor; + while (current is not null && current.Inheritor != this) + {{ + current = current.Inheritor; + }} + + if (current is null) + return; + + current.Inheritor = Inheritor; + Inheritor = null; + InvalidateCollectionCachesInChain(); + }} + private void InvalidateCollectionCachesInChain() + {{ + var current = GetRoot(); + while (current is not null) + {{ + if (current is IContainerCacheInvalidator invalidator) + invalidator.InvalidateCollectionCaches(); + + current = current.Inheritor; + }} + }} + public string AssemblyName => ""{compilation.Assembly.Name}""; public IContainer? Base {{ get; }} public IContainer? Inheritor {{ get; set; }} internal readonly object m_lock = new(); @@ -230,6 +257,7 @@ public object Resolve(Type type) public void Dispose() {{ + DetachFromBase(); if (resolvedInstances is not null) {{ foreach (var weakReference in resolvedInstances) @@ -241,7 +269,6 @@ public void Dispose() }} resolvedInstances.Clear(); }} - Base?.Dispose(); }} public bool TryResolve(Type type, out object? resolved) @@ -287,38 +314,10 @@ public bool GetBoolean(string key) }} }}"; - var booleans = dataInjections.Select(inj => inj.BooleanInjection).Where(b => b is not null) - .Select(b => b!.Key).Distinct().ToArray(); - var allArguments = booleans.Select(b => $"bool {b}").ToList(); - var justBooleans = allArguments.ToList(); - var allParameters = booleans.Select(b => $"{b}").ToList(); - var ordered = dataInjections.Reverse().ToList(); - - foreach (var injection in ordered.ToArray()) - { - log.Log(LogLevel.Debug, $"Traversing {injection.Name}"); - if (!injection.IsTestType) continue; - ordered.Remove(injection); - ordered.Add(injection); - } - - var interfaceInjectors = new Dictionary>(); - var interfaceMemberNames = new Dictionary(); - - foreach (var injection in ordered) - { - for (int i = 0; i < injection.InterfaceFullNames.Length; i++) - { - var ifaceFull = injection.InterfaceFullNames[i]; - var ifaceMember = injection.InterfaceMemberNames[i]; - if (!interfaceInjectors.ContainsKey(ifaceFull)) - { - interfaceInjectors[ifaceFull] = new List(); - interfaceMemberNames[ifaceFull] = ifaceMember; - } - interfaceInjectors[ifaceFull].Add(injection); - } - } + var booleanKeys = dataInjections.Select(inj => inj.BooleanInjection).Where(b => b is not null) + .Select(b => b!.Key).Distinct().ToArray(); + var ordered = OrderInjections(dataInjections, log); + var (interfaceInjectors, interfaceMemberNames) = BuildInterfaceInjectors(ordered); var declarations = new Dictionary(); var scopedDeclarations = new Dictionary(); @@ -339,6 +338,57 @@ public bool GetBoolean(string key) } } + var localizedParameters = new List(); + foreach (var parameter in constructorParameters.ToArray()) + { + if (!parameter.IsCollection) continue; + if (parameter.CollectionElementFullName is null) continue; + constructorParameters.Remove(parameter); + localizedParameters.Add(parameter); + } + + foreach (var parameter in constructorParameters.ToArray()) + { + if (!parameter.TypeFullName.Contains("IContainer")) continue; + log.Log(LogLevel.Debug, $"Registering {parameter.Name} as Self"); + declarations[parameter.Name] = $"private IContainer {parameter.Name} => this;"; + scopedDeclarations[parameter.Name] = $"private IContainer {parameter.Name} => this;"; + constructorParameters.Remove(parameter); + } + + ValidateExternalParameterTypes(constructorParameters); + + var booleanReservedNames = constructorParameters.Select(parameter => parameter.Name) + .Concat(localizedParameters.Select(parameter => "coll_" + parameter.CollectionElementMemberName!)) + .Concat(localizedParameters.Select(parameter => "m_coll_" + parameter.CollectionElementMemberName!)) + .Concat(interfaceMemberNames.Values) + .Concat(ordered.Select(injection => injection.Name.Replace("()", string.Empty))) + .Concat(ordered.Select(injection => injection.LazyFieldName)) + .Concat(new[] + { + "Base", + "Inheritor", + "GetRoot", + "GetTop", + "Dispose", + "Resolve", + "TryResolve", + "IsRegistered", + "GetBoolean", + "GetBooleans", + "BeginLifetimeScope", + "GetResolvedInstances", + "resolvedInstances", + "m_lock", + "m_lookup", + "m_booleans", + "m_fallback", + "fallback", + "baseContainer" + }); + var booleanIdentifiers = BuildBooleanParameterIdentifiers(booleanKeys, booleanReservedNames); + var booleanParameters = booleanKeys.Select(key => (Key: key, Identifier: booleanIdentifiers[key])).ToList(); + foreach (var ifaceFull in interfaceInjectors.Keys) { var possibilities = interfaceInjectors[ifaceFull]; @@ -366,12 +416,13 @@ public bool GetBoolean(string key) var ternary = new StringBuilder(); foreach (var key in keys) { + var keyIdentifier = booleanIdentifiers[key]; var trueValue = possibilities.LastOrDefault(p => p.BooleanInjection?.Value == true && p.BooleanInjection?.Key == key); trueValue ??= fallback; ternary.Append(key == last - ? $"{key} ? {trueValue?.Name ?? "null!"} : {fallback?.Name ?? "null!"}" - : $"{key} ? {trueValue?.Name ?? "null!"} : "); + ? $"{keyIdentifier} ? {trueValue?.Name ?? "null!"} : {fallback?.Name ?? "null!"}" + : $"{keyIdentifier} ? {trueValue?.Name ?? "null!"} : "); } if (!declarations.ContainsKey(ifaceMethodName)) @@ -383,83 +434,73 @@ public bool GetBoolean(string key) } } - var localizedParameters = new List(); var arrayDeclarations = new Dictionary(); - - foreach (var parameter in constructorParameters.ToArray()) + foreach (var pair in interfaceInjectors) { - if (!parameter.IsCollection) continue; - if (parameter.CollectionElementFullName is null) continue; - var name = "coll_" + parameter.CollectionElementMemberName!; - log.Log(LogLevel.Debug, $"Creating Collection: {name} of element type {parameter.CollectionElementFullName}"); - MakeArray(arrayDeclarations, name, parameter.CollectionElementFullName, parameter.CollectionElementMemberName!, interfaceInjectors); - constructorParameters.Remove(parameter); - localizedParameters.Add(parameter); + var name = "coll_" + interfaceMemberNames[pair.Key]; + if (arrayDeclarations.ContainsKey(name)) + continue; + log.Log(LogLevel.Debug, $"Creating Collection: {name} of element type {pair.Key}"); + MakeArray(arrayDeclarations, name, pair.Key, interfaceInjectors, booleanIdentifiers); } - var requestedUsages = new List(); - - foreach (var request in usages) + foreach (var parameter in localizedParameters) { - if (request is null) continue; - if (localizedParameters.Any(p => p.TypeFullName == request.FullName)) continue; - log.Log(LogLevel.Information, $"Creating Requested: {request.FullName}"); - log.Log(LogLevel.Debug, $"Creating Array: {request.MemberName} of type {request.ElementTypeFullName}[]"); - MakeArray(arrayDeclarations, request.MemberName, request.ElementTypeFullName, request.ElementTypeMemberName, interfaceInjectors, true); - requestedUsages.Add(request); + var name = "coll_" + parameter.CollectionElementMemberName!; + if (arrayDeclarations.ContainsKey(name)) + continue; + log.Log(LogLevel.Debug, $"Creating Collection: {name} of element type {parameter.CollectionElementFullName}"); + MakeArray(arrayDeclarations, name, parameter.CollectionElementFullName!, interfaceInjectors, booleanIdentifiers); } - foreach (var parameter in constructorParameters.ToArray()) - { - if (!parameter.TypeFullName.Contains("IContainer")) continue; - log.Log(LogLevel.Debug, $"Registering {parameter.Name} as Self"); - declarations[parameter.Name] = $"private IContainer {parameter.Name} => this;"; - scopedDeclarations[parameter.Name] = $"private IContainer {parameter.Name} => this;"; - constructorParameters.Remove(parameter); - } + var externalParameters = constructorParameters.OrderBy(parameter => parameter.TypeFullName).ToList(); + var allArguments = booleanParameters.Select(parameter => $"bool {parameter.Identifier}").ToList(); + allArguments.AddRange(externalParameters.Select(parameter => $"{parameter.TypeFullName} {parameter.Name}").Distinct()); - var arguments = constructorParameters.OrderBy(p => p.TypeFullName).Select(p => $"{p.TypeFullName} {p.Name}").Distinct(); - var parameters = constructorParameters.OrderBy(p => p.TypeFullName).Select(p => p.Name).Distinct(); - allArguments.AddRange(arguments); var lifetimeArguments = allArguments.ToList(); - allParameters.AddRange(parameters); + lifetimeArguments.Insert(0, "IContainer? baseContainer"); + lifetimeArguments.Insert(0, $"{ClassName} fallback"); + var lifetimeParameters = new List { "this", "baseContainer" }; + lifetimeParameters.AddRange(booleanParameters.Select(parameter => parameter.Identifier)); + lifetimeParameters.AddRange(externalParameters.Select(parameter => $"baseContainer != null ? baseContainer.Resolve<{parameter.TypeFullName}>() : {parameter.Name}")); var constructor = "(" + string.Join(", ", allArguments) + ")"; - lifetimeArguments.Insert(0, $"{ClassName} fallback"); - allParameters.Insert(0, "this"); var lifetimeConstructor = "(" + string.Join(", ", lifetimeArguments) + ")"; - var lifetimeParameters = string.Join(", ", allParameters); + var lifetimeParameterValues = string.Join(", ", lifetimeParameters); log.Log(LogLevel.Debug, $"Resulting Constructor: {constructor}"); var constructorFields = string.Join("\n\t", allArguments.Select(arg => "internal " + arg + ";")); var constructorAssignments = string.Join("\n\t\t", allArguments.Select(arg => arg.Split(' ').Last()).Select(arg => $"this.{arg} = {arg};")); var resolvedConstructorAssignments = string.Join("\n\t\t", - allArguments.Select(a => a.Split(' ')).Where(a => a[0] != "bool") - .Select(a => $"this.{a[1]} = Base.Resolve<{a[0]}>();")); + externalParameters.Select(parameter => $"this.{parameter.Name} = Base.Resolve<{parameter.TypeFullName}>();")); var interfacePairs = interfaceInjectors.Keys.Select(k => (TypeName: k, MemberName: interfaceMemberNames[k])).ToList(); // ReadOnlySpan is a ref struct and cannot be placed in the lookup dictionary var localizedForDict = localizedParameters.Where(p => p.CollectionKind != CollectionKind.ReadOnlySpan).ToList(); - var localizedPairs = localizedForDict + var localizedPairs = DistinctByTypeName(localizedForDict .Select(p => (TypeName: p.TypeFullName, Expression: CollectionDictExpression(p.CollectionKind, "coll_" + p.CollectionElementMemberName!))) + .ToList(), pair => pair.TypeName); + var localizedTypes = new HashSet(localizedPairs.Select(pair => pair.TypeName)); + var enumerablePairs = interfaceInjectors.Keys + .Select(key => (TypeName: $"System.Collections.Generic.IEnumerable<{key}>", Expression: "coll_" + interfaceMemberNames[key])) + .Where(pair => !localizedTypes.Contains(pair.TypeName)) .ToList(); - var requestedPairs = requestedUsages.Select(u => (TypeName: u.FullName, MemberName: u.MemberName)).ToList(); - var constructorPairs = constructorParameters.Select(p => (TypeName: p.TypeFullName, Expression: p.Name)).ToList(); + var constructorPairs = DistinctByTypeName(externalParameters.Select(p => (TypeName: p.TypeFullName, Expression: p.Name)).ToList(), pair => pair.TypeName); - var dictSize = interfaceInjectors.Count + localizedForDict.Count + requestedUsages.Count + constructorParameters.Count; + var dictSize = interfacePairs.Count + localizedPairs.Count + enumerablePairs.Count + constructorPairs.Count; yield return Constructor(usingStatements, constructorFields, constructor, constructorAssignments, - dictSize, interfacePairs, localizedPairs, requestedPairs, constructorPairs, - true, ClassName, lifetimeParameters, - resolvingConstructorAssignments: resolvedConstructorAssignments, booleans: justBooleans); + dictSize, interfacePairs, localizedPairs, enumerablePairs, constructorPairs, + true, ClassName, lifetimeInvocationValues: lifetimeParameterValues, + resolvingConstructorAssignments: resolvedConstructorAssignments, booleans: booleanParameters); yield return Declarations(usingStatements, declarations, ClassName); yield return ArrayDeclarations(usingStatements, arrayDeclarations, ClassName); yield return $@"{usingStatements} [GeneratedCode(""{ToolName}"", ""{Version}"")] #nullable enable #pragma warning disable CS0169, CS0414 -public sealed partial class LifetimeScope : IContainer +public sealed partial class LifetimeScope : IContainer, IContainerScopeFactory, IContainerRegistrationMetadata, IContainerCacheInvalidator {{ #pragma warning restore CS0169, CS0414 private IContainer GetRoot() @@ -480,17 +521,79 @@ private IContainer GetTop() }} return top; }} + private void AttachToBase(IContainer baseContainer) + {{ + if (baseContainer.Inheritor is null) + {{ + baseContainer.Inheritor = this; + InvalidateCollectionCachesInChain(); + return; + }} + + var current = baseContainer.Inheritor; + while (current!.Inheritor is not null) + {{ + current = current.Inheritor; + }} + + current.Inheritor = this; + InvalidateCollectionCachesInChain(); + }} + private void DetachFromBase() + {{ + if (Base is null) + return; + + if (Base.Inheritor == this) + {{ + Base.Inheritor = Inheritor; + Inheritor = null; + InvalidateCollectionCachesInChain(); + return; + }} + + var current = Base.Inheritor; + while (current is not null && current.Inheritor != this) + {{ + current = current.Inheritor; + }} + + if (current is null) + return; + + current.Inheritor = Inheritor; + Inheritor = null; + InvalidateCollectionCachesInChain(); + }} + private void InvalidateCollectionCachesInChain() + {{ + var current = GetRoot(); + while (current is not null) + {{ + if (current is IContainerCacheInvalidator invalidator) + invalidator.InvalidateCollectionCaches(); + + current = current.Inheritor; + }} + }} + public string AssemblyName => ""{compilation.Assembly.Name}""; public IContainer? Base {{ get; }} public IContainer? Inheritor {{ get; set; }} public ILifetimeScope BeginLifetimeScope() {{ - var scope = m_fallback.BeginLifetimeScope(); + var baseContainer = Base?.BeginLifetimeScope() as IContainer; + return BeginLifetimeScope(baseContainer); + }} + public ILifetimeScope BeginLifetimeScope(IContainer? baseContainer) + {{ + var scope = m_fallback.BeginLifetimeScope(baseContainer); GetResolvedInstances().Add(new WeakReference(scope)); return scope; }} internal readonly object m_lock = new(); private {ClassName} m_fallback; + private IContainer? m_ownedBase; private Dictionary> m_lookup; private Dictionary m_booleans; private List>? resolvedInstances; @@ -523,6 +626,7 @@ public object Resolve(Type type) public void Dispose() {{ + DetachFromBase(); if (resolvedInstances is not null) {{ foreach (var weakReference in resolvedInstances) @@ -534,7 +638,9 @@ public void Dispose() }} resolvedInstances.Clear(); }} - Base?.Dispose(); + var ownedBase = m_ownedBase; + m_ownedBase = null; + ownedBase?.Dispose(); }} public bool TryResolve(Type type, out object? resolved) @@ -583,9 +689,9 @@ public bool GetBoolean(string key) "; yield return Constructor(usingStatements, constructorFields, lifetimeConstructor, constructorAssignments, - dictSize, interfacePairs, localizedPairs, requestedPairs, constructorPairs, + dictSize, interfacePairs, localizedPairs, enumerablePairs, constructorPairs, false, LifetimeName, - resolvingConstructorAssignments: resolvedConstructorAssignments, addMergingConstructor: false, booleans: justBooleans); + resolvingConstructorAssignments: resolvedConstructorAssignments, addMergingConstructor: false, booleans: booleanParameters); yield return Declarations(usingStatements, scopedDeclarations, LifetimeName); yield return ArrayDeclarations(usingStatements, arrayDeclarations, LifetimeName); @@ -633,58 +739,226 @@ internal static void Register() "; } - private static void CheckForCycles(ImmutableArray dataInjections) + private static List OrderInjections(ImmutableArray dataInjections, ILogger? log = null) { - // Build adjacency list: interface/type name → set of dependency names - var graph = new Dictionary>(); - // Map each interface name back to its concrete type for error messages - var nodeOwner = new Dictionary(); + var ordered = dataInjections.Reverse().ToList(); - foreach (var injection in dataInjections) + foreach (var injection in ordered.ToArray()) { - if (injection.Lambda != null) continue; + log?.Log(LogLevel.Debug, $"Traversing {injection.Name}"); + if (!injection.IsTestType) continue; + ordered.Remove(injection); + ordered.Add(injection); + } - var deps = new HashSet(); - foreach (var ctor in injection.Constructors) + return ordered; + } + + private static (Dictionary> InterfaceInjectors, Dictionary InterfaceMemberNames) BuildInterfaceInjectors(IEnumerable ordered) + { + var interfaceInjectors = new Dictionary>(); + var interfaceMemberNames = new Dictionary(); + + foreach (var injection in ordered) + { + for (var i = 0; i < injection.InterfaceFullNames.Length; i++) { - foreach (var parameter in ctor.Parameters) + var ifaceFull = injection.InterfaceFullNames[i]; + var ifaceMember = injection.InterfaceMemberNames[i]; + if (!interfaceInjectors.ContainsKey(ifaceFull)) { - string? depName; - if (parameter.IsCollection) - { - if (parameter.CollectionElementFullName is null) continue; - depName = parameter.CollectionElementFullName; - } - else - { - // Strip ? so nullable params resolve to their underlying type in the cycle graph - depName = parameter.IsNullable - ? parameter.TypeFullName.TrimEnd('?') - : parameter.TypeFullName; - } - - deps.Add(depName); + interfaceInjectors[ifaceFull] = new List(); + interfaceMemberNames[ifaceFull] = ifaceMember; } + + interfaceInjectors[ifaceFull].Add(injection); } + } + + return (interfaceInjectors, interfaceMemberNames); + } + + private static List GetReachableImplementations(List possibilities) + { + if (possibilities.Count == 0) + return new List(); - foreach (var ifaceName in injection.InterfaceFullNames) + if (possibilities.All(i => i.BooleanInjection == null)) + return new List { possibilities.Last() }; + + var reachable = new List(); + var fallback = possibilities.LastOrDefault(p => p.BooleanInjection == null); + var keys = possibilities.Select(p => p.BooleanInjection?.Key).OfType().Distinct(); + + foreach (var key in keys) + { + var selected = possibilities.LastOrDefault(p => p.BooleanInjection?.Value == true && p.BooleanInjection?.Key == key) ?? fallback; + if (selected is not null && !reachable.Contains(selected)) + reachable.Add(selected); + } + + if (fallback is not null && !reachable.Contains(fallback)) + reachable.Add(fallback); + + return reachable; + } + + private static IEnumerable GetCycleDependencies(InjectionData injection, ImmutableArray availableInterfaceFullNames) + { + if (injection.Lambda is not null) + yield break; + + HashSet? missing = null; + HashSet? nullableDefaults = null; + var ctor = GetBestConstructor(injection, availableInterfaceFullNames, ref missing, ref nullableDefaults); + if (ctor is null) + yield break; + + foreach (var parameter in ctor.Parameters) + { + if (parameter.IsCollection) + continue; + + var typeLookup = parameter.IsNullable + ? parameter.TypeFullName.TrimEnd('?') + : parameter.TypeFullName; + if (!availableInterfaceFullNames.Contains(typeLookup)) + continue; + + yield return typeLookup; + } + } + + private static void ValidateExternalParameterTypes(List constructorParameters) + { + var ambiguousParameters = constructorParameters + .GroupBy(parameter => parameter.TypeFullName) + .Select(group => new { - if (!graph.ContainsKey(ifaceName)) - { - graph[ifaceName] = deps; - nodeOwner[ifaceName] = injection.TypeFullName; - } - else - { - // Multiple implementations of the same interface — merge edges - foreach (var d in deps) - graph[ifaceName].Add(d); - } + TypeFullName = group.Key, + Names = group.Select(parameter => parameter.Name).Distinct().OrderBy(name => name).ToArray() + }) + .Where(group => group.Names.Length > 1) + .ToList(); + + if (ambiguousParameters.Count == 0) + return; + + var details = string.Join("; ", ambiguousParameters.Select(group => $"{group.TypeFullName} ({string.Join(", ", group.Names)})")); + throw new InvalidOperationException( + $"Multiple externally provided values of the same type are not supported because FactoryGenerator resolves external values by type. Conflicting parameters: {details}. Wrap the values in distinct types or inject a dedicated options object."); + } + + private static List DistinctByTypeName(IEnumerable values, Func typeNameSelector) + { + var distinct = new List(); + var seenTypes = new HashSet(); + + foreach (var value in values) + { + if (!seenTypes.Add(typeNameSelector(value))) + continue; + + distinct.Add(value); + } + + return distinct; + } + + private static Dictionary BuildBooleanParameterIdentifiers(IEnumerable booleanKeys, IEnumerable reservedNames) + { + var identifiers = new Dictionary(); + var usedNames = new HashSet(reservedNames); + + foreach (var booleanKey in booleanKeys) + { + var candidate = GetBooleanParameterIdentifier(booleanKey); + var suffix = 1; + while (!usedNames.Add(candidate)) + { + candidate = $"{candidate}_{suffix}"; + suffix++; + } + + identifiers[booleanKey] = candidate; + } + + return identifiers; + } + + private static Dictionary BuildExternalParameterIdentifiers(IEnumerable parameters, IEnumerable reservedNames) + { + var identifiers = new Dictionary(); + var usedNames = new HashSet(reservedNames); + + foreach (var parameter in parameters.OrderBy(parameter => parameter.TypeFullName, StringComparer.Ordinal)) + { + if (identifiers.ContainsKey(parameter.TypeFullName)) + continue; + + var candidate = GetExternalParameterIdentifier(parameter.Name); + var suffix = 1; + while (!usedNames.Add(candidate)) + { + candidate = $"{candidate}_{suffix}"; + suffix++; + } + + identifiers[parameter.TypeFullName] = candidate; + } + + return identifiers; + } + + private static string GetBooleanParameterIdentifier(string booleanKey) + { + return GetSanitizedIdentifier(booleanKey, "boolean_"); + } + + private static string GetExternalParameterIdentifier(string parameterName) + { + return GetSanitizedIdentifier(parameterName, "argument_"); + } + + private static string GetSanitizedIdentifier(string value, string prefix) + { + if (SyntaxFacts.IsValidIdentifier(value)) + return value; + + var builder = new StringBuilder(prefix); + foreach (var character in value) + builder.Append(char.IsLetterOrDigit(character) ? character : '_'); + + var candidate = builder.ToString().TrimEnd('_'); + return SyntaxFacts.IsValidIdentifier(candidate) ? candidate : prefix.TrimEnd('_'); + } + + private static void CheckForCycles(ImmutableArray dataInjections) + { + var ordered = OrderInjections(dataInjections); + var (interfaceInjectors, _) = BuildInterfaceInjectors(ordered); + var availableInterfaceFullNames = interfaceInjectors.Keys.ToImmutableArray(); + + var graph = new Dictionary>(); + var nodeOwner = new Dictionary(); + + foreach (var interfaceInjector in interfaceInjectors) + { + var ifaceName = interfaceInjector.Key; + var possibilities = interfaceInjector.Value; + var reachable = GetReachableImplementations(possibilities); + var deps = new HashSet(); + + foreach (var injection in reachable) + { + foreach (var dep in GetCycleDependencies(injection, availableInterfaceFullNames)) + deps.Add(dep); } + + graph[ifaceName] = deps; + nodeOwner[ifaceName] = string.Join(", ", reachable.Select(injection => injection.TypeFullName).Distinct()); } - // DFS-based cycle detection - // 0 = unvisited, 1 = in-progress (on current path), 2 = done var state = new Dictionary(); var path = new List(); @@ -731,15 +1005,21 @@ private static void DfsCycleCheck(string node, Dictionary interfaceTypePairs, IEnumerable<(string TypeName, string Expression)> localizedParamPairs, - IEnumerable<(string TypeName, string MemberName)> requestedPairs, IEnumerable<(string TypeName, string Expression)> constructorParamPairs, - bool addLifetimeScopeFunction, string className, string? lifetimeParameters = null, - string? fromConstructor = null, string? resolvingConstructorAssignments = null, bool addMergingConstructor = true, List booleans = null!) + IEnumerable<(string TypeName, string Expression)> enumerablePairs, IEnumerable<(string TypeName, string Expression)> constructorParamPairs, + bool addLifetimeScopeFunction, string className, string? lifetimeInvocationValues = null, + string? fromConstructor = null, string? resolvingConstructorAssignments = null, bool addMergingConstructor = true, + IReadOnlyList<(string Key, string Identifier)> booleans = null!) { var lifetimeScopeFunction = addLifetimeScopeFunction ? $@" public ILifetimeScope BeginLifetimeScope() {{ - var scope = new {LifetimeName}({lifetimeParameters}); + var baseContainer = Base?.BeginLifetimeScope() as IContainer; + return BeginLifetimeScope(baseContainer); +}} +public ILifetimeScope BeginLifetimeScope(IContainer? baseContainer) +{{ + var scope = new {LifetimeName}({lifetimeInvocationValues}); GetResolvedInstances().Add(new WeakReference(scope)); return scope; }}" : string.Empty; @@ -748,15 +1028,15 @@ public ILifetimeScope BeginLifetimeScope() public {className}(IContainer Base{fromConstructor}) {{ this.Base = Base; - Base.Inheritor = this; + AttachToBase(Base); {resolvingConstructorAssignments} -{string.Join("\n", booleans.Select(b => b.Split(' ').Last()).Select(b => $"\t this.{b} = Base.GetBoolean(\"{b}\");"))} +{string.Join("\n", booleans.Select(boolean => $"\t this.{boolean.Identifier} = Base.GetBoolean(\"{boolean.Key}\");"))} m_lookup = new({dictSize}) {{ {MakeDictionaryFromTypes(interfaceTypePairs)} {MakeDictionaryFromParams(localizedParamPairs)} -{MakeDictionaryFromTypes(requestedPairs)} +{MakeDictionaryFromParams(enumerablePairs)} {MakeDictionaryFromParams(constructorParamPairs)} }}; m_booleans = new(); @@ -767,7 +1047,11 @@ public ILifetimeScope BeginLifetimeScope() }}" : string.Empty; - var extraConstruction = addLifetimeScopeFunction ? string.Empty : "m_fallback = fallback;"; + var extraConstruction = addLifetimeScopeFunction ? string.Empty : @"m_fallback = fallback; + this.Base = baseContainer; + m_ownedBase = baseContainer; + if (baseContainer is not null) + AttachToBase(baseContainer);"; return $@"{usingStatements} public partial class {className} {{ @@ -780,12 +1064,12 @@ public partial class {className} m_lookup = new({dictSize}) {{ {MakeDictionaryFromTypes(interfaceTypePairs)} {MakeDictionaryFromParams(localizedParamPairs)} -{MakeDictionaryFromTypes(requestedPairs)} +{MakeDictionaryFromParams(enumerablePairs)} {MakeDictionaryFromParams(constructorParamPairs)} }}; m_booleans = new({booleans.Count}) {{ -{string.Join("\n", booleans.Select(b => b.Split(' ').Last()).Select(b => $"\t\t{{ \"{b}\", {b} }},"))} +{string.Join("\n", booleans.Select(boolean => $"\t\t{{ \"{boolean.Key}\", {boolean.Identifier} }},"))} }}; }} {mergingConstructor} @@ -796,12 +1080,19 @@ public partial class {className} private static string ArrayDeclarations(string usingStatements, Dictionary arrayDeclarations, string className) { + var cacheInvalidations = string.Join("\n ", arrayDeclarations.Keys.Select(name => $"m_{name} = null;")); + return $@"{usingStatements} public partial class {className} {{ {string.Join("\n\t", arrayDeclarations.Values)} -}}"; - } + + public void InvalidateCollectionCaches() + {{ + {cacheInvalidations} + }} +}}"; + } private static string Declarations(string usingStatements, Dictionary declarations, string className) { @@ -813,14 +1104,11 @@ public partial class {className} } private static void MakeArray(Dictionary declarations, string name, - string elementTypeFullName, string elementTypeMemberName, - Dictionary> interfaceInjectors, bool function = false) + string elementTypeFullName, Dictionary> interfaceInjectors, + IReadOnlyDictionary booleanIdentifiers) { var factoryName = $"new {elementTypeFullName}[0]"; var factory = string.Empty; - var functionString = function ? "()" : string.Empty; - var starter = function ? string.Empty : "get {"; - var ender = function ? string.Empty : "}"; if (interfaceInjectors.TryGetValue(elementTypeFullName, out var injections)) { factoryName = $"Create{name}()".Replace("_", ""); @@ -835,7 +1123,7 @@ private static void MakeArray(Dictionary declarations, string na var source = new List<{elementTypeFullName}>({nonBooleanInjections.Count}) {{ {string.Join(",\n\t\t\t", nonBooleanInjections.Select(i => i.Name))} }}; - {string.Join("\n\t\t\t", booleanInjections.Select(i => $"if({i.BooleanInjection!.Key}) source.Add({i.Name});"))} + {string.Join("\n\t\t\t", booleanInjections.Select(i => $"if({booleanIdentifiers[i.BooleanInjection!.Key]}) source.Add({i.Name});"))} var b = Base; while(b is not null) {{ @@ -853,21 +1141,22 @@ private static void MakeArray(Dictionary declarations, string na }}"; } declarations[name] = $@" - internal IEnumerable<{elementTypeFullName}> {name}{functionString} + internal IEnumerable<{elementTypeFullName}> {name} {{ - {starter} - var cached = m_{name}; - if (cached != null) - return cached; - - lock (m_lock) + get {{ - cached = m_{name}; + var cached = m_{name}; if (cached != null) return cached; - return m_{name} = {factoryName}; + + lock (m_lock) + {{ + cached = m_{name}; + if (cached != null) + return cached; + return m_{name} = {factoryName}; + }} }} - {ender} }} internal IEnumerable<{elementTypeFullName}>? m_{name};" + factory; } @@ -1097,6 +1386,37 @@ private static bool GetEmitStaticExtensions(AnalyzerConfigOptionsProvider provid return !string.Equals(value, "false", StringComparison.OrdinalIgnoreCase); } + private sealed class StaticExtensionSpec + { + public StaticExtensionSpec(string typeFullName, string typeMemberName, string extensionClassName, List possibilities) + { + TypeFullName = typeFullName; + TypeMemberName = typeMemberName; + ExtensionClassName = extensionClassName; + Possibilities = possibilities; + } + + public string TypeFullName { get; } + public string TypeMemberName { get; } + public string ExtensionClassName { get; } + public List Possibilities { get; } + public List BooleanKeys { get; } = new List(); + public List ExternalParameters { get; } = new List(); + public List Dependencies { get; } = new List(); + } + + private sealed class StaticDependencyReference + { + public StaticDependencyReference(string typeFullName, bool resolveAll) + { + TypeFullName = typeFullName; + ResolveAll = resolveAll; + } + + public string TypeFullName { get; } + public bool ResolveAll { get; } + } + private static void MakeStaticExtensions( SourceProductionContext context, ((ImmutableArray Injections, Compilation Compilation) Left, bool SupportsExtensions) data) @@ -1109,53 +1429,53 @@ private static void MakeStaticExtensions( private static string GenerateStaticExtensions( ImmutableArray dataInjections, Compilation compilation) { - var ordered = dataInjections.Reverse().ToList(); - foreach (var injection in ordered.ToArray()) - { - if (!injection.IsTestType) continue; - ordered.Remove(injection); - ordered.Add(injection); - } - - var interfaceInjectors = new Dictionary>(); - var interfaceMemberNames = new Dictionary(); - foreach (var injection in ordered) - { - for (var i = 0; i < injection.InterfaceFullNames.Length; i++) - { - var ifaceFull = injection.InterfaceFullNames[i]; - var ifaceMember = injection.InterfaceMemberNames[i]; - if (!interfaceInjectors.ContainsKey(ifaceFull)) - { - interfaceInjectors[ifaceFull] = new List(); - interfaceMemberNames[ifaceFull] = ifaceMember; - } - interfaceInjectors[ifaceFull].Add(injection); - } - } - + var ordered = OrderInjections(dataInjections); + var (interfaceInjectors, interfaceMemberNames) = BuildInterfaceInjectors(ordered); var availableInterfaces = interfaceInjectors.Keys.ToImmutableArray(); + var specs = BuildStaticExtensionSpecs(interfaceInjectors, interfaceMemberNames, availableInterfaces); + + var reservedNames = BuildStaticExtensionReservedNames(specs, interfaceMemberNames); + var externalIdentifiers = BuildExternalParameterIdentifiers(specs.Values.SelectMany(spec => spec.ExternalParameters), reservedNames); + var booleanIdentifiers = BuildBooleanParameterIdentifiers( + specs.Values.SelectMany(spec => spec.BooleanKeys).Distinct(), + reservedNames.Concat(externalIdentifiers.Values)); var sb = new StringBuilder(); - sb.AppendLine($@"using System.CodeDom.Compiler; + sb.AppendLine($@"using System; +using System.CodeDom.Compiler; +using System.Collections.Generic; +using System.Collections.Immutable; using System.Linq; namespace {compilation.Assembly.Name}.Generated; -#nullable enable"); +#nullable enable + +internal sealed class StaticResolveState +{{ + private readonly HashSet m_activeCollections = new(StringComparer.Ordinal); + + public bool EnterCollection(string key) + {{ + return m_activeCollections.Add(key); + }} + + public void ExitCollection(string key) + {{ + m_activeCollections.Remove(key); + }} +}}"); - foreach (var kvp in interfaceMemberNames) + foreach (var ifaceFull in interfaceMemberNames.Keys) { - var ifaceFull = kvp.Key; - var ifaceMember = kvp.Value; - var possibilities = interfaceInjectors[ifaceFull]; - var className = ifaceMember + "Extensions"; - var body = StaticExtensionBody(ifaceFull, ifaceMember, possibilities, availableInterfaces); + var spec = specs[ifaceFull]; + var helpers = BuildStaticExtensionClass(spec, specs, availableInterfaces, booleanIdentifiers, externalIdentifiers); sb.AppendLine($@"[GeneratedCode(""{ToolName}"", ""{Version}"")] -public static class {className} +public static class {spec.ExtensionClassName} {{ - extension({ifaceFull}) +{helpers} + extension({spec.TypeFullName}) {{ -{body} +{BuildStaticPublicResolveMethods(spec, booleanIdentifiers, externalIdentifiers)} }} }}"); } @@ -1163,209 +1483,641 @@ public static class {className} return sb.ToString(); } - /// - /// Produces the method declaration(s) that go inside an extension(...) block - /// for the given interface. - /// - private static string StaticExtensionBody( - string ifaceFull, - string ifaceMember, - List possibilities, + private static Dictionary BuildStaticExtensionSpecs( + Dictionary> interfaceInjectors, + Dictionary interfaceMemberNames, ImmutableArray availableInterfaces) { - var hasBooleans = possibilities.Any(p => p.BooleanInjection != null); - var chosen = possibilities.Last(); + var specs = interfaceInjectors.ToDictionary( + pair => pair.Key, + pair => CreateDirectStaticExtensionSpec(pair.Key, pair.Value, interfaceMemberNames[pair.Key], availableInterfaces, interfaceInjectors), + StringComparer.Ordinal); + + PropagateStaticExtensionRequirements(specs); - // Boolean-switched or lambda: the container already encodes all of that logic in its - // own factory method — call it directly (no dictionary) and fall back to inline - // construction when the container is absent. - if (hasBooleans || chosen.Lambda != null) + foreach (var spec in specs.Values) { - string nullFallback; - if (chosen.Lambda != null) + var orderedBooleanKeys = spec.BooleanKeys.OrderBy(key => key, StringComparer.Ordinal).ToArray(); + spec.BooleanKeys.Clear(); + spec.BooleanKeys.AddRange(orderedBooleanKeys); + + var orderedExternalParameters = spec.ExternalParameters + .OrderBy(parameter => parameter.TypeFullName, StringComparer.Ordinal) + .ToArray(); + spec.ExternalParameters.Clear(); + spec.ExternalParameters.AddRange(orderedExternalParameters); + } + + return specs; + } + + private static StaticExtensionSpec CreateDirectStaticExtensionSpec( + string typeFullName, + List possibilities, + string typeMemberName, + ImmutableArray availableInterfaces, + Dictionary> interfaceInjectors) + { + var spec = new StaticExtensionSpec(typeFullName, typeMemberName, typeMemberName + "Extensions", possibilities); + PopulateStaticExtensionDirectRequirements(spec, availableInterfaces, interfaceInjectors); + return spec; + } + + private static void PopulateStaticExtensionDirectRequirements( + StaticExtensionSpec spec, + ImmutableArray availableInterfaces, + Dictionary> interfaceInjectors) + { + foreach (var booleanKey in spec.Possibilities.Select(possibility => possibility.BooleanInjection?.Key).OfType()) + AddDistinctString(spec.BooleanKeys, booleanKey); + + foreach (var possibility in spec.Possibilities) + { + if (possibility.Lambda is LambdaData lambda) { - nullFallback = - $"throw new global::System.InvalidOperationException(" + - $"\"Cannot resolve {ifaceFull} without a container\")"; + if (interfaceInjectors.ContainsKey(lambda.ContainingTypeFullName)) + AddDistinctDependency(spec.Dependencies, new StaticDependencyReference(lambda.ContainingTypeFullName, false)); + + foreach (var parameter in lambda.MethodParameters) + AddStaticParameterRequirement(spec, parameter, interfaceInjectors); + + continue; } - else + + HashSet? missing = null; + HashSet? nullableDefaults = null; + var constructor = GetBestConstructor(possibility, availableInterfaces, ref missing, ref nullableDefaults); + if (constructor is null) + continue; + + foreach (var parameter in constructor.Parameters) + AddStaticParameterRequirement(spec, parameter, interfaceInjectors); + } + } + + private static void AddStaticParameterRequirement( + StaticExtensionSpec spec, + ParameterData parameter, + Dictionary> interfaceInjectors) + { + if (parameter.IsCollection) + { + if (parameter.CollectionElementFullName is not null && interfaceInjectors.ContainsKey(parameter.CollectionElementFullName)) + AddDistinctDependency(spec.Dependencies, new StaticDependencyReference(parameter.CollectionElementFullName, true)); + return; + } + + var typeLookup = parameter.IsNullable + ? parameter.TypeFullName.TrimEnd('?') + : parameter.TypeFullName; + + if (interfaceInjectors.ContainsKey(typeLookup)) + { + AddDistinctDependency(spec.Dependencies, new StaticDependencyReference(typeLookup, false)); + return; + } + + if (parameter.HasExplicitDefault || parameter.IsParams || parameter.IsNullable) + return; + + AddDistinctParameter(spec.ExternalParameters, parameter); + } + + private static void PropagateStaticExtensionRequirements(IReadOnlyDictionary specs) + { + bool changed; + do + { + changed = false; + + foreach (var spec in specs.Values) { - var fallbackImpl = possibilities.LastOrDefault(p => p.BooleanInjection == null); - nullFallback = fallbackImpl != null - ? InlineCreation(fallbackImpl, availableInterfaces, nullContainer: true) - : $"throw new global::System.InvalidOperationException(" + - $"\"Cannot resolve {ifaceFull} without a container\")"; + foreach (var dependency in spec.Dependencies) + { + if (!specs.TryGetValue(dependency.TypeFullName, out var dependencySpec)) + continue; + + foreach (var booleanKey in dependencySpec.BooleanKeys) + changed |= AddDistinctString(spec.BooleanKeys, booleanKey); + + foreach (var externalParameter in dependencySpec.ExternalParameters) + changed |= AddDistinctParameter(spec.ExternalParameters, externalParameter); + } } - return $" public static {ifaceFull} Resolve({ClassName}? container) =>" + - $" container?.{ifaceMember}() ?? {nullFallback};"; + } while (changed); + } + + private static bool AddDistinctString(List values, string value) + { + if (values.Contains(value)) + return false; + + values.Add(value); + return true; + } + + private static bool AddDistinctParameter(List values, ParameterData value) + { + if (values.Any(parameter => parameter.TypeFullName == value.TypeFullName)) + return false; + + values.Add(value); + return true; + } + + private static bool AddDistinctDependency(List values, StaticDependencyReference value) + { + if (values.Any(existing => existing.TypeFullName == value.TypeFullName && existing.ResolveAll == value.ResolveAll)) + return false; + + values.Add(value); + return true; + } + + private static IEnumerable BuildStaticExtensionReservedNames( + IReadOnlyDictionary specs, + IReadOnlyDictionary interfaceMemberNames) + { + return interfaceMemberNames.Values + .Concat(specs.Values.Select(spec => spec.ExtensionClassName)) + .Concat(specs.Values.SelectMany(spec => spec.Possibilities.Select(GetStaticInjectionHelperName))) + .Concat(new[] + { + "container", + "state", + "cached", + "value", + "source", + "additional", + "disposable", + "key", + "b", + "Resolve", + "ResolveCore", + "ResolveAllCore", + "StaticResolveState", + "EnterCollection", + "ExitCollection" + }); + } + + private static string BuildStaticExtensionClass( + StaticExtensionSpec spec, + IReadOnlyDictionary specs, + ImmutableArray availableInterfaces, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + var parts = new List + { + BuildStaticResolveCoreMethod(spec, booleanIdentifiers, externalIdentifiers), + BuildStaticResolveAllCoreMethod(spec, booleanIdentifiers, externalIdentifiers) + }; + + foreach (var possibility in spec.Possibilities) + { + parts.Add(BuildStaticResolveInjectionMethod(spec, possibility, booleanIdentifiers, externalIdentifiers)); + parts.Add(BuildStaticCreateInjectionMethod(spec, possibility, specs, availableInterfaces, booleanIdentifiers, externalIdentifiers)); } - var creation = InlineCreation(chosen, availableInterfaces, nullContainer: false); - var nullCreation = InlineCreation(chosen, availableInterfaces, nullContainer: true); + return string.Join("\n\n", parts); + } - // Singleton / Scoped: inline double-checked locking directly against the container's - // cache field, bypassing both the dictionary and the factory-method call. - if (chosen.Singleton || chosen.Scoped) + private static string BuildStaticPublicResolveMethods( + StaticExtensionSpec spec, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + var runtimeDeclarations = BuildStaticRuntimeArgumentDeclarations(spec, booleanIdentifiers, externalIdentifiers); + var runtimeArguments = BuildStaticRuntimeArgumentValues(spec, booleanIdentifiers, externalIdentifiers); + var containerSignature = string.IsNullOrEmpty(runtimeDeclarations) + ? $"{ClassName}? container" + : $"{ClassName}? container, {runtimeDeclarations}"; + var containerInvocation = string.IsNullOrEmpty(runtimeArguments) + ? "container, new StaticResolveState()" + : $"container, new StaticResolveState(), {runtimeArguments}"; + + var methods = new List { - if (chosen.Disposable) - { - return $@" public static {ifaceFull} Resolve({ClassName}? container) + $@" public static {spec.TypeFullName} Resolve({containerSignature}) + {{ + return ResolveCore({containerInvocation}); + }}" + }; + + if (!string.IsNullOrEmpty(runtimeDeclarations)) + { + var nullInvocation = $"null, new StaticResolveState(), {runtimeArguments}"; + methods.Add($@" public static {spec.TypeFullName} Resolve({runtimeDeclarations}) + {{ + return ResolveCore({nullInvocation}); + }}"); + } + + return string.Join("\n\n", methods); + } + + private static string BuildStaticResolveCoreMethod( + StaticExtensionSpec spec, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + var parameterList = BuildStaticMethodParameterList(spec, booleanIdentifiers, externalIdentifiers, includeContainer: true, includeState: true); + var resolveExpression = BuildStaticResolveSelectionExpression(spec, booleanIdentifiers, externalIdentifiers); + + return $@" internal static {spec.TypeFullName} ResolveCore({parameterList}) + {{ + return {resolveExpression}; + }}"; + } + + private static string BuildStaticResolveAllCoreMethod( + StaticExtensionSpec spec, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + var parameterList = BuildStaticMethodParameterList(spec, booleanIdentifiers, externalIdentifiers, includeContainer: true, includeState: true); + var resolveInvocations = spec.Possibilities + .Where(possibility => possibility.BooleanInjection == null) + .Select(possibility => BuildStaticResolveInjectionInvocation(spec, possibility, booleanIdentifiers, externalIdentifiers)) + .ToArray(); + var conditionalInvocations = spec.Possibilities + .Where(possibility => possibility.BooleanInjection is not null) + .Select(possibility => $" if ({booleanIdentifiers[possibility.BooleanInjection!.Key]}) source.Add({BuildStaticResolveInjectionInvocation(spec, possibility, booleanIdentifiers, externalIdentifiers)});") + .ToArray(); + + return $@" internal static IEnumerable<{spec.TypeFullName}> ResolveAllCore({parameterList}) + {{ + if (!state.EnterCollection(""{spec.TypeFullName}"")) + return Array.Empty<{spec.TypeFullName}>(); + + try {{ - if (container != null) + var source = new List<{spec.TypeFullName}>({resolveInvocations.Length}) {{ + {string.Join(",\n ", resolveInvocations)} + }}; +{string.Join("\n", conditionalInvocations)} + if (container is not null) {{ - var cached = container.{chosen.LazyFieldName}; - if (cached != null) return cached; - lock (container.m_lock) + var b = container.Base; + while (b is not null) {{ - cached = container.{chosen.LazyFieldName}; - if (cached != null) return cached; - var value = {creation}; - container.GetResolvedInstances().Add(new global::System.WeakReference(value)); - container.{chosen.LazyFieldName} = value; - return value; + if (b.TryResolve>(out var additional)) + source.AddRange(additional!); + b = b.Base; }} - }} - return {nullCreation}; - }}"; - } - return $@" public static {ifaceFull} Resolve({ClassName}? container) - {{ - if (container != null) - {{ - var cached = container.{chosen.LazyFieldName}; - if (cached != null) return cached; - lock (container.m_lock) + b = container.Inheritor; + while (b is not null) {{ - cached = container.{chosen.LazyFieldName}; - if (cached != null) return cached; - return container.{chosen.LazyFieldName} = {creation}; + if (b.TryResolve>(out var additional)) + source.AddRange(additional!); + b = b.Inheritor; }} }} - return {nullCreation}; - }}"; - } - // Transient + disposable: construct and register for disposal tracking. - if (chosen.Disposable) + return source; + }} + finally + {{ + state.ExitCollection(""{spec.TypeFullName}""); + }} + }}"; + } + + private static string BuildStaticResolveInjectionMethod( + StaticExtensionSpec spec, + InjectionData injection, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + var helperName = GetStaticInjectionHelperName(injection); + var parameterList = BuildStaticMethodParameterList(spec, booleanIdentifiers, externalIdentifiers, includeContainer: true, includeState: true); + var createInvocation = BuildStaticCreateInjectionInvocation(spec, injection, booleanIdentifiers, externalIdentifiers); + + if (injection.Singleton || injection.Scoped) { - if (creation != nullCreation) + if (injection.Disposable) { - return $@" public static {ifaceFull} Resolve({ClassName}? container) + return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) + {{ + if (container is not null) {{ - var value = container != null ? {creation} : {nullCreation}; - container?.GetResolvedInstances().Add(new global::System.WeakReference(value)); - return value; - }}"; + var cached = container.{injection.LazyFieldName}; + if (cached is not null) + return cached; + + lock (container.m_lock) + {{ + cached = container.{injection.LazyFieldName}; + if (cached is not null) + return cached; + + var value = {createInvocation}; + container.GetResolvedInstances().Add(new global::System.WeakReference(value)); + container.{injection.LazyFieldName} = value; + return value; + }} + }} + + return Create_{helperName}({BuildStaticInternalInvocationArguments(spec, booleanIdentifiers, externalIdentifiers, "null", "state")}); + }}"; } - return $@" public static {ifaceFull} Resolve({ClassName}? container) + + return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) + {{ + if (container is not null) {{ - var value = {creation}; - container?.GetResolvedInstances().Add(new global::System.WeakReference(value)); - return value; - }}"; + var cached = container.{injection.LazyFieldName}; + if (cached is not null) + return cached; + + lock (container.m_lock) + {{ + cached = container.{injection.LazyFieldName}; + if (cached is not null) + return cached; + + return container.{injection.LazyFieldName} = {createInvocation}; + }} + }} + + return Create_{helperName}({BuildStaticInternalInvocationArguments(spec, booleanIdentifiers, externalIdentifiers, "null", "state")}); + }}"; } - // Transient, non-disposable: pure inline factory chain. - if (creation != nullCreation) + if (injection.Disposable) { - return $" public static {ifaceFull} Resolve({ClassName}? container) => container != null ? {creation} : {nullCreation};"; + return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) + {{ + if (container is not null) + {{ + var value = {createInvocation}; + container.GetResolvedInstances().Add(new global::System.WeakReference(value)); + return value; + }} + + return Create_{helperName}({BuildStaticInternalInvocationArguments(spec, booleanIdentifiers, externalIdentifiers, "null", "state")}); + }}"; } - return $" public static {ifaceFull} Resolve({ClassName}? container) => {creation};"; + + return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) + {{ + return {createInvocation}; + }}"; } - /// - /// Builds a new ConcreteType(...) expression for the given injection where every - /// DI-registered dependency is resolved via its own static Resolve extension. - /// - private static string InlineCreation( + private static string BuildStaticCreateInjectionMethod( + StaticExtensionSpec spec, InjectionData injection, + IReadOnlyDictionary specs, ImmutableArray availableInterfaces, - bool nullContainer) + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) { - HashSet? missing = null; + var helperName = GetStaticInjectionHelperName(injection); + var parameterList = BuildStaticMethodParameterList(spec, booleanIdentifiers, externalIdentifiers, includeContainer: true, includeState: true); + var createExpression = BuildStaticCreateExpression(injection, specs, availableInterfaces, booleanIdentifiers, externalIdentifiers); + var returnStatement = createExpression.StartsWith("throw ", StringComparison.Ordinal) + ? createExpression + ";" + : "return " + createExpression + ";"; + + return $@" private static {injection.TypeFullName} Create_{helperName}({parameterList}) + {{ + {returnStatement} + }}"; + } + + private static string BuildStaticCreateExpression( + InjectionData injection, + IReadOnlyDictionary specs, + ImmutableArray availableInterfaces, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + if (injection.Lambda is LambdaData lambda) + { + if (!specs.TryGetValue(lambda.ContainingTypeFullName, out var containingSpec)) + return BuildStaticMissingImplementationExpression(lambda.ContainingTypeFullName); + + var containingInvocation = BuildStaticResolveInvocation(containingSpec, "ResolveCore", booleanIdentifiers, externalIdentifiers); + if (!lambda.IsMethod) + return $"{containingInvocation}.{lambda.MemberName}"; + + var lambdaArguments = BuildStaticArgumentExpressions(lambda.MethodParameters, specs, booleanIdentifiers, externalIdentifiers); + return $"{containingInvocation}.{lambda.MemberName}({string.Join(", ", lambdaArguments)})"; + } + + HashSet? missing = null; HashSet? nullableDefaults = null; - var ctor = GetBestConstructor(injection, availableInterfaces, ref missing, ref nullableDefaults); + var constructor = GetBestConstructor(injection, availableInterfaces, ref missing, ref nullableDefaults); + if (constructor is null) + return BuildStaticMissingImplementationExpression(injection.TypeFullName); - if (ctor == null) - return $"default! /* no constructor found for {injection.TypeFullName} */"; + var constructorArguments = BuildStaticArgumentExpressions(constructor.Parameters, specs, booleanIdentifiers, externalIdentifiers); + return $"new {injection.TypeFullName}({string.Join(", ", constructorArguments)})"; + } - var resolveArg = nullContainer ? "null" : "container"; - var args = new List(); + private static List BuildStaticArgumentExpressions( + ImmutableArray parameters, + IReadOnlyDictionary specs, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + var arguments = new List(); - foreach (var parameter in ctor.Parameters) + foreach (var parameter in parameters) { - // Nullable params with no DI registration → null - if (nullableDefaults?.Contains(parameter) == true) - { - args.Add("null"); - continue; - } + var argumentExpression = BuildStaticArgumentExpression(parameter, specs, booleanIdentifiers, externalIdentifiers); + if (argumentExpression is not null) + arguments.Add(argumentExpression); + } - var typeLookup = parameter.IsNullable - ? parameter.TypeFullName.TrimEnd('?') - : parameter.TypeFullName; + return arguments; + } - // DI-registered: recurse via the static extension - if (availableInterfaces.Contains(typeLookup)) + private static string? BuildStaticArgumentExpression( + ParameterData parameter, + IReadOnlyDictionary specs, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + if (parameter.IsCollection && parameter.CollectionElementFullName is not null) + { + if (!specs.TryGetValue(parameter.CollectionElementFullName, out var collectionSpec)) { - args.Add($"{typeLookup}.Resolve({resolveArg})"); - continue; + return BuildStaticCollectionConversion(parameter, $"global::System.Array.Empty<{parameter.CollectionElementFullName}>()"); } - // C# will supply the default value — omit the argument - if (parameter.HasExplicitDefault) - continue; + return BuildStaticCollectionConversion( + parameter, + BuildStaticResolveAllInvocation(collectionSpec, booleanIdentifiers, externalIdentifiers)); + } - // Empty params array — omit - if (parameter.IsParams) - continue; + var typeLookup = parameter.IsNullable + ? parameter.TypeFullName.TrimEnd('?') + : parameter.TypeFullName; - // Collection: convert from the container's IEnumerable lazy property to the - // exact collection type the constructor demands, mirroring CollectionConstructorArg. - if (parameter.IsCollection && parameter.CollectionElementFullName != null) - { - var elemType = parameter.CollectionElementFullName; - string collectionArg; - if (nullContainer) - { - collectionArg = parameter.CollectionKind switch - { - CollectionKind.List => $"new global::System.Collections.Generic.List<{elemType}>()", - CollectionKind.ImmutableArray => $"global::System.Collections.Immutable.ImmutableArray<{elemType}>.Empty", - CollectionKind.ReadOnlySpan => $"global::System.ReadOnlySpan<{elemType}>.Empty", - _ => $"global::System.Array.Empty<{elemType}>()", // Array or Enumerable - }; - } - else - { - var src = $"container!.coll_{parameter.CollectionElementMemberName}"; - collectionArg = parameter.CollectionKind switch - { - CollectionKind.Array => $"{src}.ToArray()", - CollectionKind.List => $"{src}.ToList()", - CollectionKind.ImmutableArray => $"global::System.Collections.Immutable.ImmutableArray.CreateRange({src})", - CollectionKind.ReadOnlySpan => $"new global::System.ReadOnlySpan<{elemType}>({src}.ToArray())", - _ => src, // Enumerable — direct - }; - } - args.Add(collectionArg); - continue; - } + if (specs.TryGetValue(typeLookup, out var dependencySpec)) + return BuildStaticResolveInvocation(dependencySpec, "ResolveCore", booleanIdentifiers, externalIdentifiers); - if (parameter.IsNullable) - { - args.Add("null"); - continue; - } + if (parameter.HasExplicitDefault || parameter.IsParams) + return null; + + if (parameter.IsNullable) + return "null"; + + if (externalIdentifiers.TryGetValue(parameter.TypeFullName, out var identifier)) + return identifier; + + return BuildStaticMissingImplementationExpression(parameter.TypeFullName); + } + + private static string BuildStaticCollectionConversion(ParameterData parameter, string sourceExpression) + { + if (!parameter.IsCollection) + return sourceExpression; + + return parameter.CollectionKind switch + { + CollectionKind.Array => $"{sourceExpression}.ToArray()", + CollectionKind.List => $"{sourceExpression}.ToList()", + CollectionKind.ImmutableArray => $"global::System.Collections.Immutable.ImmutableArray.CreateRange({sourceExpression})", + CollectionKind.ReadOnlySpan => $"new global::System.ReadOnlySpan<{parameter.CollectionElementFullName}>({sourceExpression}.ToArray())", + _ => sourceExpression + }; + } + + private static string BuildStaticResolveInvocation( + StaticExtensionSpec spec, + string methodName, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + return $"{spec.ExtensionClassName}.{methodName}({BuildStaticInternalInvocationArguments(spec, booleanIdentifiers, externalIdentifiers, "container", "state")})"; + } + + private static string BuildStaticResolveAllInvocation( + StaticExtensionSpec spec, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + return $"{spec.ExtensionClassName}.ResolveAllCore({BuildStaticInternalInvocationArguments(spec, booleanIdentifiers, externalIdentifiers, "container", "state")})"; + } - // Required non-DI parameter: sourced from the container's internal field. - // When nullContainer is true we use default! — callers accept that null-container - // mode cannot satisfy non-DI required dependencies. - args.Add(nullContainer ? "default!" : $"container!.{parameter.Name}"); + private static string BuildStaticResolveSelectionExpression( + StaticExtensionSpec spec, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + if (spec.Possibilities.All(possibility => possibility.BooleanInjection is null)) + return BuildStaticResolveInjectionInvocation(spec, spec.Possibilities.Last(), booleanIdentifiers, externalIdentifiers); + + var keys = spec.Possibilities.Select(possibility => possibility.BooleanInjection?.Key).OfType().Distinct().Reverse().ToArray(); + var fallback = spec.Possibilities.LastOrDefault(possibility => possibility.BooleanInjection is null); + var expression = fallback is not null + ? BuildStaticResolveInjectionInvocation(spec, fallback, booleanIdentifiers, externalIdentifiers) + : BuildStaticMissingImplementationExpression(spec.TypeFullName); + + foreach (var key in keys) + { + var selected = spec.Possibilities.LastOrDefault(possibility => + possibility.BooleanInjection?.Value == true && possibility.BooleanInjection?.Key == key) ?? fallback; + + expression = $"{booleanIdentifiers[key]} ? {(selected is not null ? BuildStaticResolveInjectionInvocation(spec, selected, booleanIdentifiers, externalIdentifiers) : BuildStaticMissingImplementationExpression(spec.TypeFullName))} : {expression}"; } - return $"new {injection.TypeFullName}({string.Join(", ", args)})"; + return expression; + } + + private static string BuildStaticResolveInjectionInvocation( + StaticExtensionSpec spec, + InjectionData injection, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + return $"Resolve_{GetStaticInjectionHelperName(injection)}({BuildStaticInternalInvocationArguments(spec, booleanIdentifiers, externalIdentifiers, "container", "state")})"; + } + + private static string BuildStaticCreateInjectionInvocation( + StaticExtensionSpec spec, + InjectionData injection, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + return $"Create_{GetStaticInjectionHelperName(injection)}({BuildStaticInternalInvocationArguments(spec, booleanIdentifiers, externalIdentifiers, "container", "state")})"; + } + + private static string BuildStaticMethodParameterList( + StaticExtensionSpec spec, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers, + bool includeContainer, + bool includeState) + { + var parts = new List(); + if (includeContainer) + parts.Add($"{ClassName}? container"); + if (includeState) + parts.Add("StaticResolveState state"); + + var runtimeDeclarations = BuildStaticRuntimeArgumentDeclarations(spec, booleanIdentifiers, externalIdentifiers); + if (!string.IsNullOrEmpty(runtimeDeclarations)) + parts.Add(runtimeDeclarations); + + return string.Join(", ", parts); + } + + private static string BuildStaticRuntimeArgumentDeclarations( + StaticExtensionSpec spec, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + var parts = spec.BooleanKeys.Select(key => $"bool {booleanIdentifiers[key]}") + .Concat(spec.ExternalParameters.Select(parameter => $"{parameter.TypeFullName} {externalIdentifiers[parameter.TypeFullName]}")) + .ToArray(); + + return string.Join(", ", parts); + } + + private static string BuildStaticRuntimeArgumentValues( + StaticExtensionSpec spec, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + var parts = spec.BooleanKeys.Select(key => booleanIdentifiers[key]) + .Concat(spec.ExternalParameters.Select(parameter => externalIdentifiers[parameter.TypeFullName])) + .ToArray(); + + return string.Join(", ", parts); + } + + private static string BuildStaticInternalInvocationArguments( + StaticExtensionSpec spec, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers, + string containerExpression, + string stateExpression) + { + var runtimeValues = BuildStaticRuntimeArgumentValues(spec, booleanIdentifiers, externalIdentifiers); + if (string.IsNullOrEmpty(runtimeValues)) + return $"{containerExpression}, {stateExpression}"; + + return $"{containerExpression}, {stateExpression}, {runtimeValues}"; + } + + private static string BuildStaticMissingImplementationExpression(string typeFullName) + { + return $"throw new global::System.InvalidOperationException(\"Cannot resolve {typeFullName} without a matching implementation\")"; + } + + private static string GetStaticInjectionHelperName(InjectionData injection) + { + if (injection.Lambda is null) + return injection.TypeMemberName; + + var containingType = injection.Lambda.ContainingTypeMemberName.Replace("()", string.Empty); + return $"{containingType}_{injection.Lambda.MemberName}_{injection.TypeMemberName}"; } } } \ No newline at end of file diff --git a/FactoryGenerator/Injection.cs b/FactoryGenerator/Injection.cs index 890bf13..3ca4c02 100644 --- a/FactoryGenerator/Injection.cs +++ b/FactoryGenerator/Injection.cs @@ -94,8 +94,8 @@ public static class Injection interfaces = interfaces.Add(namedTypeSymbol); interfaces = interfaces.AddRange(attributedInterfaces); - var isDisposable = namedTypeSymbol.AllInterfaces.Any(i => i.Name.Equals("IDisposable")); - var disposableIface = interfaces.FirstOrDefault(i => i.Name.Contains("IDisposable")); + var isDisposable = namedTypeSymbol.AllInterfaces.Any(i => i.SpecialType == SpecialType.System_IDisposable); + var disposableIface = interfaces.FirstOrDefault(i => i.SpecialType == SpecialType.System_IDisposable); if (disposableIface is not null) interfaces = interfaces.Remove(disposableIface); diff --git a/FactoryGenerator/InjectionData.cs b/FactoryGenerator/InjectionData.cs index 249937f..cc3d70e 100644 --- a/FactoryGenerator/InjectionData.cs +++ b/FactoryGenerator/InjectionData.cs @@ -172,32 +172,4 @@ public bool Equals(LambdaData? other) public override int GetHashCode() => ContainingTypeFullName.GetHashCode(); } - public sealed class UsageData : IEquatable - { - public string FullName { get; } - public string MemberName { get; } // SymbolUtility.MemberName(type) without "()" - public string ElementTypeFullName { get; } - public string ElementTypeMemberName { get; } // without "()" - - public UsageData(string fullName, string memberName, string elementTypeFullName, string elementTypeMemberName) - { - FullName = fullName; - MemberName = memberName; - ElementTypeFullName = elementTypeFullName; - ElementTypeMemberName = elementTypeMemberName; - } - - public bool Equals(UsageData? other) - { - if (other is null) return false; - if (ReferenceEquals(this, other)) return true; - return FullName == other.FullName - && MemberName == other.MemberName - && ElementTypeFullName == other.ElementTypeFullName - && ElementTypeMemberName == other.ElementTypeMemberName; - } - - public override bool Equals(object? obj) => obj is UsageData other && Equals(other); - public override int GetHashCode() => FullName.GetHashCode(); - } } diff --git a/README.md b/README.md index 34ddf85..3cb505c 100644 --- a/README.md +++ b/README.md @@ -182,12 +182,21 @@ var chain = ChainA.Resolve(container); Each generated `Resolve` method inlines the full construction chain directly — no dictionary lookup, no factory-method indirection. Singletons use double-checked locking against the container's cache field, while transients emit a pure `new` expression. -**Null-container mode:** Passing `null` instead of a container instance bypasses the singleton/scoped cache entirely and performs a fresh allocation on every call. This is useful for one-off instances where you want pure construction cost without any shared state: +If a resolution graph depends on runtime booleans or external constructor values, the generated static method carries those inputs explicitly: +```csharp +var switched = ISwitchableInterface.Resolve(container, testBool: true); +var built = Constructed.Resolve(container, nonInjectedClassArgument: options); +``` + +**Null-container mode:** Passing `null` instead of a container instance bypasses the singleton/scoped cache entirely and performs a fresh allocation on every call, while still requiring any runtime inputs needed by the graph: ```csharp // Fresh allocation every time — no singleton cache var fresh = ISingleton.Resolve(null); +var switched = ISwitchableInterface.Resolve(testBool: true); ``` +Collection dependencies (`IEnumerable`, arrays, `List`, `ImmutableArray`, `ReadOnlySpan`) are resolved through the same generated static pipeline, so the extension path and the normal container path construct equivalent object graphs. Direct `Resolve>()` calls are generated for every registered service type, even when no constructor or source usage referenced that collection shape ahead of time. + The static extensions are generated alongside the standard dictionary-based container and require no additional configuration. If the consuming project's language version is below C# 14, the extensions are simply not emitted. **Opting out:** If you are on C# 14+ but do not want the static extensions (for example, to reduce generated code size or avoid conflicts), set the following property in your `.csproj`: diff --git a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs index e93bf3e..b2b9567 100644 --- a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs +++ b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs @@ -1,3 +1,4 @@ +using System; using System.Net; using System.Threading.Tasks; using Microsoft.AspNetCore.Builder; @@ -32,6 +33,27 @@ public class OtherService : IOtherService public string GetValue() => "Hello from IServiceProvider"; } +public interface IRequestScopedDependency +{ + Guid Id { get; } +} + +public class RequestScopedDependency : IRequestScopedDependency +{ + public Guid Id { get; } = Guid.NewGuid(); +} + +public interface IRequestScopedFactoryService +{ + Guid GetRequestId(); +} + +[Inject] +public class RequestScopedFactoryService(IRequestScopedDependency dependency) : IRequestScopedFactoryService +{ + public Guid GetRequestId() => dependency.Id; +} + public class IntegrationTests { [Test] @@ -46,6 +68,7 @@ public async Task Middleware_Integrates_FactoryGenerator_With_RequestServices() .ConfigureServices(services => { services.AddSingleton(); + services.AddScoped(); }) .Configure(app => { @@ -63,8 +86,10 @@ public async Task Middleware_Integrates_FactoryGenerator_With_RequestServices() { var myService = context.RequestServices.GetRequiredService(); var otherService = context.RequestServices.GetRequiredService(); + var requestScopedFactoryService = context.RequestServices.GetRequiredService(); + var requestScopedDependency = context.RequestServices.GetRequiredService(); - await context.Response.WriteAsync($"{myService.GetValue()} | {otherService.GetValue()}"); + await context.Response.WriteAsync($"{myService.GetValue()} | {otherService.GetValue()} | {requestScopedFactoryService.GetRequestId()} | {requestScopedDependency.Id}"); }); }); }) @@ -78,6 +103,55 @@ public async Task Middleware_Integrates_FactoryGenerator_With_RequestServices() // Assert response.StatusCode.ShouldBe(HttpStatusCode.OK); var content = await response.Content.ReadAsStringAsync(); - content.ShouldBe("Hello from FactoryGenerator | Hello from IServiceProvider"); + var parts = content.Split(" | "); + parts.Length.ShouldBe(4); + parts[0].ShouldBe("Hello from FactoryGenerator"); + parts[1].ShouldBe("Hello from IServiceProvider"); + parts[2].ShouldBe(parts[3]); + } + + [Test] + public async Task Middleware_Uses_Current_RequestScope_For_FrameworkScopedDependencies() + { + var host = await new HostBuilder() + .ConfigureWebHost(webBuilder => + { + webBuilder + .UseTestServer() + .ConfigureServices(services => + { + services.AddScoped(); + }) + .Configure(app => + { + var adapter = app.ApplicationServices.ToContainer(); + var container = new Generated.DependencyInjectionContainer(adapter); + + app.UseFactoryGenerator(container); + + app.Run(async context => + { + var requestScopedFactoryService = context.RequestServices.GetRequiredService(); + var requestScopedDependency = context.RequestServices.GetRequiredService(); + + await context.Response.WriteAsync($"{requestScopedFactoryService.GetRequestId()}|{requestScopedDependency.Id}"); + }); + }); + }) + .StartAsync(); + + var client = host.GetTestClient(); + + var firstResponse = await client.GetStringAsync("/"); + var secondResponse = await client.GetStringAsync("/"); + + var firstParts = firstResponse.Split('|'); + var secondParts = secondResponse.Split('|'); + + firstParts.Length.ShouldBe(2); + secondParts.Length.ShouldBe(2); + firstParts[0].ShouldBe(firstParts[1]); + secondParts[0].ShouldBe(secondParts[1]); + firstParts[0].ShouldNotBe(secondParts[0]); } } diff --git a/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs b/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs index 81ff654..347a01d 100644 --- a/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs +++ b/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs @@ -45,6 +45,30 @@ public void BuildChainCreatesWorkingContainerPipeline() final.Resolve().ShouldNotBeNull(); } + [Test] + public void BuildChainWithoutAssemblyListSkipsCurrentContainerAssembly() + { + _ = typeof(ContainerEntryPoint); + var baseContainer = new DependencyInjectionContainer(default, default, new NonInjectedClass()); + + var final = ContainerRegistry.BuildChain(baseContainer); + + ReferenceEquals(final, baseContainer).ShouldBeTrue(); + final.Inheritor.ShouldBeNull(); + } + + [Test] + public void BuildChainWithExplicitAssemblyListSkipsCurrentContainerAssembly() + { + _ = typeof(ContainerEntryPoint); + var baseContainer = new DependencyInjectionContainer(default, default, new NonInjectedClass()); + + var final = ContainerRegistry.BuildChain(baseContainer, new[] { "Inheritor" }); + + ReferenceEquals(final, baseContainer).ShouldBeTrue(); + final.Inheritor.ShouldBeNull(); + } + [Test] public void ContainerEntryPointAssemblyNameIsCorrect() { diff --git a/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj b/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj index 3472792..40d2225 100644 --- a/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj +++ b/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj @@ -2,6 +2,7 @@ net10.0 + preview enable enable false @@ -14,6 +15,7 @@ + diff --git a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs index 89b986b..ca9eb7c 100644 --- a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs +++ b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs @@ -43,6 +43,27 @@ public void ResolveUsesArguments() myContainer.Resolve().NonInjectedClassArgument.ShouldBe(dummy); } + [Test] + public void StaticExtensionsPropagateDirectExternalArguments() + { + var dummy = new NonInjectedClass(); + Constructed.Resolve(dummy).NonInjectedClassArgument.ShouldBe(dummy); + } + + [Test] + public void StaticExtensionsPropagateTransitiveExternalArguments() + { + var dummy = new NonInjectedClass(); + ConstructedConsumer.Resolve(dummy).Value.NonInjectedClassArgument.ShouldBe(dummy); + } + + [Test] + public void StaticExtensionsPropagateExternalArgumentsIntoCollections() + { + var dummy = new NonInjectedClass(); + ConstructedArrayConsumer.Resolve(dummy).Items.ShouldHaveSingleItem().NonInjectedClassArgument.ShouldBe(dummy); + } + [Test] [Arguments(true, typeof(EnabledImplementation))] [Arguments(false, typeof(FallbackImplementation))] @@ -52,6 +73,27 @@ public void PickupSingleInjectionWithBoolean(bool value, System.Type expected) myContainer.Resolve().ShouldBeOfType(expected); } + [Test] + public void StaticExtensionsPropagateDirectBooleanArguments() + { + ISwitchableInterface.Resolve(true).ShouldBeOfType(); + ISwitchableInterface.Resolve(false).ShouldBeOfType(); + } + + [Test] + public void StaticExtensionsPropagateTransitiveBooleanArguments() + { + BooleanConsumer.Resolve(true).Value.ShouldBeOfType(); + BooleanConsumer.Resolve(false).Value.ShouldBeOfType(); + } + + [Test] + public void StaticExtensionsPropagateBooleanArgumentsIntoCollections() + { + SwitchableArrayConsumer.Resolve(false).Items.Count().ShouldBe(1); + SwitchableArrayConsumer.Resolve(true).Items.Count().ShouldBe(2); + } + [Test] public void PickupSingleInjectionFromMethod() { @@ -88,6 +130,18 @@ public void DontPickupIDisposable() true.ShouldBeFalse(); } + [Test] + public void InterfacesContainingIDisposableInTheNameRemainResolvable() + { + m_container.Resolve().ShouldBeOfType(); + } + + [Test] + public void NonSystemIDisposableInterfacesRemainResolvable() + { + m_container.Resolve().ShouldBeOfType(); + } + [Test] public void DontPickupExcluded() { @@ -121,6 +175,18 @@ public void InheritorsOverride() m_container.Resolve().ShouldBeOfType(); } + [Test] + public void OverrideImplementationsPreventFalsePositiveCycleDetection() + { + m_container.Resolve().ShouldBeOfType(); + } + + [Test] + public void BestConstructorsPreventFalsePositiveCycleDetection() + { + m_container.Resolve().ShouldBeOfType(); + } + [Test] public void DisposingContainerDisposesSingletons() @@ -190,12 +256,31 @@ public void ArrayExpressionsCollect() m_container.Resolve().Arrays.Count().ShouldBe(3); } + [Test] + public void StaticExtensionsResolveCollectionsInNullContainerMode() + { + ArrayConsumer.Resolve(null).Arrays.Count().ShouldBe(3); + } + [Test] public void RequestedArraysArePresent() { Program.Method().Count().ShouldBe(3); } + [Test] + public void EnumerablesAreResolvableWithoutUsageSites() + { + m_container.Resolve>().Count().ShouldBe(2); + } + + [Test] + public void DuplicateRequestedArrayUsagesDoNotDuplicateLookupKeys() + { + Program.Method().Count().ShouldBe(3); + Program.MethodAgain().Count().ShouldBe(3); + } + [Test] public void BooleanFallbackIsOverriden() { @@ -262,6 +347,37 @@ public void HierarchicalContainersPropgatesBooleansUnknownToIt() newContainer.GetBoolean("C").ShouldBe(false); } + [Test] + public void DisposingChildContainerDoesNotDisposeBaseContainer() + { + var baseContainer = new DependencyInjectionContainer(false, default, default!); + var singleton = baseContainer.Resolve().ShouldBeOfType(); + var child = new DependencyInjectionContainer(baseContainer); + + child.Dispose(); + + singleton.WasDisposed.ShouldBeFalse(); + baseContainer.Inheritor.ShouldBeNull(); + + baseContainer.Dispose(); + singleton.WasDisposed.ShouldBeTrue(); + } + + [Test] + public void DisposingChildContainerUnregistersItFromParentCollections() + { + var parent = new DependencyInjectionContainer(false, false, new NonInjectedClass()); + var child = new DependencyInjectionContainer(parent); + + parent.Resolve>().Count().ShouldBe(6); + + child.Dispose(); + + parent.Resolve>().Count().ShouldBe(3); + parent.Inheritor.ShouldBeNull(); + parent.Dispose(); + } + // ── Nullable parameter tests ────────────────────────────────────────────── [Test] diff --git a/Tests/TestData/Inherited/Types.cs b/Tests/TestData/Inherited/Types.cs index 98d5125..375ed9b 100644 --- a/Tests/TestData/Inherited/Types.cs +++ b/Tests/TestData/Inherited/Types.cs @@ -7,12 +7,20 @@ public interface IType; public interface IOverridable; +public interface IOverrideCycle; + [Inject] public class Type : IType; [Inject] public class Overriden : IOverridable; +[Inject] +public class OverrideCycleBase(IOverrideCycle self) : IOverrideCycle +{ + public IOverrideCycle Self { get; } = self; +} + public interface ISingleton; [Inject, Singleton] @@ -43,8 +51,49 @@ public class Constructed(NonInjectedClass nonInjectedClassArgument, ISingleton i public ISingleton InjectedArgument { get; } = injectedArgument; } +[Inject, Self] +public class ConstructedConsumer(Constructed value) +{ + public Constructed Value { get; } = value; +} + +[Inject, Self] +public class ConstructedArrayConsumer(IEnumerable items) +{ + public IEnumerable Items { get; } = items; +} + +[Inject, Self] +public class BooleanConsumer(ISwitchableInterface value) +{ + public ISwitchableInterface Value { get; } = value; +} + +[Inject, Self] +public class SwitchableArrayConsumer(IEnumerable items) +{ + public IEnumerable Items { get; } = items; +} + public interface IMethodResult; +public interface IMultiConstructorCycle; + +public class ExternalOnlyDependency; + +[Inject] +public class MultiConstructorCycle : IMultiConstructorCycle +{ + public MultiConstructorCycle() + { + } + + public MultiConstructorCycle(IMultiConstructorCycle self, ExternalOnlyDependency externalOnlyDependency) + : this() + { + } +} + public class MethodResult : IMethodResult; public interface IMethodSource @@ -112,8 +161,29 @@ public class RequestedArray2 : IRequestedArray; [Inject] public class RequestedArray3 : IRequestedArray; +public interface IUnrequestedEnumerable; + +[Inject] +public class UnrequestedEnumerable1 : IUnrequestedEnumerable; + +[Inject] +public class UnrequestedEnumerable2 : IUnrequestedEnumerable; + public interface IDisposer; +public interface INotIDisposable; + +[Inject] +public class NotDisposableNameMatch : INotIDisposable; + +public class CustomDisposableTypes +{ + public interface IDisposable; + + [Inject] + public class CustomDisposable : IDisposable; +} + [Inject] public class DisposableNonSingleton : IDisposer, IDisposable { diff --git a/Tests/TestData/Inheritor/Types.cs b/Tests/TestData/Inheritor/Types.cs index 25dacab..1163552 100644 --- a/Tests/TestData/Inheritor/Types.cs +++ b/Tests/TestData/Inheritor/Types.cs @@ -10,6 +10,9 @@ public class Overrider : IOverridable; [Inject] public class OverridingBoolean : IOverrideBoolean; +[Inject] +public class OverrideCycleResolved : IOverrideCycle; + [Inject] public class ChainA(ChainB B, ChainC C, ChainD D) { @@ -50,6 +53,13 @@ public static IEnumerable Method() var array = container.Resolve>(); return array; } + + public static IEnumerable MethodAgain() + { + var container = new DependencyInjectionContainer(false, false, null!); + var array = container.Resolve>(); + return array; + } } // ── Inheritor + Base array tests ───────────────────────────────────────────── // Additional ISplitArray implementations in the Inheritor project. When a child From 8aff9a31766b96afe83e575e4606f1273a5793ab Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 13:04:44 +0200 Subject: [PATCH 14/21] More tests --- .../GeneratorBehaviorTests.cs | 136 ++++++++++++++++++ 1 file changed, 136 insertions(+) create mode 100644 Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs diff --git a/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs b/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs new file mode 100644 index 0000000..7237df9 --- /dev/null +++ b/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs @@ -0,0 +1,136 @@ +using System; +using System.IO; +using System.Linq; +using FactoryGenerator.Attributes; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using Shouldly; + +namespace FactoryGenerator.Tests; + +public class GeneratorBehaviorTests +{ + [Test] + public void GeneratorRejectsMultipleExternalValuesOfSameType() + { + const string source = """ +using FactoryGenerator.Attributes; + +namespace Sample +{ +public interface IService +{ +} + +public class ExternalValue +{ +} + +[Inject] +public class FirstConsumer : IService +{ + public FirstConsumer(ExternalValue first) + { + } +} + +[Inject, Self] +public class SecondConsumer +{ + public SecondConsumer(ExternalValue second) + { + } +} +} +"""; + + var compilation = CreateCompilation(source); + var (runResult, _) = RunGenerator(compilation); + runResult.Results.Length.ShouldBe(1); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldNotBeNull(); + generatorResult.Exception!.Message.ShouldContain("Multiple externally provided values of the same type"); + generatorResult.Exception.Message.ShouldContain("Sample.ExternalValue"); + } + + [Test] + public void GeneratorSupportsBooleanKeysThatAreNotIdentifiers() + { + const string source = """ +using FactoryGenerator.Attributes; + +namespace Sample +{ +public interface IService +{ +} + +[Inject, Boolean("feature-flag")] +public class EnabledService : IService +{ +} + +[Inject] +public class FallbackService : IService +{ +} +} +"""; + + var compilation = CreateCompilation(source); + var (runResult, outputCompilation) = RunGenerator(compilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldBeNull(); + outputCompilation.GetDiagnostics() + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .ToArray() + .ShouldBeEmpty(); + + var generatedSource = string.Join(Environment.NewLine, generatorResult.GeneratedSources.Select(sourceResult => sourceResult.SourceText.ToString())); + generatedSource.ShouldContain("\"feature-flag\""); + generatedSource.ShouldNotContain("bool feature-flag"); + generatedSource.ShouldContain("Resolve(DependencyInjectionContainer? container, bool boolean_feature_flag)"); + } + + private static CSharpCompilation CreateCompilation(string source) + { + var syntaxTree = CSharpSyntaxTree.ParseText(source, new CSharpParseOptions(LanguageVersion.Preview)); + var excludedAssemblies = new[] + { + "Benchmarks", + "FactoryGenerator", + "FactoryGenerator.Attributes", + "FactoryGenerator.Extensions.AspNetCore", + "FactoryGenerator.Extensions.AspNetCore.Tests", + "FactoryGenerator.Tests", + "Inherited", + "Inheritor", + "TestWebApp" + }; + var references = ((string?)AppContext.GetData("TRUSTED_PLATFORM_ASSEMBLIES"))! + .Split(Path.PathSeparator) + .Where(path => !excludedAssemblies.Contains(Path.GetFileNameWithoutExtension(path), StringComparer.Ordinal)) + .Select(path => (MetadataReference)MetadataReference.CreateFromFile(path)) + .ToList(); + + references.Add(MetadataReference.CreateFromFile(typeof(InjectAttribute).Assembly.Location)); + + return CSharpCompilation.Create( + assemblyName: "GeneratorBehaviorTests", + syntaxTrees: new[] { syntaxTree }, + references: references, + options: new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary)); + } + + private static (GeneratorDriverRunResult RunResult, Compilation OutputCompilation) RunGenerator(CSharpCompilation compilation) + { + var parseOptions = (CSharpParseOptions)compilation.SyntaxTrees.First().Options; + GeneratorDriver driver = CSharpGeneratorDriver.Create( + [new global::FactoryGenerator.FactoryGenerator().AsSourceGenerator()], + parseOptions: parseOptions); + driver = driver.RunGeneratorsAndUpdateCompilation(compilation, out var outputCompilation, out _); + return (driver.GetRunResult(), outputCompilation); + } +} From acb3fb6ab3f5e16b76265223c189b7905d7b6d39 Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 13:25:06 +0200 Subject: [PATCH 15/21] Controllable injection priority and more correct tree resolution --- .../Attributes/InjectionPriorityAttribute.cs | 9 ++ FactoryGenerator/FactoryGenerator.cs | 72 +++++++++++++--- FactoryGenerator/Injection.cs | 24 +++++- FactoryGenerator/InjectionData.cs | 11 ++- README.md | 16 ++++ .../GeneratorBehaviorTests.cs | 82 ++++++++++++++++++- 6 files changed, 195 insertions(+), 19 deletions(-) create mode 100644 FactoryGenerator.Attributes/Attributes/InjectionPriorityAttribute.cs diff --git a/FactoryGenerator.Attributes/Attributes/InjectionPriorityAttribute.cs b/FactoryGenerator.Attributes/Attributes/InjectionPriorityAttribute.cs new file mode 100644 index 0000000..5c5d606 --- /dev/null +++ b/FactoryGenerator.Attributes/Attributes/InjectionPriorityAttribute.cs @@ -0,0 +1,9 @@ +using System; + +namespace FactoryGenerator.Attributes; + +[AttributeUsage(AttributeTargets.Assembly)] +public class InjectionPriorityAttribute(int priority) : Attribute +{ + public int Priority { get; } = priority; +} diff --git a/FactoryGenerator/FactoryGenerator.cs b/FactoryGenerator/FactoryGenerator.cs index a7f7e2c..66810bb 100644 --- a/FactoryGenerator/FactoryGenerator.cs +++ b/FactoryGenerator/FactoryGenerator.cs @@ -128,7 +128,7 @@ private static INamespaceSymbol GetGlobalNamespace(Compilation compilation, Canc private static IEnumerable GenerateCode(ImmutableArray dataInjections, Compilation compilation, ILogger log) { - CheckForCycles(dataInjections); + CheckForCycles(dataInjections, compilation); log.Log(LogLevel.Debug, "Starting Code Generation"); var usingStatements = $@" using System; @@ -316,7 +316,7 @@ public bool GetBoolean(string key) var booleanKeys = dataInjections.Select(inj => inj.BooleanInjection).Where(b => b is not null) .Select(b => b!.Key).Distinct().ToArray(); - var ordered = OrderInjections(dataInjections, log); + var ordered = OrderInjections(dataInjections, compilation, log); var (interfaceInjectors, interfaceMemberNames) = BuildInterfaceInjectors(ordered); var declarations = new Dictionary(); @@ -739,19 +739,67 @@ internal static void Register() "; } - private static List OrderInjections(ImmutableArray dataInjections, ILogger? log = null) + private static List OrderInjections(ImmutableArray dataInjections, Compilation compilation, ILogger? log = null) { var ordered = dataInjections.Reverse().ToList(); + var assemblyDistances = BuildAssemblyDistances(compilation, ordered.Select(injection => injection.AssemblyName)); - foreach (var injection in ordered.ToArray()) + foreach (var injection in ordered) + log?.Log(LogLevel.Debug, $"Traversing {injection.Name} from {injection.AssemblyName} with priority {injection.AssemblyPriority}"); + + return ordered + .OrderBy(injection => injection.AssemblyPriority) + .ThenByDescending(injection => GetAssemblyDistance(assemblyDistances, injection.AssemblyName)) + .ThenBy(injection => injection.AssemblyName, StringComparer.Ordinal) + .ToList(); + } + + private static Dictionary BuildAssemblyDistances(Compilation compilation, IEnumerable assemblyNames) + { + var relevantAssemblyNames = new HashSet(assemblyNames, StringComparer.Ordinal); + var distances = new Dictionary(StringComparer.Ordinal); + var visited = new HashSet(SymbolEqualityComparer.Default); + var queue = new Queue<(IAssemblySymbol Assembly, int Distance)>(); + + visited.Add(compilation.Assembly); + queue.Enqueue((compilation.Assembly, 0)); + + while (queue.Count > 0 && distances.Count < relevantAssemblyNames.Count) { - log?.Log(LogLevel.Debug, $"Traversing {injection.Name}"); - if (!injection.IsTestType) continue; - ordered.Remove(injection); - ordered.Add(injection); + var current = queue.Dequeue(); + if (relevantAssemblyNames.Contains(current.Assembly.Name) + && (!distances.TryGetValue(current.Assembly.Name, out var existingDistance) + || current.Distance < existingDistance)) + { + distances[current.Assembly.Name] = current.Distance; + } + + foreach (var referencedAssembly in GetReferencedAssemblies(current.Assembly)) + { + if (!visited.Add(referencedAssembly)) + continue; + + queue.Enqueue((referencedAssembly, current.Distance + 1)); + } } - return ordered; + return distances; + } + + private static IEnumerable GetReferencedAssemblies(IAssemblySymbol assembly) + { + foreach (var module in assembly.Modules) + { + foreach (var referencedAssembly in module.ReferencedAssemblySymbols) + yield return referencedAssembly; + } + } + + private static int GetAssemblyDistance(IReadOnlyDictionary assemblyDistances, string assemblyName) + { + return assemblyDistances.TryGetValue(assemblyName, out var distance) + ? distance + : int.MaxValue; } private static (Dictionary> InterfaceInjectors, Dictionary InterfaceMemberNames) BuildInterfaceInjectors(IEnumerable ordered) @@ -933,9 +981,9 @@ private static string GetSanitizedIdentifier(string value, string prefix) return SyntaxFacts.IsValidIdentifier(candidate) ? candidate : prefix.TrimEnd('_'); } - private static void CheckForCycles(ImmutableArray dataInjections) + private static void CheckForCycles(ImmutableArray dataInjections, Compilation compilation) { - var ordered = OrderInjections(dataInjections); + var ordered = OrderInjections(dataInjections, compilation); var (interfaceInjectors, _) = BuildInterfaceInjectors(ordered); var availableInterfaceFullNames = interfaceInjectors.Keys.ToImmutableArray(); @@ -1429,7 +1477,7 @@ private static void MakeStaticExtensions( private static string GenerateStaticExtensions( ImmutableArray dataInjections, Compilation compilation) { - var ordered = OrderInjections(dataInjections); + var ordered = OrderInjections(dataInjections, compilation); var (interfaceInjectors, interfaceMemberNames) = BuildInterfaceInjectors(ordered); var availableInterfaces = interfaceInjectors.Keys.ToImmutableArray(); var specs = BuildStaticExtensionSpecs(interfaceInjectors, interfaceMemberNames, availableInterfaces); diff --git a/FactoryGenerator/Injection.cs b/FactoryGenerator/Injection.cs index 3ca4c02..b2c15b6 100644 --- a/FactoryGenerator/Injection.cs +++ b/FactoryGenerator/Injection.cs @@ -41,6 +41,9 @@ public static class Injection } if (namedTypeSymbol is null) return null; + var assembly = symbol.ContainingAssembly ?? namedTypeSymbol.ContainingAssembly; + var assemblyName = assembly?.Name ?? string.Empty; + var assemblyPriority = GetAssemblyPriority(assembly); var singleInstance = false; var acquireChildInterfaces = false; @@ -121,7 +124,8 @@ public static class Injection return new InjectionData( typeFullName: namedTypeSymbol.ToString()!, typeMemberName: typeMemberName, - isTestType: namedTypeSymbol.ToString()!.Contains("Test"), + assemblyName: assemblyName, + assemblyPriority: assemblyPriority, interfaceFullNames: ifaceFullNames, interfaceMemberNames: ifaceMemberNames, singleton: singleInstance, @@ -168,5 +172,23 @@ private static ParameterData ExtractParameter(IParameterSymbol parameter) return new BooleanInjection(true, key); return null; } + + private static int GetAssemblyPriority(IAssemblySymbol? assemblySymbol) + { + if (assemblySymbol is null) + return 0; + + foreach (var attributeData in assemblySymbol.GetAttributes()) + { + if (attributeData.AttributeClass?.ToString() != "FactoryGenerator.Attributes.InjectionPriorityAttribute") + continue; + + if (attributeData.ConstructorArguments.Length == 1 + && attributeData.ConstructorArguments[0].Value is int priority) + return priority; + } + + return 0; + } } } diff --git a/FactoryGenerator/InjectionData.cs b/FactoryGenerator/InjectionData.cs index cc3d70e..b5f321d 100644 --- a/FactoryGenerator/InjectionData.cs +++ b/FactoryGenerator/InjectionData.cs @@ -10,7 +10,8 @@ public sealed class InjectionData : IEquatable { public string TypeFullName { get; } public string TypeMemberName { get; } // MemberName(type) without "()" - public bool IsTestType { get; } + public string AssemblyName { get; } + public int AssemblyPriority { get; } public ImmutableArray InterfaceFullNames { get; } public ImmutableArray InterfaceMemberNames { get; } // parallel to InterfaceFullNames, without "()" public bool Singleton { get; } @@ -25,7 +26,7 @@ public sealed class InjectionData : IEquatable public string LazyFieldName => "m_" + TypeMemberName + (Lambda?.MemberName ?? string.Empty); public InjectionData( - string typeFullName, string typeMemberName, bool isTestType, + string typeFullName, string typeMemberName, string assemblyName, int assemblyPriority, ImmutableArray interfaceFullNames, ImmutableArray interfaceMemberNames, bool singleton, bool scoped, bool disposable, BooleanInjection? booleanInjection, @@ -33,7 +34,8 @@ public InjectionData( { TypeFullName = typeFullName; TypeMemberName = typeMemberName; - IsTestType = isTestType; + AssemblyName = assemblyName; + AssemblyPriority = assemblyPriority; InterfaceFullNames = interfaceFullNames; InterfaceMemberNames = interfaceMemberNames; Singleton = singleton; @@ -50,7 +52,8 @@ public bool Equals(InjectionData? other) if (ReferenceEquals(this, other)) return true; return TypeFullName == other.TypeFullName && TypeMemberName == other.TypeMemberName - && IsTestType == other.IsTestType + && AssemblyName == other.AssemblyName + && AssemblyPriority == other.AssemblyPriority && InterfaceFullNames.SequenceEqual(other.InterfaceFullNames) && InterfaceMemberNames.SequenceEqual(other.InterfaceMemberNames) && Singleton == other.Singleton diff --git a/README.md b/README.md index 3cb505c..eb17be1 100644 --- a/README.md +++ b/README.md @@ -116,6 +116,8 @@ Of note is perhaps `Generated.DependencyInjectionContainer`, this is the Compile | ```Scoped``` | Ensures that this type will be resolved once per created scope, if you do not use IContainer.BeginLifetimeScope(), this behaves like a singleton | ```Inject``` | | ```Boolean(string key)``` | Creates a Runtime switch to decide whether this type should be the one that gets resolved
(or the best fitting fallback option, otherwise) | ```Inject``` | +`InjectionPriority(int)` is an **assembly-level** attribute rather than an injection attribute. + ### Overriding Overriding Injections, i.e if in the graph above _Dependency C_ injected an ```ISomething``` instance and that specific implementation of ```ISomething``` will not work for anything that uses _Dependency B_, then _Dependency B_ can substitute that injection by providing it's own injection of ```ISomething```. This overriding generally follows the project dependency tree, so if _Project A_ depends on _Project B_ which depends on _Project C_, A can override both B and C, but B cannot override A. @@ -123,6 +125,18 @@ Overriding Injections, i.e if in the graph above _Dependency C_ injected an ```I **Note** Overriding Injections will not work if you resolve an ```IEnumerable```, as that will net your a collection of all ```ISomething``` that have been injected. +If you need to override the normal project-graph precedence, you can assign an assembly-level priority: + +```csharp +using FactoryGenerator.Attributes; + +[assembly: InjectionPriority(9)] +``` + +Higher priority values win over lower ones, and assemblies default to priority `0`. If two assemblies have the same priority, FactoryGenerator falls back to the normal dependency graph ordering, where the current project overrides its references and direct references override deeper transitive ones. Assemblies in the same graph tier are then ordered deterministically by assembly name. + +This assembly-level `InjectionPriority` only affects **which implementation wins during generated service resolution** when multiple assemblies provide the same service. It does **not** control plugin/container chaining order in `ContainerRegistry`. + ### Unprovided Values What happens if there are some constructor values needed by certain injected implementations, such as command line arguments, that cannot be known at compile time? @@ -322,5 +336,7 @@ Alternatively, plugins can be registered with a priority for automatic ordering. ContainerRegistry.Register("MyPlugin", ContainerEntryPoint.Create, priority: 10); ``` +This `ContainerRegistry` priority is separate from `[assembly: InjectionPriority(...)]`: it only determines **where a plugin container is inserted in the runtime chain**, not which competing implementation inside a generated container wins for a service. + **Note** This system is fully AOT-compatible. No reflection is used for container discovery — ```[ModuleInitializer]``` methods run automatically when an assembly is loaded. Each generated ```ContainerEntryPoint.Create``` directly instantiates the concrete generated container without ```Activator.CreateInstance``` or type scanning. diff --git a/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs b/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs index 7237df9..abc39d6 100644 --- a/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs +++ b/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs @@ -94,7 +94,65 @@ public class FallbackService : IService generatedSource.ShouldContain("Resolve(DependencyInjectionContainer? container, bool boolean_feature_flag)"); } - private static CSharpCompilation CreateCompilation(string source) + [Test] + public void AssemblyPriorityCanOverrideProjectGraphPrecedence() + { + var baseAssemblyName = "PriorityBase" + Guid.NewGuid().ToString("N"); + var derivedAssemblyName = "PriorityDerived" + Guid.NewGuid().ToString("N"); + + var baseSource = $$""" +using FactoryGenerator.Attributes; + +[assembly: InjectionPriority(9)] + +namespace {{baseAssemblyName}} +{ +public interface IService +{ +} + +[Inject] +public class BaseService : IService +{ +} +} +"""; + + var derivedSource = $$""" +using FactoryGenerator.Attributes; +using {{baseAssemblyName}}; + +namespace {{derivedAssemblyName}} +{ +[Inject] +public class DerivedService : IService +{ +} +} +"""; + + var baseCompilation = CreateCompilation(baseAssemblyName, baseSource); +var (baseReference, _) = EmitReference(baseCompilation); + var derivedCompilation = CreateCompilation(derivedAssemblyName, derivedSource, baseReference); + + var (runResult, outputCompilation) = RunGenerator(derivedCompilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldBeNull(); + outputCompilation.GetDiagnostics() + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .ToArray() + .ShouldBeEmpty(); + + var serviceMemberName = baseAssemblyName + "_IService()"; + var prioritizedImplementationMemberName = baseAssemblyName + "_BaseService()"; + var nonPrioritizedImplementationMemberName = derivedAssemblyName + "_DerivedService()"; + var generatedSource = string.Join(Environment.NewLine, generatorResult.GeneratedSources.Select(sourceResult => sourceResult.SourceText.ToString())); + generatedSource.ShouldContain($"internal {baseAssemblyName}.IService {serviceMemberName} => {prioritizedImplementationMemberName};"); + generatedSource.ShouldNotContain($"internal {baseAssemblyName}.IService {serviceMemberName} => {nonPrioritizedImplementationMemberName};"); + } + + private static CSharpCompilation CreateCompilation(string assemblyName, string source, params MetadataReference[] additionalReferences) { var syntaxTree = CSharpSyntaxTree.ParseText(source, new CSharpParseOptions(LanguageVersion.Preview)); var excludedAssemblies = new[] @@ -116,14 +174,20 @@ private static CSharpCompilation CreateCompilation(string source) .ToList(); references.Add(MetadataReference.CreateFromFile(typeof(InjectAttribute).Assembly.Location)); + references.AddRange(additionalReferences); return CSharpCompilation.Create( - assemblyName: "GeneratorBehaviorTests", + assemblyName: assemblyName, syntaxTrees: new[] { syntaxTree }, references: references, options: new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary)); } + private static CSharpCompilation CreateCompilation(string source) + { + return CreateCompilation("GeneratorBehaviorTests", source); + } + private static (GeneratorDriverRunResult RunResult, Compilation OutputCompilation) RunGenerator(CSharpCompilation compilation) { var parseOptions = (CSharpParseOptions)compilation.SyntaxTrees.First().Options; @@ -133,4 +197,18 @@ private static (GeneratorDriverRunResult RunResult, Compilation OutputCompilatio driver = driver.RunGeneratorsAndUpdateCompilation(compilation, out var outputCompilation, out _); return (driver.GetRunResult(), outputCompilation); } + + private static (MetadataReference Reference, byte[] Image) EmitReference(CSharpCompilation compilation) + { + var image = EmitAssembly(compilation); + return (MetadataReference.CreateFromImage(image), image); + } + + private static byte[] EmitAssembly(CSharpCompilation compilation) + { + using var stream = new MemoryStream(); + var result = compilation.Emit(stream); + result.Success.ShouldBeTrue(string.Join(Environment.NewLine, result.Diagnostics)); + return stream.ToArray(); + } } From 19de34d44d1597d2635abf714c2e04f983a77d28 Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 13:56:04 +0200 Subject: [PATCH 16/21] Ensure init --- .../ContainerRegistryTests.cs | 19 +++++++++++++------ 1 file changed, 13 insertions(+), 6 deletions(-) diff --git a/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs b/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs index 347a01d..f08435d 100644 --- a/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs +++ b/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs @@ -1,3 +1,4 @@ +using System.Runtime.CompilerServices; using FactoryGenerator; using Inherited; using Inheritor.Generated; @@ -10,17 +11,16 @@ public class ContainerRegistryTests [Test] public void ContainerEntryPointRegistersOnModuleLoad() { - // Force the Inheritor assembly to load, triggering its ModuleInitializer - _ = typeof(ContainerEntryPoint); + EnsureContainerEntryPointModuleInitialized(); - // The Inheritor assembly's ModuleInitializer should have already registered - // its container factory in ContainerRegistry when the assembly was loaded. ContainerRegistry.RegisteredAssemblies.ShouldContain("Inheritor"); } [Test] public void ContainerEntryPointCreateBuildsWorkingContainer() { + EnsureContainerEntryPointModuleInitialized(); + // Create a base container var baseContainer = new DependencyInjectionContainer(default, default, new NonInjectedClass()); @@ -34,6 +34,8 @@ public void ContainerEntryPointCreateBuildsWorkingContainer() [Test] public void BuildChainCreatesWorkingContainerPipeline() { + EnsureContainerEntryPointModuleInitialized(); + // Create a base container var baseContainer = new DependencyInjectionContainer(default, default, new NonInjectedClass()); @@ -48,7 +50,7 @@ public void BuildChainCreatesWorkingContainerPipeline() [Test] public void BuildChainWithoutAssemblyListSkipsCurrentContainerAssembly() { - _ = typeof(ContainerEntryPoint); + EnsureContainerEntryPointModuleInitialized(); var baseContainer = new DependencyInjectionContainer(default, default, new NonInjectedClass()); var final = ContainerRegistry.BuildChain(baseContainer); @@ -60,7 +62,7 @@ public void BuildChainWithoutAssemblyListSkipsCurrentContainerAssembly() [Test] public void BuildChainWithExplicitAssemblyListSkipsCurrentContainerAssembly() { - _ = typeof(ContainerEntryPoint); + EnsureContainerEntryPointModuleInitialized(); var baseContainer = new DependencyInjectionContainer(default, default, new NonInjectedClass()); var final = ContainerRegistry.BuildChain(baseContainer, new[] { "Inheritor" }); @@ -74,4 +76,9 @@ public void ContainerEntryPointAssemblyNameIsCorrect() { ContainerEntryPoint.AssemblyName.ShouldBe("Inheritor"); } + + private static void EnsureContainerEntryPointModuleInitialized() + { + RuntimeHelpers.RunModuleConstructor(typeof(ContainerEntryPoint).Module.ModuleHandle); + } } From 0e7b5aebc53f95a8687d2de61860ff442fbc2c95 Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 15:28:55 +0200 Subject: [PATCH 17/21] AsyncDisposable handling and cycles for methods --- Directory.Packages.props | 1 + .../FactoryGenerator.Attributes.csproj | 4 + FactoryGenerator.Attributes/IContainer.cs | 9 + .../LifetimeScopeDisposalExtensions.cs | 19 + .../ResolvedInstanceTracker.cs | 126 +++++ .../FactoryGeneratorMiddleware.cs | 3 +- .../FactoryGeneratorServiceProvider.cs | 12 +- .../ServiceProviderAdapter.cs | 20 +- FactoryGenerator/FactoryGenerator.cs | 490 +++++++++++------- FactoryGenerator/Injection.cs | 5 + FactoryGenerator/InjectionData.cs | 5 +- FactoryGenerator/SymbolUtility.cs | 4 +- README.md | 10 +- .../IntegrationTests.cs | 110 ++++ .../GeneratorBehaviorTests.cs | 309 ++++++++++- .../InjectionDetectionTests.cs | 87 ++++ .../ResolvedInstanceTrackerTests.cs | 102 ++++ Tests/TestData/Inherited/Types.cs | 43 ++ 18 files changed, 1173 insertions(+), 186 deletions(-) create mode 100644 FactoryGenerator.Attributes/LifetimeScopeDisposalExtensions.cs create mode 100644 FactoryGenerator.Attributes/ResolvedInstanceTracker.cs create mode 100644 Tests/FactoryGenerator.Tests/ResolvedInstanceTrackerTests.cs diff --git a/Directory.Packages.props b/Directory.Packages.props index 53c8a0a..0291d13 100644 --- a/Directory.Packages.props +++ b/Directory.Packages.props @@ -5,6 +5,7 @@
+ diff --git a/FactoryGenerator.Attributes/FactoryGenerator.Attributes.csproj b/FactoryGenerator.Attributes/FactoryGenerator.Attributes.csproj index 4c79949..b0de705 100644 --- a/FactoryGenerator.Attributes/FactoryGenerator.Attributes.csproj +++ b/FactoryGenerator.Attributes/FactoryGenerator.Attributes.csproj @@ -6,4 +6,8 @@ latest + + + + diff --git a/FactoryGenerator.Attributes/IContainer.cs b/FactoryGenerator.Attributes/IContainer.cs index 8b5e9b2..b0a8f33 100644 --- a/FactoryGenerator.Attributes/IContainer.cs +++ b/FactoryGenerator.Attributes/IContainer.cs @@ -39,4 +39,13 @@ public interface IContainerRegistrationMetadata public interface IContainerCacheInvalidator { void InvalidateCollectionCaches(); +} + +public interface IContainerLocalCollectionResolver +{ + bool TryResolveLocalCollection(Type type, out object? resolved); +} + +public interface IServiceProviderBackedContainer +{ } \ No newline at end of file diff --git a/FactoryGenerator.Attributes/LifetimeScopeDisposalExtensions.cs b/FactoryGenerator.Attributes/LifetimeScopeDisposalExtensions.cs new file mode 100644 index 0000000..3945d42 --- /dev/null +++ b/FactoryGenerator.Attributes/LifetimeScopeDisposalExtensions.cs @@ -0,0 +1,19 @@ +using System; +using System.Threading.Tasks; + +namespace FactoryGenerator; + +public static class LifetimeScopeDisposalExtensions +{ + public static ValueTask DisposeAsync(this ILifetimeScope scope) + { + if (scope is null) + throw new ArgumentNullException(nameof(scope)); + + if (scope is IAsyncDisposable asyncDisposable) + return asyncDisposable.DisposeAsync(); + + scope.Dispose(); + return default; + } +} diff --git a/FactoryGenerator.Attributes/ResolvedInstanceTracker.cs b/FactoryGenerator.Attributes/ResolvedInstanceTracker.cs new file mode 100644 index 0000000..d5d65ea --- /dev/null +++ b/FactoryGenerator.Attributes/ResolvedInstanceTracker.cs @@ -0,0 +1,126 @@ +using System; +using System.Collections.Generic; +using System.Threading.Tasks; + +namespace FactoryGenerator; + +#nullable enable + +public sealed class ResolvedInstanceTracker : IDisposable, IAsyncDisposable +{ + private enum DisposalMode + { + Active = 0, + Synchronous = 1, + Asynchronous = 2 + } + + private readonly object m_lock = new object(); + private List>? m_instances = new List>(); + private DisposalMode m_disposalMode; + + public void Track(object? instance) + { + if (instance is null) + return; + + if (instance is not IDisposable && instance is not IAsyncDisposable) + return; + + DisposalMode disposalMode; + lock (m_lock) + { + disposalMode = m_disposalMode; + if (disposalMode == DisposalMode.Active) + { + m_instances!.Add(new WeakReference(instance)); + return; + } + } + + if (disposalMode == DisposalMode.Asynchronous) + { + DisposeAsynchronously(instance).AsTask().GetAwaiter().GetResult(); + return; + } + + DisposeSynchronously(instance); + } + + public void Dispose() + { + var trackedInstances = BeginSynchronousDisposal(); + if (trackedInstances is null) + return; + + for (var index = trackedInstances.Count - 1; index >= 0; index--) + { + if (trackedInstances[index].TryGetTarget(out var instance)) + DisposeSynchronously(instance); + } + } + + public async ValueTask DisposeAsync() + { + var trackedInstances = BeginAsynchronousDisposal(); + if (trackedInstances is null) + return; + + for (var index = trackedInstances.Count - 1; index >= 0; index--) + { + if (trackedInstances[index].TryGetTarget(out var instance)) + await DisposeAsynchronously(instance).ConfigureAwait(false); + } + } + + private List>? BeginSynchronousDisposal() + { + lock (m_lock) + { + if (m_disposalMode != DisposalMode.Active) + return null; + + m_disposalMode = DisposalMode.Synchronous; + var instances = m_instances; + m_instances = null; + return instances; + } + } + + private List>? BeginAsynchronousDisposal() + { + lock (m_lock) + { + if (m_disposalMode != DisposalMode.Active) + return null; + + m_disposalMode = DisposalMode.Asynchronous; + var instances = m_instances; + m_instances = null; + return instances; + } + } + + private static void DisposeSynchronously(object instance) + { + if (instance is IDisposable disposable) + { + disposable.Dispose(); + return; + } + + if (instance is IAsyncDisposable asyncDisposable) + asyncDisposable.DisposeAsync().AsTask().GetAwaiter().GetResult(); + } + + private static ValueTask DisposeAsynchronously(object instance) + { + if (instance is IAsyncDisposable asyncDisposable) + return asyncDisposable.DisposeAsync(); + + if (instance is IDisposable disposable) + disposable.Dispose(); + + return default; + } +} diff --git a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs index 9f328d5..fe33c4d 100644 --- a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs +++ b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs @@ -29,10 +29,9 @@ public async Task Invoke(HttpContext context) } finally { - // The FactoryGeneratorServiceProvider.Dispose will dispose the scope if (context.RequestServices is FactoryGeneratorServiceProvider wrapper) { - wrapper.Dispose(); + await wrapper.DisposeAsync(); } context.RequestServices = originalProvider; } diff --git a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorServiceProvider.cs b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorServiceProvider.cs index 826282b..407ab6b 100644 --- a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorServiceProvider.cs +++ b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorServiceProvider.cs @@ -1,10 +1,11 @@ using System; +using System.Threading.Tasks; using Microsoft.Extensions.DependencyInjection; namespace FactoryGenerator.Extensions.AspNetCore; #nullable enable -internal sealed class FactoryGeneratorServiceProvider : IServiceProvider, ISupportRequiredService, IDisposable +internal sealed class FactoryGeneratorServiceProvider : IServiceProvider, ISupportRequiredService, IDisposable, IAsyncDisposable { private readonly IServiceProvider _baseProvider; private readonly ILifetimeScope _scope; @@ -50,4 +51,13 @@ public void Dispose() { _scope.Dispose(); } + + public ValueTask DisposeAsync() + { + if (_scope is IAsyncDisposable asyncDisposable) + return asyncDisposable.DisposeAsync(); + + _scope.Dispose(); + return default; + } } diff --git a/FactoryGenerator.Extensions.AspNetCore/ServiceProviderAdapter.cs b/FactoryGenerator.Extensions.AspNetCore/ServiceProviderAdapter.cs index 7c4aba4..174168a 100644 --- a/FactoryGenerator.Extensions.AspNetCore/ServiceProviderAdapter.cs +++ b/FactoryGenerator.Extensions.AspNetCore/ServiceProviderAdapter.cs @@ -1,11 +1,12 @@ using System; using System.Collections.Generic; +using System.Threading.Tasks; using Microsoft.Extensions.DependencyInjection; namespace FactoryGenerator.Extensions.AspNetCore; #nullable enable -internal sealed class ServiceProviderAdapter : IContainer, IDisposable +internal sealed class ServiceProviderAdapter : IContainer, IDisposable, IAsyncDisposable, IContainerLocalCollectionResolver, IServiceProviderBackedContainer { private readonly IServiceProvider _serviceProvider; private readonly IServiceScope? _serviceScope; @@ -26,6 +27,15 @@ public void Dispose() _serviceScope?.Dispose(); } + public ValueTask DisposeAsync() + { + if (_serviceScope is IAsyncDisposable asyncDisposable) + return asyncDisposable.DisposeAsync(); + + _serviceScope?.Dispose(); + return default; + } + public T Resolve() { var service = _serviceProvider.GetService(); @@ -58,6 +68,12 @@ public bool TryResolve(out T? resolved) return false; } + public bool TryResolveLocalCollection(Type type, out object? resolved) + { + resolved = _serviceProvider.GetService(type); + return resolved is not null; + } + public bool IsRegistered(Type type) { // IServiceProvider doesn't have a reliable IsRegistered method without resolution. @@ -80,7 +96,7 @@ public bool IsRegistered(Type type) public ILifetimeScope BeginLifetimeScope() { - var scope = _serviceProvider.CreateScope(); + var scope = _serviceProvider.CreateAsyncScope(); return new ServiceProviderAdapter(scope.ServiceProvider, scope, _baseContainer); } } diff --git a/FactoryGenerator/FactoryGenerator.cs b/FactoryGenerator/FactoryGenerator.cs index 66810bb..16ab857 100644 --- a/FactoryGenerator/FactoryGenerator.cs +++ b/FactoryGenerator/FactoryGenerator.cs @@ -4,6 +4,7 @@ using System.Linq; using System.Text; using System.Threading; +using System.Threading.Tasks; using Microsoft.CodeAnalysis; using Microsoft.CodeAnalysis.CSharp; using Microsoft.CodeAnalysis.Diagnostics; @@ -135,6 +136,7 @@ private static IEnumerable GenerateCode(ImmutableArray da using System.Linq; using System.Collections.Generic; using System.Collections.Immutable; +using System.Threading.Tasks; using FactoryGenerator; using System.CodeDom.Compiler; namespace {compilation.Assembly.Name}.Generated; @@ -144,7 +146,7 @@ namespace {compilation.Assembly.Name}.Generated; [GeneratedCode(""{ToolName}"", ""{Version}"")] #nullable enable #pragma warning disable CS0169, CS0414 -public sealed partial class {ClassName} : IContainer, IContainerScopeFactory, IContainerRegistrationMetadata, IContainerCacheInvalidator +public sealed partial class {ClassName} : IContainer, IContainerScopeFactory, IContainerRegistrationMetadata, IContainerCacheInvalidator, IAsyncDisposable, IContainerLocalCollectionResolver {{ #pragma warning restore CS0169, CS0414 @@ -226,15 +228,22 @@ private void InvalidateCollectionCachesInChain() public IContainer? Inheritor {{ get; set; }} internal readonly object m_lock = new(); private Dictionary> m_lookup; + private Dictionary> m_localCollectionLookup; private Dictionary m_booleans; - private List>? resolvedInstances; + private readonly ResolvedInstanceTracker m_resolvedInstances = new(); - internal List> GetResolvedInstances() + internal void TrackResolvedInstance(object instance) => m_resolvedInstances.Track(instance); + + public bool TryResolveLocalCollection(Type type, out object? resolved) {{ - if (resolvedInstances is null) - lock (m_lock) - resolvedInstances ??= new List>(); - return resolvedInstances; + if (m_localCollectionLookup.TryGetValue(type, out var factory)) + {{ + resolved = factory(); + return true; + }} + + resolved = default; + return false; }} public T Resolve() @@ -258,17 +267,13 @@ public object Resolve(Type type) public void Dispose() {{ DetachFromBase(); - if (resolvedInstances is not null) - {{ - foreach (var weakReference in resolvedInstances) - {{ - if(weakReference.TryGetTarget(out var disposable)) - {{ - disposable.Dispose(); - }} - }} - resolvedInstances.Clear(); - }} + m_resolvedInstances.Dispose(); + }} + + public ValueTask DisposeAsync() + {{ + DetachFromBase(); + return m_resolvedInstances.DisposeAsync(); }} public bool TryResolve(Type type, out object? resolved) @@ -329,7 +334,7 @@ public bool GetBoolean(string key) declarations[injection.Name] = Declaration(injection, availableInterfaceFullNames, false); scopedDeclarations[injection.Name] = Declaration(injection, availableInterfaceFullNames, true); - var missing = GetBestConstructorMissing(injection, availableInterfaceFullNames); + var missing = GetInjectionMissingParameters(injection, availableInterfaceFullNames); foreach (var param in missing) { var key = param.TypeFullName + " " + param.Name; @@ -360,6 +365,7 @@ public bool GetBoolean(string key) var booleanReservedNames = constructorParameters.Select(parameter => parameter.Name) .Concat(localizedParameters.Select(parameter => "coll_" + parameter.CollectionElementMemberName!)) + .Concat(interfaceMemberNames.Values.Select(name => "local_coll_" + name)) .Concat(localizedParameters.Select(parameter => "m_coll_" + parameter.CollectionElementMemberName!)) .Concat(interfaceMemberNames.Values) .Concat(ordered.Select(injection => injection.Name.Replace("()", string.Empty))) @@ -373,12 +379,15 @@ public bool GetBoolean(string key) "Dispose", "Resolve", "TryResolve", + "TryResolveLocalCollection", "IsRegistered", "GetBoolean", "GetBooleans", "BeginLifetimeScope", - "GetResolvedInstances", - "resolvedInstances", + "DisposeAsync", + "TrackResolvedInstance", + "m_resolvedInstances", + "m_localCollectionLookup", "m_lock", "m_lookup", "m_booleans", @@ -410,20 +419,11 @@ public bool GetBoolean(string key) } else { - var keys = possibilities.Select(i => i.BooleanInjection?.Key).OfType().Distinct().Reverse().ToArray(); - var fallback = possibilities.LastOrDefault(p => p.BooleanInjection == null); - var last = keys.Last(); - var ternary = new StringBuilder(); - foreach (var key in keys) - { - var keyIdentifier = booleanIdentifiers[key]; - var trueValue = possibilities.LastOrDefault(p => - p.BooleanInjection?.Value == true && p.BooleanInjection?.Key == key); - trueValue ??= fallback; - ternary.Append(key == last - ? $"{keyIdentifier} ? {trueValue?.Name ?? "null!"} : {fallback?.Name ?? "null!"}" - : $"{keyIdentifier} ? {trueValue?.Name ?? "null!"} : "); - } + var ternary = BuildBooleanSelectionExpression( + ifaceFull, + possibilities, + booleanIdentifiers, + possibility => possibility.Name); if (!declarations.ContainsKey(ifaceMethodName)) { @@ -486,12 +486,15 @@ public bool GetBoolean(string key) .Select(key => (TypeName: $"System.Collections.Generic.IEnumerable<{key}>", Expression: "coll_" + interfaceMemberNames[key])) .Where(pair => !localizedTypes.Contains(pair.TypeName)) .ToList(); + var localCollectionPairs = interfaceInjectors.Keys + .Select(key => (TypeName: $"System.Collections.Generic.IEnumerable<{key}>", Expression: "local_coll_" + interfaceMemberNames[key])) + .ToList(); var constructorPairs = DistinctByTypeName(externalParameters.Select(p => (TypeName: p.TypeFullName, Expression: p.Name)).ToList(), pair => pair.TypeName); var dictSize = interfacePairs.Count + localizedPairs.Count + enumerablePairs.Count + constructorPairs.Count; yield return Constructor(usingStatements, constructorFields, constructor, constructorAssignments, - dictSize, interfacePairs, localizedPairs, enumerablePairs, constructorPairs, + dictSize, interfacePairs, localizedPairs, enumerablePairs, constructorPairs, localCollectionPairs, true, ClassName, lifetimeInvocationValues: lifetimeParameterValues, resolvingConstructorAssignments: resolvedConstructorAssignments, booleans: booleanParameters); yield return Declarations(usingStatements, declarations, ClassName); @@ -500,7 +503,7 @@ public bool GetBoolean(string key) [GeneratedCode(""{ToolName}"", ""{Version}"")] #nullable enable #pragma warning disable CS0169, CS0414 -public sealed partial class LifetimeScope : IContainer, IContainerScopeFactory, IContainerRegistrationMetadata, IContainerCacheInvalidator +public sealed partial class LifetimeScope : IContainer, IContainerScopeFactory, IContainerRegistrationMetadata, IContainerCacheInvalidator, IAsyncDisposable, IContainerLocalCollectionResolver {{ #pragma warning restore CS0169, CS0414 private IContainer GetRoot() @@ -588,22 +591,28 @@ public ILifetimeScope BeginLifetimeScope() public ILifetimeScope BeginLifetimeScope(IContainer? baseContainer) {{ var scope = m_fallback.BeginLifetimeScope(baseContainer); - GetResolvedInstances().Add(new WeakReference(scope)); + TrackResolvedInstance(scope); return scope; }} internal readonly object m_lock = new(); private {ClassName} m_fallback; - private IContainer? m_ownedBase; private Dictionary> m_lookup; + private Dictionary> m_localCollectionLookup; private Dictionary m_booleans; - private List>? resolvedInstances; + private readonly ResolvedInstanceTracker m_resolvedInstances = new(); - internal List> GetResolvedInstances() + internal void TrackResolvedInstance(object instance) => m_resolvedInstances.Track(instance); + + public bool TryResolveLocalCollection(Type type, out object? resolved) {{ - if (resolvedInstances is null) - lock (m_lock) - resolvedInstances ??= new List>(); - return resolvedInstances; + if (m_localCollectionLookup.TryGetValue(type, out var factory)) + {{ + resolved = factory(); + return true; + }} + + resolved = default; + return false; }} public T Resolve() @@ -627,20 +636,13 @@ public object Resolve(Type type) public void Dispose() {{ DetachFromBase(); - if (resolvedInstances is not null) - {{ - foreach (var weakReference in resolvedInstances) - {{ - if(weakReference.TryGetTarget(out var disposable)) - {{ - disposable.Dispose(); - }} - }} - resolvedInstances.Clear(); - }} - var ownedBase = m_ownedBase; - m_ownedBase = null; - ownedBase?.Dispose(); + m_resolvedInstances.Dispose(); + }} + + public ValueTask DisposeAsync() + {{ + DetachFromBase(); + return m_resolvedInstances.DisposeAsync(); }} public bool TryResolve(Type type, out object? resolved) @@ -689,7 +691,7 @@ public bool GetBoolean(string key) "; yield return Constructor(usingStatements, constructorFields, lifetimeConstructor, constructorAssignments, - dictSize, interfacePairs, localizedPairs, enumerablePairs, constructorPairs, + dictSize, interfacePairs, localizedPairs, enumerablePairs, constructorPairs, localCollectionPairs, false, LifetimeName, resolvingConstructorAssignments: resolvedConstructorAssignments, addMergingConstructor: false, booleans: booleanParameters); yield return Declarations(usingStatements, scopedDeclarations, LifetimeName); @@ -853,8 +855,19 @@ private static List GetReachableImplementations(List GetCycleDependencies(InjectionData injection, ImmutableArray availableInterfaceFullNames) { - if (injection.Lambda is not null) + if (injection.Lambda is LambdaData lambda) + { + if (availableInterfaceFullNames.Contains(lambda.ContainingTypeFullName)) + yield return lambda.ContainingTypeFullName; + + if (!lambda.IsMethod) + yield break; + + foreach (var dependency in GetParameterDependencies(lambda.MethodParameters, availableInterfaceFullNames)) + yield return dependency; + yield break; + } HashSet? missing = null; HashSet? nullableDefaults = null; @@ -862,7 +875,15 @@ private static IEnumerable GetCycleDependencies(InjectionData injection, if (ctor is null) yield break; - foreach (var parameter in ctor.Parameters) + foreach (var dependency in GetParameterDependencies(ctor.Parameters, availableInterfaceFullNames)) + yield return dependency; + } + + private static IEnumerable GetParameterDependencies( + ImmutableArray parameters, + ImmutableArray availableInterfaceFullNames) + { + foreach (var parameter in parameters) { if (parameter.IsCollection) continue; @@ -958,6 +979,47 @@ private static Dictionary BuildExternalParameterIdentifiers(IEnu return identifiers; } + private static string BuildBooleanSelectionExpression( + string typeFullName, + IReadOnlyList possibilities, + IReadOnlyDictionary booleanIdentifiers, + Func implementationExpression) + { + if (possibilities.All(possibility => possibility.BooleanInjection is null)) + return implementationExpression(possibilities.Last()); + + var keys = possibilities.Select(possibility => possibility.BooleanInjection?.Key) + .OfType() + .Distinct() + .Reverse() + .ToArray(); + + var fallback = possibilities.LastOrDefault(possibility => possibility.BooleanInjection is null); + var fallbackExpression = fallback is not null + ? implementationExpression(fallback) + : BuildMissingImplementationExpression(typeFullName); + + if (keys.Length == 0) + return fallbackExpression; + + var last = keys.Last(); + var selection = new StringBuilder(); + foreach (var key in keys) + { + var selected = possibilities.LastOrDefault(possibility => + possibility.BooleanInjection?.Value == true && possibility.BooleanInjection?.Key == key) ?? fallback; + var selectedExpression = selected is not null + ? implementationExpression(selected) + : BuildMissingImplementationExpression(typeFullName); + + selection.Append(key == last + ? $"{booleanIdentifiers[key]} ? {selectedExpression} : {fallbackExpression}" + : $"{booleanIdentifiers[key]} ? {selectedExpression} : "); + } + + return selection.ToString(); + } + private static string GetBooleanParameterIdentifier(string booleanKey) { return GetSanitizedIdentifier(booleanKey, "boolean_"); @@ -1054,6 +1116,7 @@ private static void DfsCycleCheck(string node, Dictionary interfaceTypePairs, IEnumerable<(string TypeName, string Expression)> localizedParamPairs, IEnumerable<(string TypeName, string Expression)> enumerablePairs, IEnumerable<(string TypeName, string Expression)> constructorParamPairs, + IEnumerable<(string TypeName, string Expression)> localCollectionPairs, bool addLifetimeScopeFunction, string className, string? lifetimeInvocationValues = null, string? fromConstructor = null, string? resolvingConstructorAssignments = null, bool addMergingConstructor = true, IReadOnlyList<(string Key, string Identifier)> booleans = null!) @@ -1068,7 +1131,7 @@ public ILifetimeScope BeginLifetimeScope() public ILifetimeScope BeginLifetimeScope(IContainer? baseContainer) {{ var scope = new {LifetimeName}({lifetimeInvocationValues}); - GetResolvedInstances().Add(new WeakReference(scope)); + TrackResolvedInstance(scope); return scope; }}" : string.Empty; @@ -1086,6 +1149,9 @@ public ILifetimeScope BeginLifetimeScope(IContainer? baseContainer) {MakeDictionaryFromParams(localizedParamPairs)} {MakeDictionaryFromParams(enumerablePairs)} {MakeDictionaryFromParams(constructorParamPairs)} + }}; + m_localCollectionLookup = new({localCollectionPairs.Count()}) {{ +{MakeDictionaryFromParams(localCollectionPairs)} }}; m_booleans = new(); foreach(var (key, value) in Base.GetBooleans()) @@ -1097,9 +1163,11 @@ public ILifetimeScope BeginLifetimeScope(IContainer? baseContainer) var extraConstruction = addLifetimeScopeFunction ? string.Empty : @"m_fallback = fallback; this.Base = baseContainer; - m_ownedBase = baseContainer; if (baseContainer is not null) - AttachToBase(baseContainer);"; + { + AttachToBase(baseContainer); + TrackResolvedInstance(baseContainer); + }"; return $@"{usingStatements} public partial class {className} {{ @@ -1115,6 +1183,9 @@ public partial class {className} {MakeDictionaryFromParams(enumerablePairs)} {MakeDictionaryFromParams(constructorParamPairs)} }}; + m_localCollectionLookup = new({localCollectionPairs.Count()}) {{ +{MakeDictionaryFromParams(localCollectionPairs)} + }}; m_booleans = new({booleans.Count}) {{ {string.Join("\n", booleans.Select(boolean => $"\t\t{{ \"{boolean.Key}\", {boolean.Identifier} }},"))} @@ -1155,39 +1226,73 @@ private static void MakeArray(Dictionary declarations, string na string elementTypeFullName, Dictionary> interfaceInjectors, IReadOnlyDictionary booleanIdentifiers) { - var factoryName = $"new {elementTypeFullName}[0]"; - var factory = string.Empty; - if (interfaceInjectors.TryGetValue(elementTypeFullName, out var injections)) - { - factoryName = $"Create{name}()".Replace("_", ""); - var nonBooleanInjections = injections.Where(i => i.BooleanInjection == null).ToList(); - var booleanInjections = injections.Where(b => b.BooleanInjection != null).ToList(); - factory = @$" - private bool Reentrant_{name}; - IEnumerable<{elementTypeFullName}> {factoryName} + var injections = interfaceInjectors.TryGetValue(elementTypeFullName, out var foundInjections) + ? foundInjections + : new List(); + var localFactoryName = $"CreateLocal{name}()".Replace("_", ""); + var factoryName = $"Create{name}()".Replace("_", ""); + var nonBooleanInjections = injections.Where(i => i.BooleanInjection == null).ToList(); + var booleanInjections = injections.Where(b => b.BooleanInjection != null).ToList(); + var factory = @$" + IEnumerable<{elementTypeFullName}> {localFactoryName} {{ - if(Reentrant_{name}) return Array.Empty<{elementTypeFullName}>(); - Reentrant_{name} = true; var source = new List<{elementTypeFullName}>({nonBooleanInjections.Count}) {{ {string.Join(",\n\t\t\t", nonBooleanInjections.Select(i => i.Name))} }}; {string.Join("\n\t\t\t", booleanInjections.Select(i => $"if({booleanIdentifiers[i.BooleanInjection!.Key]}) source.Add({i.Name});"))} + return source; + }} + private bool Reentrant_{name}; + IEnumerable<{elementTypeFullName}> {factoryName} + {{ + if(Reentrant_{name}) return Array.Empty<{elementTypeFullName}>(); + Reentrant_{name} = true; + var source = new List<{elementTypeFullName}>({localFactoryName}); var b = Base; + var frameworkCollectionSourceSeen = false; while(b is not null) {{ - if(b.TryResolve>(out var additional)) source.AddRange(additional!); + if(!(frameworkCollectionSourceSeen && b is IServiceProviderBackedContainer)) + {{ + if (b is IContainerLocalCollectionResolver localResolver) + {{ + if(localResolver.TryResolveLocalCollection(typeof(IEnumerable<{elementTypeFullName}>), out var localAdditional)) + source.AddRange((IEnumerable<{elementTypeFullName}>)localAdditional!); + }} + else if(b.TryResolve>(out var additional)) + {{ + source.AddRange(additional!); + }} + + if (b is IServiceProviderBackedContainer) + frameworkCollectionSourceSeen = true; + }} b = b.Base; }} + var inheritorFrameworkCollectionSourceSeen = false; b = Inheritor; while(b is not null) {{ - if(b.TryResolve>(out var additional)) source.AddRange(additional!); + if(!(inheritorFrameworkCollectionSourceSeen && b is IServiceProviderBackedContainer)) + {{ + if (b is IContainerLocalCollectionResolver localResolver) + {{ + if(localResolver.TryResolveLocalCollection(typeof(IEnumerable<{elementTypeFullName}>), out var localAdditional)) + source.AddRange((IEnumerable<{elementTypeFullName}>)localAdditional!); + }} + else if(b.TryResolve>(out var additional)) + {{ + source.AddRange(additional!); + }} + + if (b is IServiceProviderBackedContainer) + inheritorFrameworkCollectionSourceSeen = true; + }} b = b.Inheritor; }} Reentrant_{name} = false; return source; }}"; - } declarations[name] = $@" internal IEnumerable<{elementTypeFullName}> {name} {{ @@ -1206,6 +1311,7 @@ internal IEnumerable<{elementTypeFullName}> {name} }} }} }} + internal IEnumerable<{elementTypeFullName}> local_{name} => {localFactoryName}; internal IEnumerable<{elementTypeFullName}>? m_{name};" + factory; } @@ -1235,9 +1341,9 @@ private static string Declaration(InjectionData injection, ImmutableArray m_fallback.{name};"; if (injection.Singleton || injection.Scoped) - return SymbolUtility.SingletonFactory(injection.TypeFullName, name, lazyName, creation, injection.Disposable); + return SymbolUtility.SingletonFactory(injection.TypeFullName, name, lazyName, creation, injection.Disposable || injection.AsyncDisposable); - if (injection.Disposable) + if (injection.Disposable || injection.AsyncDisposable) return SymbolUtility.DisposableFactory(injection.TypeFullName, name, creation); return $"internal {injection.TypeFullName} {name} => {creation};"; @@ -1251,22 +1357,13 @@ private static string CreationCall(InjectionData injection, ImmutableArray? lambdaNullableDefaults = null; - if (lambda.IsMethod && lambda.MethodParameters.Length > 0) + if (lambda.IsMethod) { - foreach (var p in lambda.MethodParameters) - { - if (!p.IsNullable) continue; - var baseType = p.TypeFullName.TrimEnd('?'); - if (availableInterfaceFullNames.Contains(baseType)) continue; - lambdaNullableDefaults ??= new HashSet(); - lambdaNullableDefaults.Add(p); - } + HashSet? lambdaMissing = null; + HashSet? lambdaNullableDefaults = null; + AnalyzeParameters(lambda.MethodParameters, availableInterfaceFullNames, ref lambdaMissing, ref lambdaNullableDefaults); + return $"{lambda.ContainingTypeMemberName}.{lambda.MemberName}{MakeMethodCall(lambda.MethodParameters, lambdaMissing, lambdaNullableDefaults)}"; } - - if (lambda.IsMethod) - return $"{lambda.ContainingTypeMemberName}.{lambda.MemberName}{MakeMethodCall(lambda.MethodParameters, null, lambdaNullableDefaults)}"; else return $"{lambda.ContainingTypeMemberName}.{lambda.MemberName}"; } @@ -1289,96 +1386,137 @@ private static string CreationCall(InjectionData injection, ImmutableArray(); - var localNullableDefaults = new HashSet(); - foreach (var parameter in ctor.Parameters) - { - // For nullable params, check availability of the underlying non-nullable type - var typeLookup = parameter.IsNullable - ? parameter.TypeFullName.TrimEnd('?') - : parameter.TypeFullName; - if (availableInterfaceFullNames.Contains(typeLookup)) continue; - if (parameter.HasExplicitDefault) continue; - if (parameter.IsParams) continue; - // Collection params (IEnumerable, T[], List, etc.) are always satisfiable - // via MakeArray – add to missing for factory generation but keep constructor valid - if (parameter.IsCollection) - { - localMissing.Add(parameter); - continue; - } - // Nullable reference/value params that aren't registered default to null - if (parameter.IsNullable) - { - localNullableDefaults.Add(parameter); - continue; - } - valid = false; - localMissing.Add(parameter); - } + HashSet? localMissing = null; + HashSet? localNullableDefaults = null; + AnalyzeParameters(ctor.Parameters, availableInterfaceFullNames, ref localMissing, ref localNullableDefaults, out var valid); if (valid) { chosen = ctor; - missing = localMissing.Count > 0 ? localMissing : null; - nullableDefaults = localNullableDefaults.Count > 0 ? localNullableDefaults : null; + missing = localMissing; + nullableDefaults = localNullableDefaults; break; } - if ((missing?.Count ?? int.MaxValue) <= localMissing.Count) continue; + if ((missing?.Count ?? int.MaxValue) <= (localMissing?.Count ?? 0)) continue; chosen = ctor; missing = localMissing; - nullableDefaults = localNullableDefaults.Count > 0 ? localNullableDefaults : null; + nullableDefaults = localNullableDefaults; } return chosen; } - private static IEnumerable GetBestConstructorMissing(InjectionData injection, + private static IEnumerable GetInjectionMissingParameters(InjectionData injection, ImmutableArray availableInterfaceFullNames) { + if (injection.Lambda is LambdaData lambda) + { + if (!lambda.IsMethod) + return Enumerable.Empty(); + + HashSet? lambdaMissing = null; + HashSet? lambdaNullableDefaults = null; + AnalyzeParameters(lambda.MethodParameters, availableInterfaceFullNames, ref lambdaMissing, ref lambdaNullableDefaults); + return lambdaMissing ?? Enumerable.Empty(); + } + HashSet? missing = null; HashSet? nullableDefaults = null; GetBestConstructor(injection, availableInterfaceFullNames, ref missing, ref nullableDefaults); return missing ?? Enumerable.Empty(); } - private static string MakeConstructorCall(ConstructorData ctor, HashSet? missing, HashSet? nullableDefaults) + private static void AnalyzeParameters( + ImmutableArray parameters, + ImmutableArray availableInterfaceFullNames, + ref HashSet? missing, + ref HashSet? nullableDefaults) { - var args = new List(); - foreach (var parameter in ctor.Parameters) + AnalyzeParameters(parameters, availableInterfaceFullNames, ref missing, ref nullableDefaults, out _); + } + + private static void AnalyzeParameters( + ImmutableArray parameters, + ImmutableArray availableInterfaceFullNames, + ref HashSet? missing, + ref HashSet? nullableDefaults, + out bool valid) + { + var localMissing = new HashSet(); + var localNullableDefaults = new HashSet(); + valid = true; + + foreach (var parameter in parameters) { - if (nullableDefaults?.Contains(parameter) == true) + var typeLookup = parameter.IsNullable + ? parameter.TypeFullName.TrimEnd('?') + : parameter.TypeFullName; + + if (parameter.IsCollection) { - args.Add("null"); + localMissing.Add(parameter); continue; } - if (missing?.Contains(parameter) == true) + + if (availableInterfaceFullNames.Contains(typeLookup)) + continue; + + if (parameter.HasExplicitDefault || parameter.IsParams) + continue; + + if (parameter.IsNullable) { - args.Add(CollectionConstructorArg(parameter)); + localNullableDefaults.Add(parameter); continue; } - args.Add(parameter.TypeMemberName + "()"); + + valid = false; + localMissing.Add(parameter); } - return $"({string.Join(", ", args)})"; + + missing = localMissing.Count > 0 ? localMissing : null; + nullableDefaults = localNullableDefaults.Count > 0 ? localNullableDefaults : null; + } + + private static string MakeConstructorCall(ConstructorData ctor, HashSet? missing, HashSet? nullableDefaults) + { + return MakeInvocationCall(ctor.Parameters, missing, nullableDefaults); } private static string MakeMethodCall(ImmutableArray parameters, HashSet? missing, HashSet? nullableDefaults = null) + { + return MakeInvocationCall(parameters, missing, nullableDefaults); + } + + private static string MakeInvocationCall( + ImmutableArray parameters, + HashSet? missing, + HashSet? nullableDefaults) { var args = new List(); + var useNamedArguments = false; foreach (var parameter in parameters) { if (nullableDefaults?.Contains(parameter) == true) { - args.Add("null"); + args.Add(useNamedArguments ? $"{parameter.Name}: null" : "null"); continue; } if (missing?.Contains(parameter) == true) { - args.Add(CollectionConstructorArg(parameter)); + var argument = CollectionConstructorArg(parameter); + args.Add(useNamedArguments ? $"{parameter.Name}: {argument}" : argument); continue; } - args.Add(parameter.TypeMemberName + "()"); + + if (parameter.HasExplicitDefault || parameter.IsParams) + { + useNamedArguments = true; + continue; + } + + var resolvedArgument = parameter.TypeMemberName + "()"; + args.Add(useNamedArguments ? $"{parameter.Name}: {resolvedArgument}" : resolvedArgument); } return $"({string.Join(", ", args)})"; } @@ -1842,10 +1980,11 @@ private static string BuildStaticResolveInjectionMethod( var helperName = GetStaticInjectionHelperName(injection); var parameterList = BuildStaticMethodParameterList(spec, booleanIdentifiers, externalIdentifiers, includeContainer: true, includeState: true); var createInvocation = BuildStaticCreateInjectionInvocation(spec, injection, booleanIdentifiers, externalIdentifiers); + var tracksResolvedInstance = injection.Disposable || injection.AsyncDisposable; if (injection.Singleton || injection.Scoped) { - if (injection.Disposable) + if (tracksResolvedInstance) { return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) {{ @@ -1862,7 +2001,7 @@ private static string BuildStaticResolveInjectionMethod( return cached; var value = {createInvocation}; - container.GetResolvedInstances().Add(new global::System.WeakReference(value)); + container.TrackResolvedInstance(value); container.{injection.LazyFieldName} = value; return value; }} @@ -1894,14 +2033,14 @@ private static string BuildStaticResolveInjectionMethod( }}"; } - if (injection.Disposable) + if (tracksResolvedInstance) { return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) {{ if (container is not null) {{ var value = {createInvocation}; - container.GetResolvedInstances().Add(new global::System.WeakReference(value)); + container.TrackResolvedInstance(value); return value; }} @@ -1946,42 +2085,50 @@ private static string BuildStaticCreateExpression( if (injection.Lambda is LambdaData lambda) { if (!specs.TryGetValue(lambda.ContainingTypeFullName, out var containingSpec)) - return BuildStaticMissingImplementationExpression(lambda.ContainingTypeFullName); + return BuildMissingImplementationExpression(lambda.ContainingTypeFullName); var containingInvocation = BuildStaticResolveInvocation(containingSpec, "ResolveCore", booleanIdentifiers, externalIdentifiers); if (!lambda.IsMethod) return $"{containingInvocation}.{lambda.MemberName}"; - var lambdaArguments = BuildStaticArgumentExpressions(lambda.MethodParameters, specs, booleanIdentifiers, externalIdentifiers); - return $"{containingInvocation}.{lambda.MemberName}({string.Join(", ", lambdaArguments)})"; + var lambdaArguments = BuildStaticArgumentList(lambda.MethodParameters, specs, booleanIdentifiers, externalIdentifiers); + return $"{containingInvocation}.{lambda.MemberName}({lambdaArguments})"; } HashSet? missing = null; HashSet? nullableDefaults = null; var constructor = GetBestConstructor(injection, availableInterfaces, ref missing, ref nullableDefaults); if (constructor is null) - return BuildStaticMissingImplementationExpression(injection.TypeFullName); + return BuildMissingImplementationExpression(injection.TypeFullName); - var constructorArguments = BuildStaticArgumentExpressions(constructor.Parameters, specs, booleanIdentifiers, externalIdentifiers); - return $"new {injection.TypeFullName}({string.Join(", ", constructorArguments)})"; + var constructorArguments = BuildStaticArgumentList(constructor.Parameters, specs, booleanIdentifiers, externalIdentifiers); + return $"new {injection.TypeFullName}({constructorArguments})"; } - private static List BuildStaticArgumentExpressions( + private static string BuildStaticArgumentList( ImmutableArray parameters, IReadOnlyDictionary specs, IReadOnlyDictionary booleanIdentifiers, IReadOnlyDictionary externalIdentifiers) { var arguments = new List(); + var useNamedArguments = false; foreach (var parameter in parameters) { var argumentExpression = BuildStaticArgumentExpression(parameter, specs, booleanIdentifiers, externalIdentifiers); - if (argumentExpression is not null) - arguments.Add(argumentExpression); + if (argumentExpression is null) + { + useNamedArguments = true; + continue; + } + + arguments.Add(useNamedArguments + ? $"{parameter.Name}: {argumentExpression}" + : argumentExpression); } - return arguments; + return string.Join(", ", arguments); } private static string? BuildStaticArgumentExpression( @@ -1994,7 +2141,11 @@ private static List BuildStaticArgumentExpressions( { if (!specs.TryGetValue(parameter.CollectionElementFullName, out var collectionSpec)) { - return BuildStaticCollectionConversion(parameter, $"global::System.Array.Empty<{parameter.CollectionElementFullName}>()"); + var resolvedCollectionName = "resolvedCollection_" + parameter.Name; + var fallbackCollection = $@"container != null && container.TryResolve>(out var {resolvedCollectionName}) + ? {resolvedCollectionName}! + : global::System.Array.Empty<{parameter.CollectionElementFullName}>()"; + return BuildStaticCollectionConversion(parameter, fallbackCollection); } return BuildStaticCollectionConversion( @@ -2018,7 +2169,7 @@ private static List BuildStaticArgumentExpressions( if (externalIdentifiers.TryGetValue(parameter.TypeFullName, out var identifier)) return identifier; - return BuildStaticMissingImplementationExpression(parameter.TypeFullName); + return BuildMissingImplementationExpression(parameter.TypeFullName); } private static string BuildStaticCollectionConversion(ParameterData parameter, string sourceExpression) @@ -2058,24 +2209,11 @@ private static string BuildStaticResolveSelectionExpression( IReadOnlyDictionary booleanIdentifiers, IReadOnlyDictionary externalIdentifiers) { - if (spec.Possibilities.All(possibility => possibility.BooleanInjection is null)) - return BuildStaticResolveInjectionInvocation(spec, spec.Possibilities.Last(), booleanIdentifiers, externalIdentifiers); - - var keys = spec.Possibilities.Select(possibility => possibility.BooleanInjection?.Key).OfType().Distinct().Reverse().ToArray(); - var fallback = spec.Possibilities.LastOrDefault(possibility => possibility.BooleanInjection is null); - var expression = fallback is not null - ? BuildStaticResolveInjectionInvocation(spec, fallback, booleanIdentifiers, externalIdentifiers) - : BuildStaticMissingImplementationExpression(spec.TypeFullName); - - foreach (var key in keys) - { - var selected = spec.Possibilities.LastOrDefault(possibility => - possibility.BooleanInjection?.Value == true && possibility.BooleanInjection?.Key == key) ?? fallback; - - expression = $"{booleanIdentifiers[key]} ? {(selected is not null ? BuildStaticResolveInjectionInvocation(spec, selected, booleanIdentifiers, externalIdentifiers) : BuildStaticMissingImplementationExpression(spec.TypeFullName))} : {expression}"; - } - - return expression; + return BuildBooleanSelectionExpression( + spec.TypeFullName, + spec.Possibilities, + booleanIdentifiers, + possibility => BuildStaticResolveInjectionInvocation(spec, possibility, booleanIdentifiers, externalIdentifiers)); } private static string BuildStaticResolveInjectionInvocation( @@ -2154,7 +2292,7 @@ private static string BuildStaticInternalInvocationArguments( return $"{containerExpression}, {stateExpression}, {runtimeValues}"; } - private static string BuildStaticMissingImplementationExpression(string typeFullName) + private static string BuildMissingImplementationExpression(string typeFullName) { return $"throw new global::System.InvalidOperationException(\"Cannot resolve {typeFullName} without a matching implementation\")"; } diff --git a/FactoryGenerator/Injection.cs b/FactoryGenerator/Injection.cs index b2c15b6..f9120fa 100644 --- a/FactoryGenerator/Injection.cs +++ b/FactoryGenerator/Injection.cs @@ -98,9 +98,13 @@ public static class Injection interfaces = interfaces.AddRange(attributedInterfaces); var isDisposable = namedTypeSymbol.AllInterfaces.Any(i => i.SpecialType == SpecialType.System_IDisposable); + var isAsyncDisposable = namedTypeSymbol.AllInterfaces.Any(i => i.ToString() == "System.IAsyncDisposable"); var disposableIface = interfaces.FirstOrDefault(i => i.SpecialType == SpecialType.System_IDisposable); if (disposableIface is not null) interfaces = interfaces.Remove(disposableIface); + var asyncDisposableIface = interfaces.FirstOrDefault(i => i.ToString() == "System.IAsyncDisposable"); + if (asyncDisposableIface is not null) + interfaces = interfaces.Remove(asyncDisposableIface); interfaces = interfaces .RemoveRange(preventedInterfaces) @@ -131,6 +135,7 @@ public static class Injection singleton: singleInstance, scoped: scoped, disposable: isDisposable, + asyncDisposable: isAsyncDisposable, booleanInjection: boolean, constructors: constructors, lambda: lambdaData); diff --git a/FactoryGenerator/InjectionData.cs b/FactoryGenerator/InjectionData.cs index b5f321d..d03ef48 100644 --- a/FactoryGenerator/InjectionData.cs +++ b/FactoryGenerator/InjectionData.cs @@ -17,6 +17,7 @@ public sealed class InjectionData : IEquatable public bool Singleton { get; } public bool Scoped { get; } public bool Disposable { get; } + public bool AsyncDisposable { get; } public BooleanInjection? BooleanInjection { get; } public ImmutableArray Constructors { get; } public LambdaData? Lambda { get; } @@ -28,7 +29,7 @@ public sealed class InjectionData : IEquatable public InjectionData( string typeFullName, string typeMemberName, string assemblyName, int assemblyPriority, ImmutableArray interfaceFullNames, ImmutableArray interfaceMemberNames, - bool singleton, bool scoped, bool disposable, + bool singleton, bool scoped, bool disposable, bool asyncDisposable, BooleanInjection? booleanInjection, ImmutableArray constructors, LambdaData? lambda) { @@ -41,6 +42,7 @@ public InjectionData( Singleton = singleton; Scoped = scoped; Disposable = disposable; + AsyncDisposable = asyncDisposable; BooleanInjection = booleanInjection; Constructors = constructors; Lambda = lambda; @@ -59,6 +61,7 @@ public bool Equals(InjectionData? other) && Singleton == other.Singleton && Scoped == other.Scoped && Disposable == other.Disposable + && AsyncDisposable == other.AsyncDisposable && Equals(BooleanInjection, other.BooleanInjection) && Constructors.SequenceEqual(other.Constructors) && Equals(Lambda, other.Lambda); diff --git a/FactoryGenerator/SymbolUtility.cs b/FactoryGenerator/SymbolUtility.cs index 3028e43..80ef8fb 100644 --- a/FactoryGenerator/SymbolUtility.cs +++ b/FactoryGenerator/SymbolUtility.cs @@ -118,7 +118,7 @@ public static string SingletonFactory(string typeName, string name, string lazyN if (cached != null) return cached; var value = {creation}; - GetResolvedInstances().Add(new WeakReference(value)); + TrackResolvedInstance(value); {lazyName} = value; return value; }} @@ -150,7 +150,7 @@ internal static string DisposableFactory(string typeName, string name, string cr internal {typeName} {name} {{ var value = {creationCall}; - GetResolvedInstances().Add(new WeakReference(value)); + TrackResolvedInstance(value); return value; }}"; } diff --git a/README.md b/README.md index eb17be1..d7bdd82 100644 --- a/README.md +++ b/README.md @@ -103,6 +103,8 @@ public class Program ``` Of note is perhaps `Generated.DependencyInjectionContainer`, this is the Compile-time created implementation of our IoC container, it implements the interface `FactoryGenerator.IContainer`. +Generated containers also implement `IAsyncDisposable`. If you are working with an `ILifetimeScope` or `IContainer` reference, `await scope.DisposeAsync()` is the preferred path, but a synchronous `Dispose()` will also block until async-only services finish disposing. + ### Attributes | Attribute | Description | Requires | @@ -118,6 +120,8 @@ Of note is perhaps `Generated.DependencyInjectionContainer`, this is the Compile `InjectionPriority(int)` is an **assembly-level** attribute rather than an injection attribute. +If every implementation of a service is guarded by `[Boolean(...)]` and no ungated fallback exists, resolving that service throws `InvalidOperationException` when none of the booleans select an implementation. + ### Overriding Overriding Injections, i.e if in the graph above _Dependency C_ injected an ```ISomething``` instance and that specific implementation of ```ISomething``` will not work for anything that uses _Dependency B_, then _Dependency B_ can substitute that injection by providing it's own injection of ```ISomething```. This overriding generally follows the project dependency tree, so if _Project A_ depends on _Project B_ which depends on _Project C_, A can override both B and C, but B cannot override A. @@ -183,6 +187,8 @@ public class Provider : IProvider ``` With this code, it is now possible to do `container.Resolve()`, which will effectively return the result of `new Provider().Method()`, although, since `Method` is `[Inject]`ed as a `[Singleton]`, the result will be cached and the same instance will be returned at every call to `Resolve` as well as shared between all Injected implementations that require a `IResultType`. +Injected method parameters follow the same rules as constructor parameters: unresolved external values are surfaced on the generated container constructor, optional parameters keep their declared defaults, and `params` collections are supplied from the container when possible. + ### Static Extensions (C# 14 / .NET 10+) When targeting C# 14 or later, FactoryGenerator automatically emits [static extension methods](https://learn.microsoft.com/en-us/dotnet/csharp/whats-new/csharp-14#extension-members) for every registered interface. This provides a dictionary-free, inline resolution path that the JIT can aggressively optimize. @@ -209,7 +215,7 @@ var fresh = ISingleton.Resolve(null); var switched = ISwitchableInterface.Resolve(testBool: true); ``` -Collection dependencies (`IEnumerable`, arrays, `List`, `ImmutableArray`, `ReadOnlySpan`) are resolved through the same generated static pipeline, so the extension path and the normal container path construct equivalent object graphs. Direct `Resolve>()` calls are generated for every registered service type, even when no constructor or source usage referenced that collection shape ahead of time. +Collection dependencies (`IEnumerable`, arrays, `List`, `ImmutableArray`, `ReadOnlySpan`) are resolved through the same generated static pipeline, so the extension path and the normal container path construct equivalent object graphs. Direct `Resolve>()` calls are generated for every registered service type, even when no constructor or source usage referenced that collection shape ahead of time. If the current generated container has no local implementations for a collection element type, collection resolution still falls back to base containers and ASP.NET Core `IServiceProvider` sources. The static extensions are generated alongside the standard dictionary-based container and require no additional configuration. If the consuming project's language version is below C# 14, the extensions are simply not emitted. @@ -254,6 +260,8 @@ app.Run(); This integration allows you to inject standard framework services (like `IConfiguration` or `ILogger`) into your `[Inject]`ed classes, and ensures that `[Scoped]` services are correctly disposed of at the end of each HTTP request. +If a request-scoped FactoryGenerator service only implements `IAsyncDisposable`, the middleware disposes it asynchronously at the end of the request. + ### Plugin / AOT Container Loading FactoryGenerator supports a plugin architecture where multiple assemblies each generate their own container, and these containers are chained together at runtime. This works seamlessly with AOT-compiled assemblies since no reflection is involved — discovery is entirely push-based via ```[ModuleInitializer]```. diff --git a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs index b2b9567..84cc1f0 100644 --- a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs +++ b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs @@ -1,5 +1,7 @@ using System; +using System.Collections.Generic; using System.Net; +using System.Linq; using System.Threading.Tasks; using Microsoft.AspNetCore.Builder; using Microsoft.AspNetCore.Hosting; @@ -54,6 +56,41 @@ public class RequestScopedFactoryService(IRequestScopedDependency dependency) : public Guid GetRequestId() => dependency.Id; } +public interface IAsyncRequestScopedService +{ +} + +[Inject, Scoped] +public class AsyncRequestScopedService : IAsyncRequestScopedService, IAsyncDisposable +{ + public static int DisposeAsyncCount { get; private set; } + + public static void Reset() + { + DisposeAsyncCount = 0; + } + + public ValueTask DisposeAsync() + { + DisposeAsyncCount++; + return default; + } +} + +public interface IFrameworkOnlyCollectionItem +{ +} + +public sealed class FrameworkOnlyCollectionItem : IFrameworkOnlyCollectionItem +{ +} + +[Inject, Self] +public class FrameworkOnlyCollectionConsumer(IEnumerable items) +{ + public IEnumerable Items { get; } = items; +} + public class IntegrationTests { [Test] @@ -154,4 +191,77 @@ public async Task Middleware_Uses_Current_RequestScope_For_FrameworkScopedDepend secondParts[0].ShouldBe(secondParts[1]); firstParts[0].ShouldNotBe(secondParts[0]); } + + [Test] + public async Task Middleware_Disposes_AsyncOnly_FactoryScopedServices() + { + AsyncRequestScopedService.Reset(); + + var host = await new HostBuilder() + .ConfigureWebHost(webBuilder => + { + webBuilder + .UseTestServer() + .ConfigureServices(services => + { + services.AddScoped(); + }) + .Configure(app => + { + var adapter = app.ApplicationServices.ToContainer(); + var container = new Generated.DependencyInjectionContainer(adapter); + + app.UseFactoryGenerator(container); + + app.Run(context => + { + _ = context.RequestServices.GetRequiredService(); + return context.Response.WriteAsync("ok"); + }); + }); + }) + .StartAsync(); + + var client = host.GetTestClient(); + var response = await client.GetAsync("/"); + + response.StatusCode.ShouldBe(HttpStatusCode.OK); + AsyncRequestScopedService.DisposeAsyncCount.ShouldBe(1); + } + + [Test] + public async Task Middleware_Uses_Framework_Collections_When_No_Local_Implementations_Exist() + { + var host = await new HostBuilder() + .ConfigureWebHost(webBuilder => + { + webBuilder + .UseTestServer() + .ConfigureServices(services => + { + services.AddScoped(); + services.AddSingleton(); + services.AddSingleton(); + }) + .Configure(app => + { + var adapter = app.ApplicationServices.ToContainer(); + var container = new Generated.DependencyInjectionContainer(adapter); + + app.UseFactoryGenerator(container); + + app.Run(async context => + { + var consumer = context.RequestServices.GetRequiredService(); + await context.Response.WriteAsync(consumer.Items.Count().ToString()); + }); + }); + }) + .StartAsync(); + + var client = host.GetTestClient(); + var response = await client.GetStringAsync("/"); + + response.ShouldBe("2"); + } } diff --git a/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs b/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs index abc39d6..a641c79 100644 --- a/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs +++ b/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs @@ -1,6 +1,7 @@ using System; using System.IO; using System.Linq; +using FactoryGenerator; using FactoryGenerator.Attributes; using Microsoft.CodeAnalysis; using Microsoft.CodeAnalysis.CSharp; @@ -94,6 +95,312 @@ public class FallbackService : IService generatedSource.ShouldContain("Resolve(DependencyInjectionContainer? container, bool boolean_feature_flag)"); } + [Test] + public void BooleanOnlyImplementationsThrowInsteadOfResolvingNull() + { + var assemblyName = "BooleanOnly" + Guid.NewGuid().ToString("N"); + var source = $$""" +using FactoryGenerator.Attributes; + +namespace {{assemblyName}} +{ +public interface IService +{ +} + +[Inject, Boolean("enabled")] +public class EnabledService : IService +{ +} +} +"""; + + var compilation = CreateCompilation(assemblyName, source); + var (runResult, outputCompilation) = RunGenerator(compilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldBeNull(); + outputCompilation.GetDiagnostics() + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .ToArray() + .ShouldBeEmpty(); + + var expectedMessage = $"Cannot resolve {assemblyName}.IService without a matching implementation"; + var declarations = generatorResult.GeneratedSources + .Single(sourceResult => sourceResult.HintName == "DependencyInjectionContainer.Declarations.g.cs") + .SourceText + .ToString(); + declarations.ShouldContain(expectedMessage); + declarations.ShouldNotContain("null!"); + + var staticExtensions = generatorResult.GeneratedSources + .Single(sourceResult => sourceResult.HintName == "DependencyInjectionContainer.StaticExtensions.g.cs") + .SourceText + .ToString(); + staticExtensions.ShouldContain(expectedMessage); + + var assembly = System.Reflection.Assembly.Load(EmitAssembly(outputCompilation)); + var containerType = assembly.GetType($"{assemblyName}.Generated.DependencyInjectionContainer"); + var serviceType = assembly.GetType($"{assemblyName}.IService"); + + containerType.ShouldNotBeNull(); + serviceType.ShouldNotBeNull(); + + var container = (IContainer)Activator.CreateInstance(containerType!, new object[] { false })!; + var exception = Should.Throw(() => container.Resolve(serviceType!)); + exception.Message.ShouldContain(expectedMessage); + } + + [Test] + public void GeneratorDetectsCyclesThroughInjectedMethods() + { + const string source = """ +using FactoryGenerator.Attributes; + +namespace Sample +{ +public interface IResult +{ +} + +public class Result : IResult +{ +} + +public interface IFactory +{ + [Inject] + IResult Create(); +} + +[Inject] +public class Factory : IFactory +{ + public Factory(IResult result) + { + } + + public IResult Create() => new Result(); +} +} +"""; + + var compilation = CreateCompilation(source); + var (runResult, _) = RunGenerator(compilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldNotBeNull(); + generatorResult.Exception!.Message.ShouldContain("Cyclic Dependency Detected"); + generatorResult.Exception.Message.ShouldContain("Sample.IResult"); + generatorResult.Exception.Message.ShouldContain("Sample.IFactory"); + } + + [Test] + public void GeneratorDetectsCyclesThroughInjectedProperties() + { + const string source = """ +using FactoryGenerator.Attributes; + +namespace Sample +{ +public interface IResult +{ +} + +public class Result : IResult +{ +} + +public interface IFactory +{ + [Inject] + IResult Value { get; } +} + +[Inject] +public class Factory : IFactory +{ + public Factory(IResult result) + { + } + + public IResult Value => new Result(); +} +} +"""; + + var compilation = CreateCompilation(source); + var (runResult, _) = RunGenerator(compilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldNotBeNull(); + generatorResult.Exception!.Message.ShouldContain("Cyclic Dependency Detected"); + generatorResult.Exception.Message.ShouldContain("Sample.IResult"); + generatorResult.Exception.Message.ShouldContain("Sample.IFactory"); + } + + [Test] + public void InjectedMethodsSurfaceExternalParametersAndHonorOptionalAndParamsArguments() + { + var assemblyName = "InjectedMethod" + Guid.NewGuid().ToString("N"); + var source = $$""" +using FactoryGenerator.Attributes; + +namespace {{assemblyName}} +{ +public sealed class ExternalInput +{ + public ExternalInput(string name) + { + Name = name; + } + + public string Name { get; } +} + +public interface IPart +{ +} + +[Inject] +public class PartOne : IPart +{ +} + +[Inject] +public class PartTwo : IPart +{ +} + +public interface IResult +{ +} + +public sealed class Result : IResult +{ + public Result(string summary, int partCount) + { + Summary = summary; + PartCount = partCount; + } + + public string Summary { get; } + public int PartCount { get; } +} + +public interface IFactory +{ + [Inject] + IResult Create(ExternalInput input, string label = "default", params IPart[] parts); +} + +[Inject] +public class Factory : IFactory +{ + public IResult Create(ExternalInput input, string label = "default", params IPart[] parts) + { + return new Result(input.Name + ":" + label, parts.Length); + } +} +} +"""; + + var compilation = CreateCompilation(assemblyName, source); + var (runResult, outputCompilation) = RunGenerator(compilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldBeNull(); + outputCompilation.GetDiagnostics() + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .ToArray() + .ShouldBeEmpty(); + + var assembly = System.Reflection.Assembly.Load(EmitAssembly(outputCompilation)); + var containerType = assembly.GetType($"{assemblyName}.Generated.DependencyInjectionContainer"); + var externalType = assembly.GetType($"{assemblyName}.ExternalInput"); + var serviceType = assembly.GetType($"{assemblyName}.IResult"); + + containerType.ShouldNotBeNull(); + externalType.ShouldNotBeNull(); + serviceType.ShouldNotBeNull(); + + var constructor = containerType!.GetConstructors() + .Single(ctor => + { + var parameters = ctor.GetParameters(); + return parameters.Length == 1 && parameters[0].ParameterType == externalType; + }); + + var external = Activator.CreateInstance(externalType!, "runtime"); + var container = (IContainer)constructor.Invoke(new[] { external! }); + var resolved = container.Resolve(serviceType!); + + resolved.GetType().GetProperty("Summary")!.GetValue(resolved).ShouldBe("runtime:default"); + resolved.GetType().GetProperty("PartCount")!.GetValue(resolved).ShouldBe(2); + } + + [Test] + public void InjectedConstructorsHonorOptionalAndParamsArguments() + { + var assemblyName = "InjectedConstructor" + Guid.NewGuid().ToString("N"); + var source = $$""" +using FactoryGenerator.Attributes; + +namespace {{assemblyName}} +{ +public interface IPart +{ +} + +[Inject] +public class PartOne : IPart +{ +} + +[Inject] +public class PartTwo : IPart +{ +} + +[Inject, Self] +public sealed class Consumer +{ + public Consumer(string label = "default", params IPart[] parts) + { + Summary = label; + PartCount = parts.Length; + } + + public string Summary { get; } + public int PartCount { get; } +} +} +"""; + + var compilation = CreateCompilation(assemblyName, source); + var (runResult, outputCompilation) = RunGenerator(compilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldBeNull(); + outputCompilation.GetDiagnostics() + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .ToArray() + .ShouldBeEmpty(); + + var assembly = System.Reflection.Assembly.Load(EmitAssembly(outputCompilation)); + var containerType = assembly.GetType($"{assemblyName}.Generated.DependencyInjectionContainer"); + var consumerType = assembly.GetType($"{assemblyName}.Consumer"); + + containerType.ShouldNotBeNull(); + consumerType.ShouldNotBeNull(); + + var container = (IContainer)Activator.CreateInstance(containerType!)!; + var resolved = container.Resolve(consumerType!); + + resolved.GetType().GetProperty("Summary")!.GetValue(resolved).ShouldBe("default"); + resolved.GetType().GetProperty("PartCount")!.GetValue(resolved).ShouldBe(2); + } + [Test] public void AssemblyPriorityCanOverrideProjectGraphPrecedence() { @@ -204,7 +511,7 @@ private static (MetadataReference Reference, byte[] Image) EmitReference(CSharpC return (MetadataReference.CreateFromImage(image), image); } - private static byte[] EmitAssembly(CSharpCompilation compilation) + private static byte[] EmitAssembly(Compilation compilation) { using var stream = new MemoryStream(); var result = compilation.Emit(stream); diff --git a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs index ca9eb7c..4332eb4 100644 --- a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs +++ b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs @@ -2,6 +2,7 @@ using Inheritor; using Inheritor.Generated; using Shouldly; +using System.Threading.Tasks; using Type = Inherited.Type; namespace FactoryGenerator.Tests; @@ -130,6 +131,21 @@ public void DontPickupIDisposable() true.ShouldBeFalse(); } + [Test] + public void DontPickupIAsyncDisposable() + { + try + { + m_container.Resolve(); + } + catch (Exception) + { + return; + } + + true.ShouldBeFalse(); + } + [Test] public void InterfacesContainingIDisposableInTheNameRemainResolvable() { @@ -250,6 +266,45 @@ public void DisposingContainerDoesNotDisposesUnreferencedSingletons() using var myContainer = new DependencyInjectionContainer(false, default, default!); } + [Test] + public void DisposingContainerSynchronouslyWaitsForAsyncOnlyServices() + { + var myContainer = new DependencyInjectionContainer(false, default, default!); + var singleton = myContainer.Resolve(); + + myContainer.Dispose(); + + singleton.ShouldBeOfType(); + singleton.WasDisposed.ShouldBeTrue(); + } + + [Test] + public async Task AsyncDisposingContainerDisposesAsyncSingletons() + { + IAsyncSingletonDisposer singleton; + var myContainer = new DependencyInjectionContainer(false, default, default!); + singleton = myContainer.Resolve(); + + await myContainer.DisposeAsync(); + + singleton.ShouldBeOfType(); + singleton.WasDisposed.ShouldBeTrue(); + } + + [Test] + public async Task AsyncDisposingLifetimeContainerDisposesAsyncScoped() + { + var myContainer = new DependencyInjectionContainer(false, default, default!); + var lifetime = myContainer.BeginLifetimeScope(); + var scoped = lifetime.Resolve(); + + await lifetime.DisposeAsync(); + + scoped.ShouldBeOfType(); + scoped.WasDisposed.ShouldBeTrue(); + myContainer.Dispose(); + } + [Test] public void ArrayExpressionsCollect() { @@ -327,6 +382,27 @@ public void HierarchicalContainersResolveUsesFallBackIfItCannotFindImplementatio newContainer.Resolve().ShouldBe(DummyContainer.DummyText); } + [Test] + public void HierarchicalContainersResolveCollectionsFromBaseWhenNoLocalImplementationExists() + { + var newContainer = new DependencyInjectionContainer(new DummyContainer()); + newContainer.Resolve().Items.Count().ShouldBe(2); + } + + [Test] + public void DirectCollectionResolveUsesBaseWhenNoLocalImplementationExists() + { + var newContainer = new DependencyInjectionContainer(new DummyContainer()); + newContainer.Resolve>().Count().ShouldBe(2); + } + + [Test] + public void StaticExtensionsUseContainerFallbackForCollectionsWithoutLocalImplementations() + { + var newContainer = new DependencyInjectionContainer(new DummyContainer()); + FallbackCollectionConsumer.Resolve(newContainer).Items.Count().ShouldBe(2); + } + [Test] public void ContainerPropgatesRelevantBooleansCreateItself() { @@ -465,6 +541,11 @@ public void BaseContainerSeesInheritorArraysAfterLinking() private class DummyContainer : IContainer { public const string DummyText = "I am a bit of text"; + private static readonly IFallbackCollectionItem[] s_fallbackCollectionItems = + [ + new DummyFallbackCollectionItem(), + new DummyFallbackCollectionItem() + ]; public static NonInjectedClass m_dummy = new(); public IContainer? Base => null; @@ -498,12 +579,14 @@ public bool IsRegistered() public T Resolve() { if (typeof(T) == typeof(string)) return (T) (object) DummyText; + if (typeof(T) == typeof(IEnumerable)) return (T) (object) s_fallbackCollectionItems; return (T) (object) m_dummy; } public object Resolve(System.Type type) { if (type == typeof(string)) return DummyText; + if (type == typeof(IEnumerable)) return s_fallbackCollectionItems; return m_dummy; } @@ -511,6 +594,7 @@ public bool TryResolve(System.Type type, out object? resolved) { resolved = null; if (type == typeof(string)) resolved = DummyText; + if (type == typeof(IEnumerable)) resolved = s_fallbackCollectionItems; return resolved != null; } @@ -518,6 +602,7 @@ public bool TryResolve(out T? resolved) { resolved = default; if (typeof(T) == typeof(string)) resolved = (T) (object) DummyText; + if (typeof(T) == typeof(IEnumerable)) resolved = (T) (object) s_fallbackCollectionItems; return resolved != null; } public IEnumerable<(string Key, bool Value)> GetBooleans() @@ -526,4 +611,6 @@ public bool TryResolve(out T? resolved) } } + + private sealed class DummyFallbackCollectionItem : IFallbackCollectionItem; } \ No newline at end of file diff --git a/Tests/FactoryGenerator.Tests/ResolvedInstanceTrackerTests.cs b/Tests/FactoryGenerator.Tests/ResolvedInstanceTrackerTests.cs new file mode 100644 index 0000000..a65ccab --- /dev/null +++ b/Tests/FactoryGenerator.Tests/ResolvedInstanceTrackerTests.cs @@ -0,0 +1,102 @@ +using System.Collections.Concurrent; +using System.Linq; +using System.Threading; +using System.Threading.Tasks; +using Shouldly; + +namespace FactoryGenerator.Tests; + +public class ResolvedInstanceTrackerTests +{ + [Test] + public void ConcurrentTrackingDuringDisposeDisposesEveryInstance() + { + const int count = 128; + var tracker = new ResolvedInstanceTracker(); + var instances = new ConcurrentBag(); + var start = new ManualResetEventSlim(false); + + var tasks = Enumerable.Range(0, count) + .Select(_ => Task.Run(() => + { + var instance = new SyncDisposableProbe(); + instances.Add(instance); + start.Wait(); + tracker.Track(instance); + })) + .ToArray(); + + start.Set(); + tracker.Dispose(); + Task.WhenAll(tasks).GetAwaiter().GetResult(); + + instances.Count.ShouldBe(count); + instances.All(instance => instance.WasDisposed).ShouldBeTrue(); + } + + [Test] + public void SynchronousDisposeWaitsForAsyncOnlyInstances() + { + var tracker = new ResolvedInstanceTracker(); + var instance = new AsyncDisposableProbe(); + tracker.Track(instance); + + tracker.Dispose(); + + instance.WasDisposed.ShouldBeTrue(); + } + + [Test] + public async Task AsynchronousDisposeDisposesAsyncOnlyInstances() + { + var tracker = new ResolvedInstanceTracker(); + var instance = new AsyncDisposableProbe(); + tracker.Track(instance); + + await tracker.DisposeAsync(); + + instance.WasDisposed.ShouldBeTrue(); + } + + [Test] + public void TrackingDoesNotKeepObjectsAlive() + { + var tracker = new ResolvedInstanceTracker(); + var weakReference = CreateTrackedWeakReference(tracker); + + GC.Collect(); + GC.WaitForPendingFinalizers(); + GC.Collect(); + + weakReference.TryGetTarget(out _).ShouldBeFalse(); + tracker.Dispose(); + } + + private static WeakReference CreateTrackedWeakReference(ResolvedInstanceTracker tracker) + { + var instance = new SyncDisposableProbe(); + tracker.Track(instance); + return new WeakReference(instance); + } + + private sealed class SyncDisposableProbe : IDisposable + { + public bool WasDisposed { get; private set; } + + public void Dispose() + { + WasDisposed = true; + } + } + + private sealed class AsyncDisposableProbe : IAsyncDisposable + { + public bool WasDisposed { get; private set; } + + public ValueTask DisposeAsync() + { + WasDisposed = true; + return default; + } + } +} diff --git a/Tests/TestData/Inherited/Types.cs b/Tests/TestData/Inherited/Types.cs index 375ed9b..4729fea 100644 --- a/Tests/TestData/Inherited/Types.cs +++ b/Tests/TestData/Inherited/Types.cs @@ -1,5 +1,6 @@ using FactoryGenerator.Attributes; using System.Collections.Immutable; +using System.Threading.Tasks; namespace Inherited; @@ -169,6 +170,14 @@ public class UnrequestedEnumerable1 : IUnrequestedEnumerable; [Inject] public class UnrequestedEnumerable2 : IUnrequestedEnumerable; +public interface IFallbackCollectionItem; + +[Inject, Self] +public class FallbackCollectionConsumer(IEnumerable items) +{ + public IEnumerable Items { get; } = items; +} + public interface IDisposer; public interface INotIDisposable; @@ -208,6 +217,23 @@ public void Dispose() } } +public interface IAsyncSingletonDisposer +{ + bool WasDisposed { get; } +} + +[Inject, Singleton] +public class AsyncDisposableSingleton : IAsyncSingletonDisposer, IAsyncDisposable +{ + public bool WasDisposed { get; private set; } + + public ValueTask DisposeAsync() + { + WasDisposed = true; + return default; + } +} + public interface IOverrideBoolean; [Inject, Boolean("A")] @@ -227,6 +253,11 @@ public interface IScoped bool WasDisposed { get; } } +public interface IAsyncScoped +{ + bool WasDisposed { get; } +} + public interface ISelfish; public interface ISelfReferentialFactory @@ -260,6 +291,18 @@ public void Dispose() } } +[Inject, Scoped] +public class AsyncScoped : IAsyncScoped, IAsyncDisposable +{ + public bool WasDisposed { get; private set; } + + public ValueTask DisposeAsync() + { + WasDisposed = true; + return default; + } +} + // ── Nullable parameter tests ───────────────────────────────────────────────── /// Interface with no [Inject] implementation — intentionally unregistered. From ed94836b021459d57192f1b0a89a292bd60892fd Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 15:50:50 +0200 Subject: [PATCH 18/21] Generator benchmarks --- .github/workflows/benchmark.yml | 14 +- .github/workflows/build.yml | 14 +- Benchmarking/Benchmarks/Benchmarks.csproj | 3 + .../Benchmarks/GeneratorBenchmarks.cs | 751 ++++++++++++++++++ Benchmarking/Benchmarks/Program.cs | 2 +- 5 files changed, 781 insertions(+), 3 deletions(-) create mode 100644 Benchmarking/Benchmarks/GeneratorBenchmarks.cs diff --git a/.github/workflows/benchmark.yml b/.github/workflows/benchmark.yml index c0a86dc..74f3b7c 100644 --- a/.github/workflows/benchmark.yml +++ b/.github/workflows/benchmark.yml @@ -27,7 +27,7 @@ jobs: - name: Run benchmark run: cd Benchmarking/Benchmarks && dotnet run -c Release --exporters json --filter '*' - - name: Store benchmark result + - name: Store runtime benchmark result uses: rhysd/github-action-benchmark@v1 with: name: Benchmark.Net Benchmark @@ -39,3 +39,15 @@ jobs: # Show alert with commit comment on detecting possible performance regression alert-threshold: '200%' comment-on-alert: true + + - name: Store generator benchmark result + uses: rhysd/github-action-benchmark@v1 + with: + name: Generator Benchmark.Net Benchmark + tool: 'benchmarkdotnet' + output-file-path: Benchmarking/Benchmarks/BenchmarkDotNet.Artifacts/results/Benchmarks.GeneratorBenchmarks-report-full-compressed.json + github-token: ${{ secrets.GITHUB_TOKEN }} + auto-push: true + summary-always: true + alert-threshold: '200%' + comment-on-alert: true diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index c07a597..6f7fc4c 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -114,7 +114,7 @@ jobs: - name: Run benchmark run: cd Benchmarking/Benchmarks && dotnet run -c Release --exporters json --filter '*' - - name: Store benchmark result + - name: Store runtime benchmark result uses: rhysd/github-action-benchmark@v1 with: name: Benchmark.Net Benchmark @@ -126,3 +126,15 @@ jobs: alert-threshold: '200%' comment-on-alert: true fail-on-alert: true + + - name: Store generator benchmark result + uses: rhysd/github-action-benchmark@v1 + with: + name: Generator Benchmark.Net Benchmark + tool: 'benchmarkdotnet' + output-file-path: Benchmarking/Benchmarks/BenchmarkDotNet.Artifacts/results/Benchmarks.GeneratorBenchmarks-report-full-compressed.json + github-token: ${{ secrets.GITHUB_TOKEN }} + summary-always: true + alert-threshold: '200%' + comment-on-alert: true + fail-on-alert: true diff --git a/Benchmarking/Benchmarks/Benchmarks.csproj b/Benchmarking/Benchmarks/Benchmarks.csproj index ef7a918..78c4269 100644 --- a/Benchmarking/Benchmarks/Benchmarks.csproj +++ b/Benchmarking/Benchmarks/Benchmarks.csproj @@ -10,9 +10,12 @@ + + + diff --git a/Benchmarking/Benchmarks/GeneratorBenchmarks.cs b/Benchmarking/Benchmarks/GeneratorBenchmarks.cs new file mode 100644 index 0000000..14a136a --- /dev/null +++ b/Benchmarking/Benchmarks/GeneratorBenchmarks.cs @@ -0,0 +1,751 @@ +using System.Collections.Immutable; +using System.IO; +using System.Text; +using BenchmarkDotNet.Attributes; +using FactoryGenerator.Attributes; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using Microsoft.CodeAnalysis.Diagnostics; + +namespace Benchmarks; + +[MemoryDiagnoser] +[ShortRunJob] +[JsonExporterAttribute.Full] +[JsonExporterAttribute.FullCompressed] +public class GeneratorBenchmarks +{ + private ColdGeneratorScenario m_constructorGraph = null!; + private ColdGeneratorScenario m_noiseHeavyProject = null!; + private ColdGeneratorScenario m_featureRichStaticExtensionsDisabled = null!; + private ColdGeneratorScenario m_featureRichStaticExtensionsEnabled = null!; + private ColdGeneratorScenario m_multiAssemblyOverrideGraph = null!; + private FeatureRichIncrementalScenario m_featureRichIncremental = null!; + private IncrementalGeneratorScenario m_referenceAssemblyIncremental = null!; + + [GlobalSetup] + public void Setup() + { + m_constructorGraph = GeneratorBenchmarkScenarioFactory.CreateConstructorGraph(serviceCount: 250); + m_noiseHeavyProject = GeneratorBenchmarkScenarioFactory.CreateNoiseHeavyProject(serviceCount: 64, noiseTypeCount: 2000); + m_featureRichStaticExtensionsDisabled = GeneratorBenchmarkScenarioFactory.CreateFeatureRichGraph(emitStaticExtensions: false); + m_featureRichStaticExtensionsEnabled = GeneratorBenchmarkScenarioFactory.CreateFeatureRichGraph(emitStaticExtensions: true); + m_multiAssemblyOverrideGraph = GeneratorBenchmarkScenarioFactory.CreateMultiAssemblyOverrideGraph(baseServiceCount: 128, overrideCount: 16); + m_featureRichIncremental = GeneratorBenchmarkScenarioFactory.CreateFeatureRichIncrementalScenario(); + m_referenceAssemblyIncremental = GeneratorBenchmarkScenarioFactory.CreateReferenceAssemblyIncrementalScenario(); + + GeneratorBenchmarkHarness.Validate(m_constructorGraph); + GeneratorBenchmarkHarness.Validate(m_noiseHeavyProject); + GeneratorBenchmarkHarness.Validate(m_featureRichStaticExtensionsDisabled); + GeneratorBenchmarkHarness.Validate(m_featureRichStaticExtensionsEnabled); + GeneratorBenchmarkHarness.Validate(m_multiAssemblyOverrideGraph); + } + + [Benchmark] + public int Cold_ConstructorGraph() => GeneratorBenchmarkHarness.RunCold(m_constructorGraph); + + [Benchmark] + public int Cold_NoiseHeavyProject() => GeneratorBenchmarkHarness.RunCold(m_noiseHeavyProject); + + [Benchmark] + public int Cold_FeatureRichGraph_StaticExtensionsDisabled() => GeneratorBenchmarkHarness.RunCold(m_featureRichStaticExtensionsDisabled); + + [Benchmark] + public int Cold_FeatureRichGraph_StaticExtensionsEnabled() => GeneratorBenchmarkHarness.RunCold(m_featureRichStaticExtensionsEnabled); + + [Benchmark] + public int Cold_MultiAssemblyOverrideGraph() => GeneratorBenchmarkHarness.RunCold(m_multiAssemblyOverrideGraph); +} + +internal sealed class ColdGeneratorScenario(CSharpCompilation compilation, AnalyzerConfigOptionsProvider optionsProvider) +{ + public CSharpCompilation Compilation { get; } = compilation; + public AnalyzerConfigOptionsProvider OptionsProvider { get; } = optionsProvider; +} + +internal sealed class FeatureRichIncrementalScenario( + GeneratorDriver warmDriver, + CSharpCompilation baselineCompilation, + CSharpCompilation unrelatedEditCompilation, + CSharpCompilation injectedSignatureEditCompilation, + CSharpCompilation addInjectCompilation) +{ + public GeneratorDriver WarmDriver { get; } = warmDriver; + public CSharpCompilation BaselineCompilation { get; } = baselineCompilation; + public CSharpCompilation UnrelatedEditCompilation { get; } = unrelatedEditCompilation; + public CSharpCompilation InjectedSignatureEditCompilation { get; } = injectedSignatureEditCompilation; + public CSharpCompilation AddInjectCompilation { get; } = addInjectCompilation; +} + +internal sealed class IncrementalGeneratorScenario(GeneratorDriver warmDriver, CSharpCompilation changedCompilation) +{ + public GeneratorDriver WarmDriver { get; } = warmDriver; + public CSharpCompilation ChangedCompilation { get; } = changedCompilation; +} + +internal static class GeneratorBenchmarkHarness +{ + public static int RunCold(ColdGeneratorScenario scenario) + { + var driver = CreateDriver(scenario.Compilation, scenario.OptionsProvider); + return RunAndSummarize(driver, scenario.Compilation); + } + + public static int RunIncremental(GeneratorDriver warmDriver, CSharpCompilation compilation) + { + return RunAndSummarize(warmDriver, compilation); + } + + public static GeneratorDriver WarmAndValidate(ColdGeneratorScenario scenario) + { + var driver = CreateDriver(scenario.Compilation, scenario.OptionsProvider); + return RunAndValidate(driver, scenario.Compilation); + } + + public static void Validate(ColdGeneratorScenario scenario) + { + _ = WarmAndValidate(scenario); + } + + public static void Validate(GeneratorDriver warmDriver, CSharpCompilation compilation) + { + _ = RunAndValidate(warmDriver, compilation); + } + + private static GeneratorDriver CreateDriver(CSharpCompilation compilation, AnalyzerConfigOptionsProvider optionsProvider) + { + return CSharpGeneratorDriver.Create( + [new global::FactoryGenerator.FactoryGenerator().AsSourceGenerator()], + parseOptions: (CSharpParseOptions) compilation.SyntaxTrees.First().Options, + optionsProvider: optionsProvider); + } + + private static GeneratorDriver RunAndValidate(GeneratorDriver driver, CSharpCompilation compilation) + { + driver = driver.RunGeneratorsAndUpdateCompilation(compilation, out var outputCompilation, out _); + var runResult = driver.GetRunResult(); + var exception = runResult.Results + .Select(result => result.Exception) + .FirstOrDefault(resultException => resultException is not null); + + if (exception is not null) + throw exception; + + var errors = outputCompilation.GetDiagnostics() + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .Select(diagnostic => diagnostic.ToString()) + .ToArray(); + + if (errors.Length != 0) + throw new InvalidOperationException(string.Join(Environment.NewLine, errors)); + + return driver; + } + + private static int RunAndSummarize(GeneratorDriver driver, CSharpCompilation compilation) + { + driver = driver.RunGeneratorsAndUpdateCompilation(compilation, out _, out _); + var runResult = driver.GetRunResult(); + + return runResult.Results.Sum(result => result.GeneratedSources.Sum(source => source.SourceText.Length)); + } +} + +internal static class GeneratorBenchmarkScenarioFactory +{ + private static readonly ImmutableArray s_metadataReferences = CreateMetadataReferences(); + private static readonly AnalyzerConfigOptionsProvider s_staticExtensionsEnabledOptions = new BenchmarkAnalyzerConfigOptionsProvider(true); + private static readonly AnalyzerConfigOptionsProvider s_staticExtensionsDisabledOptions = new BenchmarkAnalyzerConfigOptionsProvider(false); + + public static ColdGeneratorScenario CreateConstructorGraph(int serviceCount) + { + var compilation = CreateCompilation( + "GeneratorConstructorGraphBenchmarks", + new BenchmarkSourceDocument("ConstructorGraph.cs", BuildConstructorGraphSource("GeneratorConstructorGraphInput", serviceCount))); + + return new ColdGeneratorScenario(compilation, s_staticExtensionsDisabledOptions); + } + + public static ColdGeneratorScenario CreateNoiseHeavyProject(int serviceCount, int noiseTypeCount) + { + var compilation = CreateCompilation( + "GeneratorNoiseHeavyBenchmarks", + new BenchmarkSourceDocument("ConstructorGraph.cs", BuildConstructorGraphSource("GeneratorNoiseInput", serviceCount)), + new BenchmarkSourceDocument("Noise.cs", BuildNoiseSource("GeneratorNoiseInput", noiseTypeCount))); + + return new ColdGeneratorScenario(compilation, s_staticExtensionsDisabledOptions); + } + + public static ColdGeneratorScenario CreateFeatureRichGraph(bool emitStaticExtensions) + { + var compilation = CreateCompilation( + emitStaticExtensions ? "GeneratorFeatureRichStaticExtensionsBenchmarks" : "GeneratorFeatureRichBenchmarks", + new BenchmarkSourceDocument( + "FeatureGraph.cs", + BuildFeatureRichSource("GeneratorFeatureRichInput", includeAdditionalExternalParameter: false, includeExtraWidgetInjection: false, labelDefault: "default", retryCountDefault: 3)), + new BenchmarkSourceDocument("Utilities.cs", BuildUtilitySource("GeneratorFeatureRichInput", utilitySuffix: "Baseline"))); + + return new ColdGeneratorScenario(compilation, emitStaticExtensions ? s_staticExtensionsEnabledOptions : s_staticExtensionsDisabledOptions); + } + + public static ColdGeneratorScenario CreateMultiAssemblyOverrideGraph(int baseServiceCount, int overrideCount) + { + const string baseAssemblyName = "GeneratorOverrideBase"; + const string derivedAssemblyName = "GeneratorOverrideDerived"; + + var baseCompilation = CreateCompilation( + baseAssemblyName, + new BenchmarkSourceDocument("BaseServices.cs", BuildOverrideBaseSource(baseAssemblyName, baseServiceCount))); + var baseReference = EmitReference(baseCompilation); + var derivedCompilation = CreateCompilation( + derivedAssemblyName, + baseReference, + new BenchmarkSourceDocument("DerivedServices.cs", BuildOverrideDerivedSource(baseAssemblyName, derivedAssemblyName, baseServiceCount, overrideCount))); + + return new ColdGeneratorScenario(derivedCompilation, s_staticExtensionsDisabledOptions); + } + + public static FeatureRichIncrementalScenario CreateFeatureRichIncrementalScenario() + { + const string assemblyName = "GeneratorFeatureRichIncremental"; + + var baselineCompilation = CreateCompilation( + assemblyName, + new BenchmarkSourceDocument( + "FeatureGraph.cs", BuildFeatureRichSource(assemblyName, includeAdditionalExternalParameter: false, includeExtraWidgetInjection: false, labelDefault: "default", retryCountDefault: 3)), + new BenchmarkSourceDocument("Utilities.cs", BuildUtilitySource(assemblyName, utilitySuffix: "Baseline"))); + var unrelatedEditCompilation = CreateCompilation( + assemblyName, + new BenchmarkSourceDocument( + "FeatureGraph.cs", BuildFeatureRichSource(assemblyName, includeAdditionalExternalParameter: false, includeExtraWidgetInjection: false, labelDefault: "default", retryCountDefault: 3)), + new BenchmarkSourceDocument("Utilities.cs", BuildUtilitySource(assemblyName, utilitySuffix: "Edited"))); + var injectedSignatureEditCompilation = CreateCompilation( + assemblyName, + new BenchmarkSourceDocument( + "FeatureGraph.cs", BuildFeatureRichSource(assemblyName, includeAdditionalExternalParameter: true, includeExtraWidgetInjection: false, labelDefault: "edited", retryCountDefault: 5)), + new BenchmarkSourceDocument("Utilities.cs", BuildUtilitySource(assemblyName, utilitySuffix: "Baseline"))); + var addInjectCompilation = CreateCompilation( + assemblyName, + new BenchmarkSourceDocument( + "FeatureGraph.cs", BuildFeatureRichSource(assemblyName, includeAdditionalExternalParameter: false, includeExtraWidgetInjection: true, labelDefault: "default", retryCountDefault: 3)), + new BenchmarkSourceDocument("Utilities.cs", BuildUtilitySource(assemblyName, utilitySuffix: "Baseline"))); + + var baselineScenario = new ColdGeneratorScenario(baselineCompilation, s_staticExtensionsEnabledOptions); + var warmDriver = GeneratorBenchmarkHarness.WarmAndValidate(baselineScenario); + GeneratorBenchmarkHarness.Validate(GeneratorBenchmarkHarness.WarmAndValidate(baselineScenario), unrelatedEditCompilation); + GeneratorBenchmarkHarness.Validate(GeneratorBenchmarkHarness.WarmAndValidate(baselineScenario), injectedSignatureEditCompilation); + GeneratorBenchmarkHarness.Validate(GeneratorBenchmarkHarness.WarmAndValidate(baselineScenario), addInjectCompilation); + + return new FeatureRichIncrementalScenario( + warmDriver, + baselineCompilation, + unrelatedEditCompilation, + injectedSignatureEditCompilation, + addInjectCompilation); + } + + public static IncrementalGeneratorScenario CreateReferenceAssemblyIncrementalScenario() + { + const string baseAssemblyName = "GeneratorReferenceBase"; + const string derivedAssemblyName = "GeneratorReferenceDerived"; + + var baseCompilation = CreateCompilation( + baseAssemblyName, + new BenchmarkSourceDocument("BaseServices.cs", BuildReferenceBaseSource(baseAssemblyName, includeSecondBasePart: false))); + var changedBaseCompilation = CreateCompilation( + baseAssemblyName, + new BenchmarkSourceDocument("BaseServices.cs", BuildReferenceBaseSource(baseAssemblyName, includeSecondBasePart: true))); + + var baselineCompilation = CreateCompilation( + derivedAssemblyName, + EmitReference(baseCompilation), + new BenchmarkSourceDocument("DerivedServices.cs", BuildReferenceDerivedSource(baseAssemblyName, derivedAssemblyName))); + var changedCompilation = CreateCompilation( + derivedAssemblyName, + EmitReference(changedBaseCompilation), + new BenchmarkSourceDocument("DerivedServices.cs", BuildReferenceDerivedSource(baseAssemblyName, derivedAssemblyName))); + + var baselineScenario = new ColdGeneratorScenario(baselineCompilation, s_staticExtensionsEnabledOptions); + var warmDriver = GeneratorBenchmarkHarness.WarmAndValidate(baselineScenario); + GeneratorBenchmarkHarness.Validate(GeneratorBenchmarkHarness.WarmAndValidate(baselineScenario), changedCompilation); + + return new IncrementalGeneratorScenario(warmDriver, changedCompilation); + } + + private static CSharpCompilation CreateCompilation(string assemblyName, params BenchmarkSourceDocument[] documents) + { + return CreateCompilation(assemblyName, s_metadataReferences, documents); + } + + private static CSharpCompilation CreateCompilation( + string assemblyName, + MetadataReference additionalReference, + params BenchmarkSourceDocument[] documents) + { + return CreateCompilation(assemblyName, s_metadataReferences.Add(additionalReference), documents); + } + + private static CSharpCompilation CreateCompilation( + string assemblyName, + ImmutableArray references, + params BenchmarkSourceDocument[] documents) + { + var syntaxTrees = documents + .Select(document => CSharpSyntaxTree.ParseText( + document.Source, + new CSharpParseOptions(LanguageVersion.Preview), + path: document.FileName)) + .ToArray(); + + return CSharpCompilation.Create( + assemblyName, + syntaxTrees, + references, + new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary)); + } + + private static MetadataReference EmitReference(Compilation compilation) + { + using var stream = new MemoryStream(); + var result = compilation.Emit(stream); + if (!result.Success) + { + throw new InvalidOperationException( + string.Join(Environment.NewLine, result.Diagnostics.Select(diagnostic => diagnostic.ToString()))); + } + + return MetadataReference.CreateFromImage(stream.ToArray()); + } + + private static ImmutableArray CreateMetadataReferences() + { + var excludedAssemblies = new HashSet(StringComparer.Ordinal) + { + "Benchmarks", + "FactoryGenerator", + "FactoryGenerator.Attributes", + "FactoryGenerator.Extensions.AspNetCore", + "FactoryGenerator.Extensions.AspNetCore.Tests", + "FactoryGenerator.Tests", + "Inherited", + "Inheritor", + "TestWebApp" + }; + + return + [ + .. ((string?) AppContext.GetData("TRUSTED_PLATFORM_ASSEMBLIES"))! + .Split(Path.PathSeparator) + .Where(path => !excludedAssemblies.Contains(Path.GetFileNameWithoutExtension(path))) + .Select(path => (MetadataReference) MetadataReference.CreateFromFile(path)), + + MetadataReference.CreateFromFile(typeof(InjectAttribute).Assembly.Location) + ]; + } + + private static string BuildConstructorGraphSource(string namespaceName, int serviceCount) + { + var sb = new StringBuilder(); + sb.AppendLine("using FactoryGenerator.Attributes;"); + sb.AppendLine(); + sb.AppendLine($"namespace {namespaceName}"); + sb.AppendLine("{"); + + for (var i = 0; i < serviceCount; i++) + { + sb.AppendLine($"public interface IService{i}"); + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine(); + sb.AppendLine("[Inject]"); + if (i == 0) + { + sb.AppendLine($"public sealed class Service{i} : IService{i}"); + sb.AppendLine("{"); + sb.AppendLine("}"); + } + else + { + sb.AppendLine($"public sealed class Service{i}(IService{i - 1} previous) : IService{i}"); + sb.AppendLine("{"); + sb.AppendLine("}"); + } + + sb.AppendLine(); + } + + sb.AppendLine("[Inject]"); + sb.AppendLine($"public sealed class RootConsumer(IService{serviceCount - 1} root)"); + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine("}"); + + return sb.ToString(); + } + + private static string BuildNoiseSource(string namespaceName, int noiseTypeCount) + { + var sb = new StringBuilder(); + sb.AppendLine("using System;"); + sb.AppendLine(); + sb.AppendLine($"namespace {namespaceName}"); + sb.AppendLine("{"); + + for (var i = 0; i < noiseTypeCount; i++) + { + sb.AppendLine($"public sealed class NoiseType{i}"); + sb.AppendLine("{"); + sb.AppendLine($" public int Compute(int value) => value + {i};"); + sb.AppendLine($" public string Name => \"NoiseType{i}\";"); + sb.AppendLine(" public DateTime Timestamp => DateTime.UnixEpoch;"); + sb.AppendLine("}"); + sb.AppendLine(); + } + + sb.AppendLine("}"); + return sb.ToString(); + } + + private static string BuildFeatureRichSource( + string namespaceName, + bool includeAdditionalExternalParameter, + bool includeExtraWidgetInjection, + string labelDefault, + int retryCountDefault) + { + var additionalExternalParameter = includeAdditionalExternalParameter ? ", AdditionalExternalDependency additional" : string.Empty; + var widgetCAttribute = includeExtraWidgetInjection ? "[Inject]\n" : string.Empty; + + return $$""" + using System.Collections.Generic; + using FactoryGenerator.Attributes; + + namespace {{namespaceName}} + { + public sealed class ExternalDependency + { + } + + public sealed class AdditionalExternalDependency + { + } + + public interface IFlaggedFeature + { + } + + [Inject, Boolean("feature_enabled")] + public sealed class EnabledFeature : IFlaggedFeature + { + } + + [Inject] + public sealed class FallbackFeature : IFlaggedFeature + { + } + + public interface IWidget + { + } + + [Inject] + public sealed class WidgetA : IWidget + { + } + + [Inject] + public sealed class WidgetB : IWidget + { + } + + {{widgetCAttribute}}public sealed class WidgetC : IWidget + { + } + + public interface IPropertyResult + { + } + + public sealed class PropertyResult(IFlaggedFeature feature) : IPropertyResult + { + public IFlaggedFeature Feature { get; } = feature; + } + + public interface IPropertyFactory + { + [Inject] + IPropertyResult Value { get; } + } + + [Inject] + public sealed class PropertyFactory(IFlaggedFeature feature) : IPropertyFactory + { + public IPropertyResult Value => new PropertyResult(feature); + } + + public interface IMethodResult + { + } + + public sealed class MethodResult( + IFlaggedFeature feature, + IEnumerable widgets, + ExternalDependency external, + string label, + int retryCount) : IMethodResult + { + public IFlaggedFeature Feature { get; } = feature; + public IEnumerable Widgets { get; } = widgets; + public ExternalDependency External { get; } = external; + public string Label { get; } = label; + public int RetryCount { get; } = retryCount; + } + + public interface IFeatureFactory + { + [Inject] + IMethodResult Create( + ExternalDependency external{{additionalExternalParameter}}, + string label = "{{labelDefault}}", + int retryCount = {{retryCountDefault}}, + params IWidget[] widgets); + } + + [Inject] + public sealed class FeatureFactory(IFlaggedFeature feature) : IFeatureFactory + { + public IMethodResult Create( + ExternalDependency external{{additionalExternalParameter}}, + string label = "{{labelDefault}}", + int retryCount = {{retryCountDefault}}, + params IWidget[] widgets) + { + return new MethodResult(feature, widgets, external, label, retryCount); + } + } + + [Inject] + public sealed class FeatureGraphConsumer( + IMethodResult methodResult, + IPropertyResult propertyResult, + IEnumerable widgets, + IFlaggedFeature feature) + { + public IMethodResult MethodResult { get; } = methodResult; + public IPropertyResult PropertyResult { get; } = propertyResult; + public IEnumerable Widgets { get; } = widgets; + public IFlaggedFeature Feature { get; } = feature; + } + } + """; + } + + private static string BuildUtilitySource(string namespaceName, string utilitySuffix) + { + return $$""" + namespace {{namespaceName}} + { + public static class UtilityValues + { + public const string Marker = "{{utilitySuffix}}"; + + public static string Combine(string prefix) + { + return prefix + Marker; + } + } + } + """; + } + + private static string BuildOverrideBaseSource(string assemblyName, int baseServiceCount) + { + var sb = new StringBuilder(); + sb.AppendLine("using FactoryGenerator.Attributes;"); + sb.AppendLine(); + sb.AppendLine("[assembly: InjectionPriority(9)]"); + sb.AppendLine(); + sb.AppendLine($"namespace {assemblyName}"); + sb.AppendLine("{"); + sb.AppendLine("public interface ISharedService"); + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine(); + sb.AppendLine("[Inject]"); + sb.AppendLine("public sealed class BaseSharedService : ISharedService"); + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine(); + + for (var i = 0; i < baseServiceCount; i++) + { + sb.AppendLine($"public interface INode{i}"); + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine(); + sb.AppendLine("[Inject]"); + if (i == 0) + { + sb.AppendLine($"public sealed class BaseNode{i}(ISharedService sharedService) : INode{i}"); + } + else + { + sb.AppendLine($"public sealed class BaseNode{i}(INode{i - 1} previous) : INode{i}"); + } + + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine(); + } + + sb.AppendLine("}"); + return sb.ToString(); + } + + private static string BuildOverrideDerivedSource(string baseAssemblyName, string derivedAssemblyName, int baseServiceCount, int overrideCount) + { + var sb = new StringBuilder(); + sb.AppendLine($"using {baseAssemblyName};"); + sb.AppendLine("using FactoryGenerator.Attributes;"); + sb.AppendLine(); + sb.AppendLine($"namespace {derivedAssemblyName}"); + sb.AppendLine("{"); + + for (var i = 0; i < overrideCount; i++) + { + var serviceIndex = i * Math.Max(1, baseServiceCount / overrideCount); + sb.AppendLine("[Inject]"); + if (serviceIndex == 0) + { + sb.AppendLine($"public sealed class DerivedNode{serviceIndex}(ISharedService sharedService) : INode{serviceIndex}"); + } + else + { + sb.AppendLine($"public sealed class DerivedNode{serviceIndex}(INode{serviceIndex - 1} previous) : INode{serviceIndex}"); + } + + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine(); + } + + sb.AppendLine("[Inject]"); + sb.AppendLine($"public sealed class DerivedRoot(ISharedService sharedService, INode0 firstNode, INode{baseServiceCount - 1} lastNode)"); + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine("}"); + return sb.ToString(); + } + + private static string BuildReferenceBaseSource(string assemblyName, bool includeSecondBasePart) + { + var secondPart = includeSecondBasePart + ? """ + + [Inject] + public sealed class BasePartTwo : IBasePart + { + } + """ + : string.Empty; + + return $$""" + using System.Collections.Generic; + using FactoryGenerator.Attributes; + + namespace {{assemblyName}} + { + public interface IBasePart + { + } + + [Inject] + public sealed class BasePartOne : IBasePart + { + } + {{secondPart}} + + public interface IBaseService + { + } + + [Inject] + public sealed class BaseService(IEnumerable parts) : IBaseService + { + public IEnumerable Parts { get; } = parts; + } + } + """; + } + + private static string BuildReferenceDerivedSource(string baseAssemblyName, string derivedAssemblyName) + { + return $$""" + using System.Collections.Generic; + using FactoryGenerator.Attributes; + using {{baseAssemblyName}}; + + namespace {{derivedAssemblyName}} + { + [Inject] + public sealed class DerivedPart : IBasePart + { + } + + [Inject] + public sealed class DerivedConsumer(IBaseService service, IEnumerable parts) + { + public IBaseService Service { get; } = service; + public IEnumerable Parts { get; } = parts; + } + } + """; + } +} + +internal sealed class BenchmarkSourceDocument(string fileName, string source) +{ + public string FileName { get; } = fileName; + public string Source { get; } = source; +} + +internal sealed class BenchmarkAnalyzerConfigOptionsProvider(bool emitStaticExtensions) : AnalyzerConfigOptionsProvider +{ + private readonly AnalyzerConfigOptions m_globalOptions = new DictionaryAnalyzerConfigOptions( + new Dictionary(StringComparer.OrdinalIgnoreCase) + { + ["build_property.FactoryGenerator_EmitStaticExtensions"] = emitStaticExtensions ? "true" : "false" + }); + + public override AnalyzerConfigOptions GlobalOptions => m_globalOptions; + + public override AnalyzerConfigOptions GetOptions(SyntaxTree tree) => EmptyAnalyzerConfigOptions.Instance; + + public override AnalyzerConfigOptions GetOptions(AdditionalText textFile) => EmptyAnalyzerConfigOptions.Instance; +} + +internal sealed class DictionaryAnalyzerConfigOptions(IReadOnlyDictionary values) : AnalyzerConfigOptions +{ + public override bool TryGetValue(string key, out string value) + { + if (values.TryGetValue(key, out var foundValue)) + { + value = foundValue; + return true; + } + + value = string.Empty; + return false; + } +} + +internal sealed class EmptyAnalyzerConfigOptions : AnalyzerConfigOptions +{ + public static EmptyAnalyzerConfigOptions Instance { get; } = new(); + + public override bool TryGetValue(string key, out string value) + { + value = string.Empty; + return false; + } +} \ No newline at end of file diff --git a/Benchmarking/Benchmarks/Program.cs b/Benchmarking/Benchmarks/Program.cs index dc59dcb..15e6b33 100644 --- a/Benchmarking/Benchmarks/Program.cs +++ b/Benchmarking/Benchmarks/Program.cs @@ -69,5 +69,5 @@ public class ResolveBenchmarks internal static class Program { private static void Main(string[] args) => - BenchmarkRunner.Run(); + BenchmarkSwitcher.FromAssembly(typeof(Program).Assembly).Run(args); } \ No newline at end of file From 369b8b84017dd9ba52ffee8736c787dd1fa6f798 Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 16:51:50 +0200 Subject: [PATCH 19/21] Small benchmark fix --- Benchmarking/Benchmarks/Program.cs | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/Benchmarking/Benchmarks/Program.cs b/Benchmarking/Benchmarks/Program.cs index 15e6b33..5d93083 100644 --- a/Benchmarking/Benchmarks/Program.cs +++ b/Benchmarking/Benchmarks/Program.cs @@ -35,7 +35,12 @@ public class ResolveBenchmarks public IContainer Create() => new DependencyInjectionContainer(default, default, default!); [Benchmark] - public IContainer CreateFromSelf() => new DependencyInjectionContainer(m_container); + public void CreateFromSelf() + { + // Child containers attach to their base until disposed, so each benchmark + // invocation must clean up or the inheritor chain grows across operations. + using var child = new DependencyInjectionContainer(m_container); + } // ── Static-extension resolution (C# 14 / .NET 10+ path) ───────────────────── // From 89274cfe4f92929c2577207c48aaadbef6daafd0 Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 17:14:03 +0200 Subject: [PATCH 20/21] CI fix --- .github/workflows/build.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index 6f7fc4c..06593e7 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -124,7 +124,7 @@ jobs: summary-always: true # Show alert with commit comment on detecting possible performance regression alert-threshold: '200%' - comment-on-alert: true + comment-on-alert: false fail-on-alert: true - name: Store generator benchmark result @@ -136,5 +136,5 @@ jobs: github-token: ${{ secrets.GITHUB_TOKEN }} summary-always: true alert-threshold: '200%' - comment-on-alert: true + comment-on-alert: false fail-on-alert: true From 1eacad39a05e3c015df00f87bcda561e91aaf436 Mon Sep 17 00:00:00 2001 From: Carl Andersson Date: Fri, 24 Jul 2026 17:24:21 +0200 Subject: [PATCH 21/21] Don't fail on deviatio --- .github/workflows/build.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index 06593e7..cb76aff 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -125,7 +125,7 @@ jobs: # Show alert with commit comment on detecting possible performance regression alert-threshold: '200%' comment-on-alert: false - fail-on-alert: true + fail-on-alert: false - name: Store generator benchmark result uses: rhysd/github-action-benchmark@v1 @@ -137,4 +137,4 @@ jobs: summary-always: true alert-threshold: '200%' comment-on-alert: false - fail-on-alert: true + fail-on-alert: false