diff --git a/.github/workflows/benchmark.yml b/.github/workflows/benchmark.yml index 74f3b7c..7a7ff31 100644 --- a/.github/workflows/benchmark.yml +++ b/.github/workflows/benchmark.yml @@ -28,26 +28,35 @@ jobs: run: cd Benchmarking/Benchmarks && dotnet run -c Release --exporters json --filter '*' - name: Store runtime benchmark result - uses: rhysd/github-action-benchmark@v1 + uses: benchmark-action/github-action-benchmark@v1 with: name: Benchmark.Net Benchmark tool: 'benchmarkdotnet' output-file-path: Benchmarking/Benchmarks/BenchmarkDotNet.Artifacts/results/Benchmarks.ResolveBenchmarks-report-full-compressed.json github-token: ${{ secrets.GITHUB_TOKEN }} - auto-push: true 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: false + auto-push: true + + - name: Reset gh-pages branch for second benchmark + run: | + git fetch origin gh-pages:gh-pages --force || true - name: Store generator benchmark result - uses: rhysd/github-action-benchmark@v1 + uses: benchmark-action/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 + comment-on-alert: false + fail-on-alert: false + # The prior step already fetched and locally committed to gh-pages in this same job; + # re-fetching here would be rejected as non-fast-forward against that local commit. + skip-fetch-gh-pages: true + auto-push: true diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index cb76aff..874b311 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -101,40 +101,3 @@ jobs: dotnet nuget push $file --api-key "${{ secrets.NUGET_APIKEY }}" --source https://api.nuget.org/v3/index.json done - - benchmark: - name: Performance regression check - runs-on: ubuntu-latest - steps: - - uses: actions/checkout@v4 - - name: Setup dotnet ${{ matrix.dotnet-version }} - uses: actions/setup-dotnet@v4 - with: - dotnet-version: "10.0.x" - - name: Run benchmark - run: cd Benchmarking/Benchmarks && dotnet run -c Release --exporters json --filter '*' - - - name: Store runtime benchmark result - uses: rhysd/github-action-benchmark@v1 - with: - name: Benchmark.Net Benchmark - tool: 'benchmarkdotnet' - output-file-path: Benchmarking/Benchmarks/BenchmarkDotNet.Artifacts/results/Benchmarks.ResolveBenchmarks-report-full-compressed.json - github-token: ${{ secrets.GITHUB_TOKEN }} - summary-always: true - # Show alert with commit comment on detecting possible performance regression - alert-threshold: '200%' - comment-on-alert: false - fail-on-alert: false - - - 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: false - fail-on-alert: false diff --git a/Benchmarking/Benchmarks/BenchmarkConfigs.cs b/Benchmarking/Benchmarks/BenchmarkConfigs.cs new file mode 100644 index 0000000..d44bd9d --- /dev/null +++ b/Benchmarking/Benchmarks/BenchmarkConfigs.cs @@ -0,0 +1,64 @@ +using BenchmarkDotNet.Configs; +using BenchmarkDotNet.Jobs; +using Perfolizer.Mathematics.OutlierDetection; + +namespace Benchmarks; + +/// +/// Job configuration for the source-generator ("cold start") benchmarks in . +/// +/// Each Cold_* benchmark drives a full Roslyn compilation plus a generator run, costing anywhere from +/// ~1ms to ~20ms per invocation. The previous [ShortRunJob] preset fixed the sample count at +/// exactly 3 iterations (after 3 warmups), which is far too few at this scale: GC pauses, JIT tiering, +/// and OS thread-scheduling noise are all large relative to a single iteration — several Cold_* +/// results measured a standard error larger than the mean itself. This config instead: +/// - keeps a single process launch (LaunchCount=1) — relaunching the whole process mainly +/// re-pays JIT/compilation startup cost, which the warmup stage already amortizes, so a second +/// launch buys little extra accuracy for a much longer total run; +/// - increases warmup to 6 iterations so the JIT has fully tiered up before measurement begins; +/// - replaces the fixed iteration count with an adaptive 15-30 range, giving the engine enough +/// samples to converge on a stable estimate instead of stopping after 3; +/// - removes outliers on both sides (), since GC/JIT blips can +/// push individual iterations slower (common) or faster (rarer) than the true steady-state cost. +/// +public sealed class AccurateColdStartConfig : ManualConfig +{ + public AccurateColdStartConfig() + { + AddJob(new Job("Accurate") + .WithLaunchCount(1) + .WithWarmupCount(6) + .WithMinIterationCount(15) + .WithMaxIterationCount(30) + .WithOutlierMode(OutlierMode.RemoveAll)); + } +} + +/// +/// Job configuration for the runtime resolve/construction micro-benchmarks in . +/// +/// These benchmarks measure single-digit-to-low-hundreds of nanoseconds per call, so BenchmarkDotNet +/// unrolls each iteration into millions of invocations. At that scale, a single background GC +/// collection (workstation GC's concurrent/background mode can run mid-measurement) is enough to +/// visibly skew an iteration — this is exactly what showed up as periodic outlier spikes (e.g. +/// ResolveChain jumping from ~55ns to 100+ns on isolated iterations) in earlier runs. This config: +/// - disables concurrent/background GC (WithGcConcurrent(false)) +/// so a collection cannot preempt a measurement iteration on a background thread; workstation GC +/// still runs non-concurrently, it simply can no longer interrupt the benchmarked thread mid-iteration; +/// - widens the iteration bounds (15-25) so the dynamic stopping criteria has more samples to work +/// with before it decides the estimate has converged; +/// - removes outliers on both sides () to further suppress any +/// remaining scheduling noise. +/// +public sealed class AccurateMicroBenchmarkConfig : ManualConfig +{ + public AccurateMicroBenchmarkConfig() + { + AddJob(new Job("Accurate") + .WithGcServer(false) + .WithGcConcurrent(false) + .WithMinIterationCount(15) + .WithMaxIterationCount(25) + .WithOutlierMode(OutlierMode.RemoveAll)); + } +} diff --git a/Benchmarking/Benchmarks/GeneratorBenchmarks.cs b/Benchmarking/Benchmarks/GeneratorBenchmarks.cs index 14a136a..0a46207 100644 --- a/Benchmarking/Benchmarks/GeneratorBenchmarks.cs +++ b/Benchmarking/Benchmarks/GeneratorBenchmarks.cs @@ -10,16 +10,18 @@ namespace Benchmarks; [MemoryDiagnoser] -[ShortRunJob] +[Config(typeof(AccurateColdStartConfig))] [JsonExporterAttribute.Full] [JsonExporterAttribute.FullCompressed] public class GeneratorBenchmarks { private ColdGeneratorScenario m_constructorGraph = null!; + private ColdGeneratorScenario m_constructorGraphWithStaticExtensions = null!; private ColdGeneratorScenario m_noiseHeavyProject = null!; private ColdGeneratorScenario m_featureRichStaticExtensionsDisabled = null!; private ColdGeneratorScenario m_featureRichStaticExtensionsEnabled = null!; private ColdGeneratorScenario m_multiAssemblyOverrideGraph = null!; + private ColdGeneratorScenario m_manyAssembliesGraph = null!; private FeatureRichIncrementalScenario m_featureRichIncremental = null!; private IncrementalGeneratorScenario m_referenceAssemblyIncremental = null!; @@ -27,23 +29,36 @@ public class GeneratorBenchmarks public void Setup() { m_constructorGraph = GeneratorBenchmarkScenarioFactory.CreateConstructorGraph(serviceCount: 250); + m_constructorGraphWithStaticExtensions = GeneratorBenchmarkScenarioFactory.CreateConstructorGraphWithStaticExtensions(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_manyAssembliesGraph = GeneratorBenchmarkScenarioFactory.CreateManyAssembliesGraph(assemblyCount: 25, typesPerAssembly: 2); m_featureRichIncremental = GeneratorBenchmarkScenarioFactory.CreateFeatureRichIncrementalScenario(); m_referenceAssemblyIncremental = GeneratorBenchmarkScenarioFactory.CreateReferenceAssemblyIncrementalScenario(); GeneratorBenchmarkHarness.Validate(m_constructorGraph); + GeneratorBenchmarkHarness.Validate(m_constructorGraphWithStaticExtensions); GeneratorBenchmarkHarness.Validate(m_noiseHeavyProject); GeneratorBenchmarkHarness.Validate(m_featureRichStaticExtensionsDisabled); GeneratorBenchmarkHarness.Validate(m_featureRichStaticExtensionsEnabled); GeneratorBenchmarkHarness.Validate(m_multiAssemblyOverrideGraph); + GeneratorBenchmarkHarness.Validate(m_manyAssembliesGraph); } [Benchmark] public int Cold_ConstructorGraph() => GeneratorBenchmarkHarness.RunCold(m_constructorGraph); + /// + /// Same 250-service linear dependency chain as , but with + /// static extensions enabled — isolates PropagateStaticExtensionRequirements's fixed-point-loop + /// cost (a hypothesized, previously-untested scaling risk for long chains) from the ordinary + /// per-injection processing cost already captured by Cold_ConstructorGraph. + /// + [Benchmark] + public int Cold_ConstructorGraph_StaticExtensionsEnabled() => GeneratorBenchmarkHarness.RunCold(m_constructorGraphWithStaticExtensions); + [Benchmark] public int Cold_NoiseHeavyProject() => GeneratorBenchmarkHarness.RunCold(m_noiseHeavyProject); @@ -55,6 +70,56 @@ public void Setup() [Benchmark] public int Cold_MultiAssemblyOverrideGraph() => GeneratorBenchmarkHarness.RunCold(m_multiAssemblyOverrideGraph); + + /// + /// 25 small assemblies where each layer cumulatively references every prior layer + /// (O(n^2) reference edges), isolating GetRelevantAssemblies's assembly-reachability + /// BFS/DFS cost from per-type scanning cost (already covered by ). + /// + [Benchmark] + public int Cold_ManyAssembliesGraph() => GeneratorBenchmarkHarness.RunCold(m_manyAssembliesGraph); + + /// + /// Re-runs the warmed driver against the exact same compilation it was warmed with. Floor/ + /// reference point for the Incremental_* benchmarks below: since nothing at all changed, this + /// is the fastest possible incremental re-run and isolates Roslyn's own driver-level overhead + /// from any FactoryGenerator-specific recomputation. + /// + [Benchmark] + public int Incremental_NoOpRerun() => GeneratorBenchmarkHarness.RunIncremental(m_featureRichIncremental.WarmDriver, m_featureRichIncremental.BaselineCompilation); + + /// + /// Re-runs the warmed driver after only Utilities.cs changed — a file with zero + /// injectable types, entirely unrelated to dependency injection. In a well-incrementalized + /// generator this should cost close to ; if + /// FactoryGenerator.Initialize()'s direct use of context.CompilationProvider (threaded through + /// GetInjectionScanScope, and combined in again for the analysis/RegisterSourceOutput stages) + /// poisons Roslyn's per-stage caching, this should instead cost close to a full cold run. + /// + [Benchmark] + public int Incremental_UnrelatedEdit() => GeneratorBenchmarkHarness.RunIncremental(m_featureRichIncremental.WarmDriver, m_featureRichIncremental.UnrelatedEditCompilation); + + /// + /// Re-runs the warmed driver after an injected constructor's parameters/defaults changed — a + /// legitimate, relevant edit that should cost something regardless of pipeline architecture. + /// + [Benchmark] + public int Incremental_InjectedSignatureEdit() => GeneratorBenchmarkHarness.RunIncremental(m_featureRichIncremental.WarmDriver, m_featureRichIncremental.InjectedSignatureEditCompilation); + + /// + /// Re-runs the warmed driver after a new [Inject] attribute was added — another + /// legitimate, relevant edit that should cost something regardless of pipeline architecture. + /// + [Benchmark] + public int Incremental_AddInjection() => GeneratorBenchmarkHarness.RunIncremental(m_featureRichIncremental.WarmDriver, m_featureRichIncremental.AddInjectCompilation); + + /// + /// Re-runs a warmed driver after a *referenced assembly's* source changed (not the current + /// compilation's own source). Exercises the metadata-symbol scanning path (GetRelevantAssemblies/ + /// GetCandidateTypes over referenced assemblies) rather than the own-compilation discovery path. + /// + [Benchmark] + public int Incremental_ReferenceAssemblyChange() => GeneratorBenchmarkHarness.RunIncremental(m_referenceAssemblyIncremental.WarmDriver, m_referenceAssemblyIncremental.ChangedCompilation); } internal sealed class ColdGeneratorScenario(CSharpCompilation compilation, AnalyzerConfigOptionsProvider optionsProvider) @@ -166,6 +231,22 @@ public static ColdGeneratorScenario CreateConstructorGraph(int serviceCount) return new ColdGeneratorScenario(compilation, s_staticExtensionsDisabledOptions); } + /// + /// Same long linear dependency chain as , but with static + /// extensions enabled. Exists to directly measure whether + /// PropagateStaticExtensionRequirements's fixed-point loop (which can take one iteration + /// per hop of a dependency chain to converge) scales poorly with chain length, rather than + /// leaving that as an untested hypothesis. + /// + public static ColdGeneratorScenario CreateConstructorGraphWithStaticExtensions(int serviceCount) + { + var compilation = CreateCompilation( + "GeneratorConstructorGraphStaticExtensionsBenchmarks", + new BenchmarkSourceDocument("ConstructorGraph.cs", BuildConstructorGraphSource("GeneratorConstructorGraphStaticExtensionsInput", serviceCount))); + + return new ColdGeneratorScenario(compilation, s_staticExtensionsEnabledOptions); + } + public static ColdGeneratorScenario CreateNoiseHeavyProject(int serviceCount, int noiseTypeCount) { var compilation = CreateCompilation( @@ -205,6 +286,35 @@ public static ColdGeneratorScenario CreateMultiAssemblyOverrideGraph(int baseSer return new ColdGeneratorScenario(derivedCompilation, s_staticExtensionsDisabledOptions); } + /// + /// A layered graph of many small assemblies where each layer cumulatively references every + /// prior layer (fan-in), producing O(assemblyCount^2) reference edges rather than one edge per + /// assembly. Exists to measure GetRelevantAssemblies's assembly-reachability BFS cost in + /// isolation, since only involves 2 custom + /// assemblies and can't show a signal for that specific cost. + /// + public static ColdGeneratorScenario CreateManyAssembliesGraph(int assemblyCount, int typesPerAssembly) + { + var priorReferences = ImmutableArray.Empty; + CSharpCompilation compilation = null!; + for (var i = 0; i < assemblyCount; i++) + { + var assemblyName = $"GeneratorManyAssembliesLayer{i}"; + var references = priorReferences.IsDefaultOrEmpty ? s_metadataReferences : s_metadataReferences.AddRange(priorReferences); + compilation = CreateCompilation( + assemblyName, + references, + new BenchmarkSourceDocument($"Layer{i}.cs", BuildManyAssembliesLayerSource(assemblyName, i, typesPerAssembly))); + + // The final layer doesn't need to be emitted; only earlier layers need a real + // MetadataReference so later layers can reference them. + if (i < assemblyCount - 1) + priorReferences = priorReferences.Add(EmitReference(compilation)); + } + + return new ColdGeneratorScenario(compilation, s_staticExtensionsDisabledOptions); + } + public static FeatureRichIncrementalScenario CreateFeatureRichIncrementalScenario() { const string assemblyName = "GeneratorFeatureRichIncremental"; @@ -637,6 +747,27 @@ private static string BuildOverrideDerivedSource(string baseAssemblyName, string return sb.ToString(); } + private static string BuildManyAssembliesLayerSource(string assemblyName, int layerIndex, int typesPerLayer) + { + var sb = new StringBuilder(); + sb.AppendLine("using FactoryGenerator.Attributes;"); + sb.AppendLine(); + sb.AppendLine($"namespace {assemblyName}"); + sb.AppendLine("{"); + + for (var i = 0; i < typesPerLayer; i++) + { + sb.AppendLine("[Inject]"); + sb.AppendLine($"public sealed class Layer{layerIndex}Service{i}"); + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine(); + } + + sb.AppendLine("}"); + return sb.ToString(); + } + private static string BuildReferenceBaseSource(string assemblyName, bool includeSecondBasePart) { var secondPart = includeSecondBasePart diff --git a/Benchmarking/Benchmarks/Program.cs b/Benchmarking/Benchmarks/Program.cs index 5d93083..62c7da6 100644 --- a/Benchmarking/Benchmarks/Program.cs +++ b/Benchmarking/Benchmarks/Program.cs @@ -10,11 +10,19 @@ namespace Benchmarks; // ── Dictionary-based resolution (existing path) ────────────────────────────── [MemoryDiagnoser] +[Config(typeof(AccurateMicroBenchmarkConfig))] [JsonExporterAttribute.Full] [JsonExporterAttribute.FullCompressed] public class ResolveBenchmarks { private readonly DependencyInjectionContainer m_container = new(default, default, new NonInjectedClass()); + private ILifetimeScope m_scope = null!; + + [GlobalSetup] + public void Setup() => m_scope = m_container.BeginLifetimeScope(); + + [GlobalCleanup] + public void Cleanup() => m_scope.Dispose(); [Benchmark] public ChainA ResolveChain() => m_container.Resolve(); @@ -42,6 +50,32 @@ public void CreateFromSelf() using var child = new DependencyInjectionContainer(m_container); } + [Benchmark] + public void CreateLifetimeScope() + { + // Like CreateFromSelf above, a scope attaches to m_container's Inheritor chain until + // disposed, so each invocation must clean up or the chain grows across operations. + // LifetimeScope is now a thin subclass of DependencyInjectionContainer (it inherits every + // factory/lookup member instead of duplicating them) — this measures that construction path. + using var scope = m_container.BeginLifetimeScope(); + } + + // ── Resolution through a LifetimeScope ─────────────────────────────────────── + // + // m_scope is a single long-lived scope (created in Setup, disposed in Cleanup), so these + // benchmark steady-state resolve cost, not scope creation (see CreateLifetimeScope above). + // + // ResolveSingletonThroughScope exercises the owner-forwarding check added to singleton members + // (`if (m_singletonOwner != this) return m_singletonOwner.X();`) so every singleton resolves to + // the one instance owned by the root container, regardless of which scope resolves it. + // ResolveScopedThroughScope is the control case: [Scoped] members are cached per-scope-instance + // and never forward, so this exercises the unchanged local-cache path for comparison. + [Benchmark] + public ISingleton ResolveSingletonThroughScope() => m_scope.Resolve(); + + [Benchmark] + public IScoped ResolveScopedThroughScope() => m_scope.Resolve(); + // ── Static-extension resolution (C# 14 / .NET 10+ path) ───────────────────── // // Each Resolve(container?) call inlines the full construction chain directly — diff --git a/FactoryGenerator/FactoryGenerator.cs b/FactoryGenerator/FactoryGenerator.cs index 16ab857..6c47e72 100644 --- a/FactoryGenerator/FactoryGenerator.cs +++ b/FactoryGenerator/FactoryGenerator.cs @@ -18,26 +18,34 @@ public class LoggingOptions } [Generator] - public class FactoryGenerator : IIncrementalGenerator + public partial class FactoryGenerator : IIncrementalGenerator { private const string ToolName = nameof(FactoryGenerator); - private const string Version = "1.0.0"; + private const string Version = "2.1.0"; public void Initialize(IncrementalGeneratorInitializationContext context) { var logProvider = SetupLog(context); - var references = context.CompilationProvider.Select(GetGlobalNamespace); - var rest = references.SelectMany(FindMethods); + var scanScopes = context.CompilationProvider.Select(GetInjectionScanScope); + var rest = scanScopes.SelectMany(FindMethods); var attributes = rest.Collect(); var compilation = context.CompilationProvider; - var combined = attributes.Combine(compilation).Combine(logProvider); + + // Ordering + interface-indexing is identical work for both consumers below (the + // dictionary-based container and the static-extensions generator). Computing it once + // here means it's derived exactly once per compilation instead of once per consumer — + // see InjectionAnalysis/BuildInjectionAnalysis in FactoryGenerator.InjectionOrdering.cs. + var analysis = attributes.Combine(compilation) + .Select(static (pair, token) => BuildInjectionAnalysis(pair.Left, pair.Right, token)); + + var combined = analysis.Combine(compilation).Combine(logProvider); context.RegisterSourceOutput(combined, MakeAutofacModule); var supportsStaticExtensions = context.ParseOptionsProvider.Select(IsAtLeastCSharp14); 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); + var extensionData = analysis.Combine(compilation).Combine(staticExtensionsEnabled); context.RegisterSourceOutput(extensionData, MakeStaticExtensions); } @@ -59,2251 +67,19 @@ public void Initialize(IncrementalGeneratorInitializationContext context) } private void MakeAutofacModule(SourceProductionContext context, - ((ImmutableArray Injections, Compilation Compilation) Left, LoggingOptions? log) data) + ((InjectionAnalysis Analysis, Compilation Compilation) Left, LoggingOptions? log) data) { - var injections = data.Left.Injections; + var analysis = data.Left.Analysis; 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, 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]); - context.AddSource("DependencyInjectionContainer.EnumerableDeclarations.g.cs", source[3]); - context.AddSource("LifetimeScope.Lookup.g.cs", source[4]); - context.AddSource("LifetimeScope.Constructor.g.cs", source[5]); - context.AddSource("LifetimeScope.Declarations.g.cs", source[6]); - context.AddSource("LifetimeScope.EnumerableDeclarations.g.cs", source[7]); - context.AddSource("ContainerEntryPoint.g.cs", source[8]); - } - - private static IEnumerable FindMethods(INamespaceSymbol namespaceSymbol, CancellationToken token) - { - foreach (var type in SymbolUtility.GetAllTypes(namespaceSymbol)) - { - token.ThrowIfCancellationRequested(); - if (type.TypeKind != TypeKind.Class && type.TypeKind != TypeKind.Interface) continue; - var typeAttributes = type.GetAttributes().Concat(type.AllInterfaces.SelectMany(i => i.GetAttributes())) - .ToImmutableArray(); - if (typeAttributes.Any(IsInjection)) - { - var info = Injection.Create(type, typeAttributes, token); - if (info is not null) yield return info; - } - - foreach (var method in type.GetMembers().OfType() - .Where(method => method.DeclaredAccessibility == Accessibility.Public)) - { - var attributes = method.GetAttributes(); - if (!attributes.Any(IsInjection)) - continue; - var info = Injection.Create(method, attributes, token); - if (info is not null) yield return info; - } - - foreach (var property in type.GetMembers().OfType() - .Where(property => property.DeclaredAccessibility == Accessibility.Public)) - { - var attributes = property.GetAttributes(); - if (!attributes.Any(IsInjection)) - continue; - var info = Injection.Create(property, attributes, token); - if (info is not null) yield return info; - } - } + foreach (var injection in analysis.Ordered) + log.Log(LogLevel.Debug, $"Traversing {injection.Name} from {injection.AssemblyName} with priority {injection.AssemblyPriority}"); - bool IsInjection(AttributeData attribute) - { - return attribute.AttributeClass?.Name.Contains("Inject") == true && attribute.AttributeClass.ToString().StartsWith("FactoryGenerator.Attributes"); - } - } - - private static INamespaceSymbol GetGlobalNamespace(Compilation compilation, CancellationToken token) - { - return compilation.GlobalNamespace; + GenerateCode(analysis, compilation, log, context); } private const string ClassName = "DependencyInjectionContainer"; private const string LifetimeName = "LifetimeScope"; - - private static IEnumerable GenerateCode(ImmutableArray dataInjections, - Compilation compilation, ILogger log) - { - CheckForCycles(dataInjections, compilation); - log.Log(LogLevel.Debug, "Starting Code Generation"); - var usingStatements = $@" -using System; -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; -#nullable enable"; - - yield return $@"{usingStatements} -[GeneratedCode(""{ToolName}"", ""{Version}"")] -#nullable enable -#pragma warning disable CS0169, CS0414 -public sealed partial class {ClassName} : IContainer, IContainerScopeFactory, IContainerRegistrationMetadata, IContainerCacheInvalidator, IAsyncDisposable, IContainerLocalCollectionResolver -{{ - -#pragma warning restore CS0169, CS0414 - private IContainer GetRoot() - {{ - IContainer root = this; - while(root.Base != null) - {{ - root = root.Base; - }} - return root; - }} - private IContainer GetTop() - {{ - IContainer top = this; - while(top.Inheritor != null) - {{ - top = top.Inheritor; - }} - 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(); - private Dictionary> m_lookup; - private Dictionary> m_localCollectionLookup; - private Dictionary m_booleans; - private readonly ResolvedInstanceTracker m_resolvedInstances = new(); - - internal void TrackResolvedInstance(object instance) => m_resolvedInstances.Track(instance); - - public bool TryResolveLocalCollection(Type type, out object? resolved) - {{ - if (m_localCollectionLookup.TryGetValue(type, out var factory)) - {{ - resolved = factory(); - return true; - }} - - resolved = default; - return false; - }} - - public T Resolve() - {{ - if (m_lookup.TryGetValue(typeof(T), out var factory)) - return (T)factory(); - if (Base is not null) - return Base.Resolve(); - throw new KeyNotFoundException($""The type {{typeof(T)}} has not been registered, and thus cannot be resolved""); - }} - - public object Resolve(Type type) - {{ - if (m_lookup.TryGetValue(type, out var factory)) - return factory(); - if (Base is not null) - return Base.Resolve(type); - throw new KeyNotFoundException($""The type {{type}} has not been registered, and thus cannot be resolved""); - }} - - public void Dispose() - {{ - DetachFromBase(); - m_resolvedInstances.Dispose(); - }} - - public ValueTask DisposeAsync() - {{ - DetachFromBase(); - return m_resolvedInstances.DisposeAsync(); - }} - - public bool TryResolve(Type type, out object? resolved) - {{ - if(m_lookup.TryGetValue(type, out var factory)) - {{ - resolved = factory(); - return true; - }} - if(Base is not null) - return Base.TryResolve(type, out resolved); - resolved = default; - return false; - }} - - public bool TryResolve(out T? resolved) - {{ - if(m_lookup.TryGetValue(typeof(T), out var factory)) - {{ - resolved = (T)factory(); - return true; - }} - if(Base is not null) - return Base.TryResolve(out resolved); - resolved = default; - return false; - }} - public bool IsRegistered(Type type) - {{ - return m_lookup.ContainsKey(type) || Base?.IsRegistered(type) == true; - }} - public bool IsRegistered() => IsRegistered(typeof(T)); - public bool GetBoolean(string key) - {{ - return m_booleans.TryGetValue(key, out var value) && value; - }} - public IEnumerable<(string Key, bool Value)> GetBooleans() - {{ - foreach(var pair in m_booleans) - {{ - yield return (pair.Key, pair.Value); - }} - }} -}}"; - - var booleanKeys = dataInjections.Select(inj => inj.BooleanInjection).Where(b => b is not null) - .Select(b => b!.Key).Distinct().ToArray(); - var ordered = OrderInjections(dataInjections, compilation, log); - var (interfaceInjectors, interfaceMemberNames) = BuildInterfaceInjectors(ordered); - - var declarations = new Dictionary(); - var scopedDeclarations = new Dictionary(); - var availableInterfaceFullNames = interfaceInjectors.Keys.ToImmutableArray(); - var constructorParameters = new List(); - - foreach (var injection in ordered) - { - declarations[injection.Name] = Declaration(injection, availableInterfaceFullNames, false); - scopedDeclarations[injection.Name] = Declaration(injection, availableInterfaceFullNames, true); - - var missing = GetInjectionMissingParameters(injection, availableInterfaceFullNames); - foreach (var param in missing) - { - var key = param.TypeFullName + " " + param.Name; - if (constructorParameters.All(p => p.TypeFullName + " " + p.Name != key)) - constructorParameters.Add(param); - } - } - - 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(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))) - .Concat(ordered.Select(injection => injection.LazyFieldName)) - .Concat(new[] - { - "Base", - "Inheritor", - "GetRoot", - "GetTop", - "Dispose", - "Resolve", - "TryResolve", - "TryResolveLocalCollection", - "IsRegistered", - "GetBoolean", - "GetBooleans", - "BeginLifetimeScope", - "DisposeAsync", - "TrackResolvedInstance", - "m_resolvedInstances", - "m_localCollectionLookup", - "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]; - var ifaceMember = interfaceMemberNames[ifaceFull]; - var ifaceMethodName = ifaceMember + "()"; - - if (possibilities.All(i => i.BooleanInjection == null)) - { - var chosen = possibilities.Last(); - if (ifaceMethodName != chosen.Name) - { - if (!declarations.ContainsKey(ifaceMethodName)) - { - log.Log(LogLevel.Information, $"Selecting {chosen.Name} for {ifaceFull}"); - declarations[ifaceMethodName] = $"internal {ifaceFull} {ifaceMethodName} => {chosen.Name};"; - scopedDeclarations[ifaceMethodName] = $"internal {ifaceFull} {ifaceMethodName} => {chosen.Name};"; - } - } - } - else - { - var ternary = BuildBooleanSelectionExpression( - ifaceFull, - possibilities, - booleanIdentifiers, - possibility => possibility.Name); - - if (!declarations.ContainsKey(ifaceMethodName)) - { - log.Log(LogLevel.Information, $"Selecting {ternary} for {ifaceFull}"); - declarations[ifaceMethodName] = $"internal {ifaceFull} {ifaceMethodName} => {ternary};"; - scopedDeclarations[ifaceMethodName] = $"internal {ifaceFull} {ifaceMethodName} => {ternary};"; - } - } - } - - var arrayDeclarations = new Dictionary(); - foreach (var pair in interfaceInjectors) - { - 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); - } - - foreach (var parameter in localizedParameters) - { - 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); - } - - 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 lifetimeArguments = allArguments.ToList(); - 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) + ")"; - var lifetimeConstructor = "(" + string.Join(", ", lifetimeArguments) + ")"; - 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", - 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 = 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 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, localCollectionPairs, - 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, IContainerScopeFactory, IContainerRegistrationMetadata, IContainerCacheInvalidator, IAsyncDisposable, IContainerLocalCollectionResolver -{{ -#pragma warning restore CS0169, CS0414 - private IContainer GetRoot() - {{ - IContainer root = this; - while(root.Base != null) - {{ - root = root.Base; - }} - return root; - }} - private IContainer GetTop() - {{ - IContainer top = this; - while(top.Inheritor != null) - {{ - top = top.Inheritor; - }} - 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 baseContainer = Base?.BeginLifetimeScope() as IContainer; - return BeginLifetimeScope(baseContainer); - }} - public ILifetimeScope BeginLifetimeScope(IContainer? baseContainer) - {{ - var scope = m_fallback.BeginLifetimeScope(baseContainer); - TrackResolvedInstance(scope); - return scope; - }} - internal readonly object m_lock = new(); - private {ClassName} m_fallback; - private Dictionary> m_lookup; - private Dictionary> m_localCollectionLookup; - private Dictionary m_booleans; - private readonly ResolvedInstanceTracker m_resolvedInstances = new(); - - internal void TrackResolvedInstance(object instance) => m_resolvedInstances.Track(instance); - - public bool TryResolveLocalCollection(Type type, out object? resolved) - {{ - if (m_localCollectionLookup.TryGetValue(type, out var factory)) - {{ - resolved = factory(); - return true; - }} - - resolved = default; - return false; - }} - - public T Resolve() - {{ - if (m_lookup.TryGetValue(typeof(T), out var factory)) - return (T)factory(); - if (Base is not null) - return Base.Resolve(); - throw new KeyNotFoundException($""The type {{typeof(T)}} has not been registered, and thus cannot be resolved""); - }} - - public object Resolve(Type type) - {{ - if (m_lookup.TryGetValue(type, out var factory)) - return factory(); - if (Base is not null) - return Base.Resolve(type); - throw new KeyNotFoundException($""The type {{type}} has not been registered, and thus cannot be resolved""); - }} - - public void Dispose() - {{ - DetachFromBase(); - m_resolvedInstances.Dispose(); - }} - - public ValueTask DisposeAsync() - {{ - DetachFromBase(); - return m_resolvedInstances.DisposeAsync(); - }} - - public bool TryResolve(Type type, out object? resolved) - {{ - if(m_lookup.TryGetValue(type, out var factory)) - {{ - resolved = factory(); - return true; - }} - if(Base is not null) - return Base.TryResolve(type, out resolved); - resolved = default; - return false; - }} - - public bool TryResolve(out T? resolved) - {{ - if(m_lookup.TryGetValue(typeof(T), out var factory)) - {{ - resolved = (T)factory(); - return true; - }} - if(Base is not null) - return Base.TryResolve(out resolved); - resolved = default; - return false; - }} - public bool IsRegistered(Type type) - {{ - return m_lookup.ContainsKey(type) || Base?.IsRegistered(type) == true; - }} - public bool IsRegistered() => IsRegistered(typeof(T)); - - public bool GetBoolean(string key) - {{ - return m_booleans.TryGetValue(key, out var value) && value; - }} - public IEnumerable<(string Key, bool Value)> GetBooleans() - {{ - foreach(var pair in m_booleans) - {{ - yield return (pair.Key, pair.Value); - }} - }} -}} -"; - yield return Constructor(usingStatements, constructorFields, - lifetimeConstructor, constructorAssignments, - dictSize, interfacePairs, localizedPairs, enumerablePairs, constructorPairs, localCollectionPairs, - false, LifetimeName, - resolvingConstructorAssignments: resolvedConstructorAssignments, addMergingConstructor: false, booleans: booleanParameters); - yield return Declarations(usingStatements, scopedDeclarations, LifetimeName); - yield return ArrayDeclarations(usingStatements, arrayDeclarations, LifetimeName); - - // Emit the static factory + module initializer for plugin container registration - yield return $@" -using System; -using System.Runtime.CompilerServices; -using FactoryGenerator; - -#if !NET5_0_OR_GREATER -namespace System.Runtime.CompilerServices -{{ - [AttributeUsage(AttributeTargets.Method, AllowMultiple = false)] - internal sealed class ModuleInitializerAttribute : Attribute {{ }} -}} -#endif - -namespace {compilation.Assembly.Name}.Generated -{{ - /// - /// Provides a static factory for the generated container and auto-registers it in the ContainerRegistry on assembly load. - /// - public static class ContainerEntryPoint - {{ - /// - /// Creates a new DependencyInjectionContainer that chains on top of the given base container. - /// - public static IContainer Create(IContainer baseContainer) - {{ - return new {ClassName}(baseContainer); - }} - - /// - /// The assembly name this container was generated for. - /// - public static string AssemblyName => ""{compilation.Assembly.Name}""; - - [ModuleInitializer] - internal static void Register() - {{ - ContainerRegistry.Register(""{compilation.Assembly.Name}"", Create); - }} - }} -}} -"; - } - - 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) - 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) - { - 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 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) - { - 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); - } - } - - return (interfaceInjectors, interfaceMemberNames); - } - - private static List GetReachableImplementations(List possibilities) - { - if (possibilities.Count == 0) - return new List(); - - 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 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; - var ctor = GetBestConstructor(injection, availableInterfaceFullNames, ref missing, ref nullableDefaults); - if (ctor is null) - yield break; - - 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; - - 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 - { - 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 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_"); - } - - 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, Compilation compilation) - { - var ordered = OrderInjections(dataInjections, compilation); - 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()); - } - - 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, - IEnumerable<(string TypeName, string MemberName)> 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!) - { - var lifetimeScopeFunction = addLifetimeScopeFunction - ? $@" -public ILifetimeScope BeginLifetimeScope() -{{ - var baseContainer = Base?.BeginLifetimeScope() as IContainer; - return BeginLifetimeScope(baseContainer); -}} -public ILifetimeScope BeginLifetimeScope(IContainer? baseContainer) -{{ - var scope = new {LifetimeName}({lifetimeInvocationValues}); - TrackResolvedInstance(scope); - return scope; -}}" : string.Empty; - - var mergingConstructor = addMergingConstructor ? $@" -public {className}(IContainer Base{fromConstructor}) -{{ - this.Base = Base; - AttachToBase(Base); - {resolvingConstructorAssignments} - -{string.Join("\n", booleans.Select(boolean => $"\t this.{boolean.Identifier} = Base.GetBoolean(\"{boolean.Key}\");"))} - - m_lookup = new({dictSize}) {{ -{MakeDictionaryFromTypes(interfaceTypePairs)} -{MakeDictionaryFromParams(localizedParamPairs)} -{MakeDictionaryFromParams(enumerablePairs)} -{MakeDictionaryFromParams(constructorParamPairs)} - }}; - m_localCollectionLookup = new({localCollectionPairs.Count()}) {{ -{MakeDictionaryFromParams(localCollectionPairs)} - }}; - m_booleans = new(); - foreach(var (key, value) in Base.GetBooleans()) - {{ - m_booleans[key] = value; - }} -}}" : string.Empty; - - - var extraConstruction = addLifetimeScopeFunction ? string.Empty : @"m_fallback = fallback; - this.Base = baseContainer; - if (baseContainer is not null) - { - AttachToBase(baseContainer); - TrackResolvedInstance(baseContainer); - }"; - return $@"{usingStatements} -public partial class {className} -{{ - {constructorFields} - public {className}{constructor} - {{ - {extraConstruction} - {constructorAssignments} - - m_lookup = new({dictSize}) {{ -{MakeDictionaryFromTypes(interfaceTypePairs)} -{MakeDictionaryFromParams(localizedParamPairs)} -{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} }},"))} - }}; - }} - {mergingConstructor} - {lifetimeScopeFunction} - -}}"; - } - - 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) - { - return $@"{usingStatements} -public partial class {className} -{{ - {string.Join("\n\t", declarations.Values)} -}}"; - } - - private static void MakeArray(Dictionary declarations, string name, - string elementTypeFullName, Dictionary> interfaceInjectors, - IReadOnlyDictionary booleanIdentifiers) - { - 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} - {{ - 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(!(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(!(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} - {{ - get - {{ - var cached = m_{name}; - if (cached != null) - return cached; - - lock (m_lock) - {{ - cached = m_{name}; - if (cached != null) - return cached; - return m_{name} = {factoryName}; - }} - }} - }} - internal IEnumerable<{elementTypeFullName}> local_{name} => {localFactoryName}; - internal IEnumerable<{elementTypeFullName}>? m_{name};" + factory; - } - - private static string MakeDictionaryFromTypes(IEnumerable<(string TypeName, string MemberName)> pairs) - { - var builder = new StringBuilder(); - foreach (var (typeName, memberName) in pairs) - builder.AppendLine($"\t\t\t{{ typeof({typeName}),{memberName} }},"); - return builder.ToString(); - } - - private static string MakeDictionaryFromParams(IEnumerable<(string TypeName, string Expression)> pairs) - { - var builder = new StringBuilder(); - foreach (var (typeName, expression) in pairs) - builder.AppendLine($"\t\t\t{{ typeof({typeName}), () => {expression} }},"); - return builder.ToString(); - } - - private static string Declaration(InjectionData injection, ImmutableArray availableInterfaceFullNames, bool forLifetimeScope) - { - var name = injection.Name; - var lazyName = injection.LazyFieldName; - var creation = CreationCall(injection, availableInterfaceFullNames); - - if (forLifetimeScope && injection.Singleton) - return $"internal {injection.TypeFullName} {name} => m_fallback.{name};"; - - if (injection.Singleton || injection.Scoped) - return SymbolUtility.SingletonFactory(injection.TypeFullName, name, lazyName, creation, injection.Disposable || injection.AsyncDisposable); - - if (injection.Disposable || injection.AsyncDisposable) - return SymbolUtility.DisposableFactory(injection.TypeFullName, name, creation); - - return $"internal {injection.TypeFullName} {name} => {creation};"; - } - - private static string CreationCall(InjectionData injection, ImmutableArray availableInterfaceFullNames) - { - if (injection.Lambda is LambdaData lambda) - { - if (!availableInterfaceFullNames.Contains(lambda.ContainingTypeFullName)) - throw new Exception( - $"Could not find any [Inject]ed implementations of {lambda.ContainingTypeFullName} to use as the source for the injection of {lambda.ContainingTypeFullName}.{lambda.MemberName}. Please provide at least one injection of the type {lambda.ContainingTypeFullName}."); - - if (lambda.IsMethod) - { - HashSet? lambdaMissing = null; - HashSet? lambdaNullableDefaults = null; - AnalyzeParameters(lambda.MethodParameters, availableInterfaceFullNames, ref lambdaMissing, ref lambdaNullableDefaults); - return $"{lambda.ContainingTypeMemberName}.{lambda.MemberName}{MakeMethodCall(lambda.MethodParameters, lambdaMissing, lambdaNullableDefaults)}"; - } - else - return $"{lambda.ContainingTypeMemberName}.{lambda.MemberName}"; - } - - HashSet? missing = null; - HashSet? nullableDefaults = null; - var ctor = GetBestConstructor(injection, availableInterfaceFullNames, ref missing, ref nullableDefaults); - if (ctor is null) - throw new Exception($"No Construction method for {injection.TypeFullName}. Lambda was null."); - - return $"new {injection.TypeFullName}{MakeConstructorCall(ctor, missing, nullableDefaults)}"; - } - - private static ConstructorData? GetBestConstructor(InjectionData injection, - ImmutableArray availableInterfaceFullNames, ref HashSet? missing, - ref HashSet? nullableDefaults) - { - missing = null; - nullableDefaults = null; - ConstructorData? chosen = null; - foreach (var ctor in injection.Constructors) - { - HashSet? localMissing = null; - HashSet? localNullableDefaults = null; - AnalyzeParameters(ctor.Parameters, availableInterfaceFullNames, ref localMissing, ref localNullableDefaults, out var valid); - - if (valid) - { - chosen = ctor; - missing = localMissing; - nullableDefaults = localNullableDefaults; - break; - } - - if ((missing?.Count ?? int.MaxValue) <= (localMissing?.Count ?? 0)) continue; - chosen = ctor; - missing = localMissing; - nullableDefaults = localNullableDefaults; - } - return chosen; - } - - 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 void AnalyzeParameters( - ImmutableArray parameters, - ImmutableArray availableInterfaceFullNames, - ref HashSet? missing, - ref HashSet? nullableDefaults) - { - 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) - { - var typeLookup = parameter.IsNullable - ? parameter.TypeFullName.TrimEnd('?') - : parameter.TypeFullName; - - if (parameter.IsCollection) - { - localMissing.Add(parameter); - continue; - } - - if (availableInterfaceFullNames.Contains(typeLookup)) - continue; - - if (parameter.HasExplicitDefault || parameter.IsParams) - continue; - - if (parameter.IsNullable) - { - localNullableDefaults.Add(parameter); - continue; - } - - valid = false; - localMissing.Add(parameter); - } - - 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(useNamedArguments ? $"{parameter.Name}: null" : "null"); - continue; - } - if (missing?.Contains(parameter) == true) - { - var argument = CollectionConstructorArg(parameter); - args.Add(useNamedArguments ? $"{parameter.Name}: {argument}" : argument); - continue; - } - - if (parameter.HasExplicitDefault || parameter.IsParams) - { - useNamedArguments = true; - continue; - } - - var resolvedArgument = parameter.TypeMemberName + "()"; - args.Add(useNamedArguments ? $"{parameter.Name}: {resolvedArgument}" : resolvedArgument); - } - return $"({string.Join(", ", args)})"; - } - - /// - /// Returns the expression to use when passing a collection (or plain-missing) parameter - /// 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) - { - if (!parameter.IsCollection) - return parameter.Name; - var memberName = "coll_" + parameter.CollectionElementMemberName!; - return parameter.CollectionKind switch - { - 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 - /// for a localized (collection) parameter. - /// - private static string CollectionDictExpression(CollectionKind kind, string factoryName) => - kind switch - { - CollectionKind.Array => $"{factoryName}.ToArray()", - CollectionKind.List => $"{factoryName}.ToList()", - 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 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 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) - { - 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 = OrderInjections(dataInjections, compilation); - 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; -using System.CodeDom.Compiler; -using System.Collections.Generic; -using System.Collections.Immutable; -using System.Linq; -namespace {compilation.Assembly.Name}.Generated; -#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 ifaceFull in interfaceMemberNames.Keys) - { - var spec = specs[ifaceFull]; - var helpers = BuildStaticExtensionClass(spec, specs, availableInterfaces, booleanIdentifiers, externalIdentifiers); - - sb.AppendLine($@"[GeneratedCode(""{ToolName}"", ""{Version}"")] -public static class {spec.ExtensionClassName} -{{ -{helpers} - extension({spec.TypeFullName}) - {{ -{BuildStaticPublicResolveMethods(spec, booleanIdentifiers, externalIdentifiers)} - }} -}}"); - } - - return sb.ToString(); - } - - private static Dictionary BuildStaticExtensionSpecs( - Dictionary> interfaceInjectors, - Dictionary interfaceMemberNames, - ImmutableArray availableInterfaces) - { - var specs = interfaceInjectors.ToDictionary( - pair => pair.Key, - pair => CreateDirectStaticExtensionSpec(pair.Key, pair.Value, interfaceMemberNames[pair.Key], availableInterfaces, interfaceInjectors), - StringComparer.Ordinal); - - PropagateStaticExtensionRequirements(specs); - - foreach (var spec in specs.Values) - { - 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) - { - if (interfaceInjectors.ContainsKey(lambda.ContainingTypeFullName)) - AddDistinctDependency(spec.Dependencies, new StaticDependencyReference(lambda.ContainingTypeFullName, false)); - - foreach (var parameter in lambda.MethodParameters) - AddStaticParameterRequirement(spec, parameter, interfaceInjectors); - - continue; - } - - 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) - { - 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); - } - } - } 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)); - } - - return string.Join("\n\n", parts); - } - - 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 - { - $@" 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 - {{ - var source = new List<{spec.TypeFullName}>({resolveInvocations.Length}) {{ - {string.Join(",\n ", resolveInvocations)} - }}; -{string.Join("\n", conditionalInvocations)} - if (container is not null) - {{ - var b = container.Base; - while (b is not null) - {{ - if (b.TryResolve>(out var additional)) - source.AddRange(additional!); - b = b.Base; - }} - - b = container.Inheritor; - while (b is not null) - {{ - if (b.TryResolve>(out var additional)) - source.AddRange(additional!); - b = b.Inheritor; - }} - }} - - 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); - var tracksResolvedInstance = injection.Disposable || injection.AsyncDisposable; - - if (injection.Singleton || injection.Scoped) - { - if (tracksResolvedInstance) - { - return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) - {{ - if (container is not null) - {{ - 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.TrackResolvedInstance(value); - container.{injection.LazyFieldName} = value; - return value; - }} - }} - - return Create_{helperName}({BuildStaticInternalInvocationArguments(spec, booleanIdentifiers, externalIdentifiers, "null", "state")}); - }}"; - } - - return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) - {{ - if (container is not null) - {{ - 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")}); - }}"; - } - - if (tracksResolvedInstance) - { - return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) - {{ - if (container is not null) - {{ - var value = {createInvocation}; - container.TrackResolvedInstance(value); - return value; - }} - - return Create_{helperName}({BuildStaticInternalInvocationArguments(spec, booleanIdentifiers, externalIdentifiers, "null", "state")}); - }}"; - } - - return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) - {{ - return {createInvocation}; - }}"; - } - - private static string BuildStaticCreateInjectionMethod( - StaticExtensionSpec spec, - InjectionData injection, - IReadOnlyDictionary specs, - ImmutableArray availableInterfaces, - IReadOnlyDictionary booleanIdentifiers, - IReadOnlyDictionary externalIdentifiers) - { - 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 BuildMissingImplementationExpression(lambda.ContainingTypeFullName); - - var containingInvocation = BuildStaticResolveInvocation(containingSpec, "ResolveCore", booleanIdentifiers, externalIdentifiers); - if (!lambda.IsMethod) - return $"{containingInvocation}.{lambda.MemberName}"; - - 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 BuildMissingImplementationExpression(injection.TypeFullName); - - var constructorArguments = BuildStaticArgumentList(constructor.Parameters, specs, booleanIdentifiers, externalIdentifiers); - return $"new {injection.TypeFullName}({constructorArguments})"; - } - - 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 null) - { - useNamedArguments = true; - continue; - } - - arguments.Add(useNamedArguments - ? $"{parameter.Name}: {argumentExpression}" - : argumentExpression); - } - - return string.Join(", ", arguments); - } - - 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)) - { - 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( - parameter, - BuildStaticResolveAllInvocation(collectionSpec, booleanIdentifiers, externalIdentifiers)); - } - - var typeLookup = parameter.IsNullable - ? parameter.TypeFullName.TrimEnd('?') - : parameter.TypeFullName; - - if (specs.TryGetValue(typeLookup, out var dependencySpec)) - return BuildStaticResolveInvocation(dependencySpec, "ResolveCore", booleanIdentifiers, externalIdentifiers); - - if (parameter.HasExplicitDefault || parameter.IsParams) - return null; - - if (parameter.IsNullable) - return "null"; - - if (externalIdentifiers.TryGetValue(parameter.TypeFullName, out var identifier)) - return identifier; - - return BuildMissingImplementationExpression(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")})"; - } - - private static string BuildStaticResolveSelectionExpression( - StaticExtensionSpec spec, - IReadOnlyDictionary booleanIdentifiers, - IReadOnlyDictionary externalIdentifiers) - { - return BuildBooleanSelectionExpression( - spec.TypeFullName, - spec.Possibilities, - booleanIdentifiers, - possibility => BuildStaticResolveInjectionInvocation(spec, possibility, booleanIdentifiers, externalIdentifiers)); - } - - 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 BuildMissingImplementationExpression(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/FactoryGenerator.csproj b/FactoryGenerator/FactoryGenerator.csproj index a27e13a..3746997 100644 --- a/FactoryGenerator/FactoryGenerator.csproj +++ b/FactoryGenerator/FactoryGenerator.csproj @@ -3,7 +3,7 @@ netstandard2.0 true - 9 + 11 enable false true diff --git a/FactoryGenerator/Generation/FactoryGenerator.ConstructorResolution.cs b/FactoryGenerator/Generation/FactoryGenerator.ConstructorResolution.cs new file mode 100644 index 0000000..2a0d752 --- /dev/null +++ b/FactoryGenerator/Generation/FactoryGenerator.ConstructorResolution.cs @@ -0,0 +1,266 @@ +using System; +using System.Collections.Generic; +using System.Collections.Immutable; +using System.Linq; +using System.Text; +using System.Threading; +using System.Threading.Tasks; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using Microsoft.CodeAnalysis.Diagnostics; + +namespace FactoryGenerator +{ + /// + /// Resolves how an injection is constructed: picks the best available constructor (or lambda + /// member) given the interfaces available for injection, and builds the resulting creation + /// expression plus any parameters that must be supplied externally. + /// + public partial class FactoryGenerator + { + private static string Declaration(InjectionData injection, string creation) + { + var name = injection.Name; + var lazyName = injection.LazyFieldName; + + if (injection.Singleton || injection.Scoped) + return SymbolUtility.SingletonFactory(injection.TypeFullName, name, lazyName, creation, injection.Disposable || injection.AsyncDisposable, forwardToOwner: injection.Singleton); + + if (injection.Disposable || injection.AsyncDisposable) + return SymbolUtility.DisposableFactory(injection.TypeFullName, name, creation); + + return $"internal {injection.TypeFullName} {name} => {creation};"; + } + + /// + /// The outcome of resolving how an injection is constructed: the expression used to create + /// it, any parameters it could not satisfy from the available interfaces (destined to + /// become externally-supplied constructor parameters of the generated container), and the + /// chosen constructor's full parameter list. + /// + /// + /// exists purely so cycle-detection (GetCycleDependencies) + /// can reuse the exact constructor already picked instead of + /// calling a second time for the same injection. It is + /// for lambda-based injections (no constructor to choose). + /// + private readonly struct InjectionResolution + { + public InjectionResolution(string creation, IEnumerable missingParameters, ImmutableArray? constructorParameters) + { + Creation = creation; + MissingParameters = missingParameters; + ConstructorParameters = constructorParameters; + } + + public string Creation { get; } + public IEnumerable MissingParameters { get; } + public ImmutableArray? ConstructorParameters { get; } + } + + /// + /// Resolves the creation expression, missing parameters, and chosen constructor for every + /// injection in a single pass. Computed once per GenerateCode run and reused by both + /// cycle-detection and the declarations loop, avoiding repeated constructor selection/analysis + /// for the same injection. + /// + private static Dictionary ResolveInjections( + IEnumerable ordered, HashSet availableInterfaceFullNames) + { + var resolutions = new Dictionary(); + foreach (var injection in ordered) + resolutions[injection] = ResolveInjection(injection, availableInterfaceFullNames); + return resolutions; + } + + private static InjectionResolution ResolveInjection(InjectionData injection, HashSet availableInterfaceFullNames) + { + if (injection.Lambda is LambdaData lambda) + { + if (!availableInterfaceFullNames.Contains(lambda.ContainingTypeFullName)) + throw new Exception( + $"Could not find any [Inject]ed implementations of {lambda.ContainingTypeFullName} to use as the source for the injection of {lambda.ContainingTypeFullName}.{lambda.MemberName}. Please provide at least one injection of the type {lambda.ContainingTypeFullName}."); + + if (!lambda.IsMethod) + return new InjectionResolution($"{lambda.ContainingTypeMemberName}.{lambda.MemberName}", Enumerable.Empty(), null); + + HashSet? lambdaMissing = null; + HashSet? lambdaNullableDefaults = null; + AnalyzeParameters(lambda.MethodParameters, availableInterfaceFullNames, ref lambdaMissing, ref lambdaNullableDefaults); + var lambdaCreation = $"{lambda.ContainingTypeMemberName}.{lambda.MemberName}{MakeMethodCall(lambda.MethodParameters, lambdaMissing, lambdaNullableDefaults)}"; + return new InjectionResolution(lambdaCreation, (IEnumerable?) lambdaMissing ?? Enumerable.Empty(), null); + } + + HashSet? missing = null; + HashSet? nullableDefaults = null; + var ctor = GetBestConstructor(injection, availableInterfaceFullNames, ref missing, ref nullableDefaults); + if (ctor is null) + throw new Exception($"No Construction method for {injection.TypeFullName}. Lambda was null."); + + var creation = $"new {injection.TypeFullName}{MakeConstructorCall(ctor, missing, nullableDefaults)}"; + return new InjectionResolution(creation, (IEnumerable?) missing ?? Enumerable.Empty(), ctor.Parameters); + } + + private static ConstructorData? GetBestConstructor(InjectionData injection, + HashSet availableInterfaceFullNames, ref HashSet? missing, + ref HashSet? nullableDefaults) + { + missing = null; + nullableDefaults = null; + ConstructorData? chosen = null; + foreach (var ctor in injection.Constructors) + { + HashSet? localMissing = null; + HashSet? localNullableDefaults = null; + AnalyzeParameters(ctor.Parameters, availableInterfaceFullNames, ref localMissing, ref localNullableDefaults, out var valid); + + if (valid) + { + chosen = ctor; + missing = localMissing; + nullableDefaults = localNullableDefaults; + break; + } + + if ((missing?.Count ?? int.MaxValue) <= (localMissing?.Count ?? 0)) continue; + chosen = ctor; + missing = localMissing; + nullableDefaults = localNullableDefaults; + } + return chosen; + } + + private static void AnalyzeParameters( + ImmutableArray parameters, + HashSet availableInterfaceFullNames, + ref HashSet? missing, + ref HashSet? nullableDefaults) + { + AnalyzeParameters(parameters, availableInterfaceFullNames, ref missing, ref nullableDefaults, out _); + } + + private static void AnalyzeParameters( + ImmutableArray parameters, + HashSet availableInterfaceFullNames, + ref HashSet? missing, + ref HashSet? nullableDefaults, + out bool valid) + { + // Allocated lazily: the overwhelmingly common case is a fully-satisfied constructor + // with zero missing/nullable-default parameters, so most calls should allocate neither + // HashSet at all instead of two unconditionally per candidate constructor tried. + HashSet? localMissing = null; + HashSet? localNullableDefaults = null; + valid = true; + + foreach (var parameter in parameters) + { + var typeLookup = parameter.IsNullable + ? parameter.TypeFullName.TrimEnd('?') + : parameter.TypeFullName; + + if (parameter.IsCollection) + { + (localMissing ??= new HashSet()).Add(parameter); + continue; + } + + if (availableInterfaceFullNames.Contains(typeLookup)) + continue; + + if (parameter.HasExplicitDefault || parameter.IsParams) + continue; + + if (parameter.IsNullable) + { + (localNullableDefaults ??= new HashSet()).Add(parameter); + continue; + } + + valid = false; + (localMissing ??= new HashSet()).Add(parameter); + } + + // A HashSet is only ever allocated above when something is actually added to it, so + // "allocated" already implies non-empty — no need to re-check .Count here. + missing = localMissing; + nullableDefaults = localNullableDefaults; + } + + 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(useNamedArguments ? $"{parameter.Name}: null" : "null"); + continue; + } + if (missing?.Contains(parameter) == true) + { + var argument = CollectionConstructorArg(parameter); + args.Add(useNamedArguments ? $"{parameter.Name}: {argument}" : argument); + continue; + } + + if (parameter.HasExplicitDefault || parameter.IsParams) + { + useNamedArguments = true; + continue; + } + + var resolvedArgument = parameter.TypeMemberName + "()"; + args.Add(useNamedArguments ? $"{parameter.Name}: {resolvedArgument}" : resolvedArgument); + } + return $"({string.Join(", ", args)})"; + } + + /// + /// Returns the expression to use when passing a collection (or plain-missing) parameter + /// 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) + { + if (!parameter.IsCollection) + return parameter.Name; + var memberName = "coll_" + parameter.CollectionElementMemberName!; + return parameter.CollectionKind switch + { + 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 + /// for a localized (collection) parameter. + /// + private static string CollectionDictExpression(CollectionKind kind, string factoryName) => + kind switch + { + CollectionKind.Array => $"{factoryName}.ToArray()", + CollectionKind.List => $"{factoryName}.ToList()", + CollectionKind.ImmutableArray => $"ImmutableArray.CreateRange({factoryName})", + _ => factoryName, // Enumerable → direct + }; + } +} diff --git a/FactoryGenerator/Generation/FactoryGenerator.ContainerCodeTemplates.cs b/FactoryGenerator/Generation/FactoryGenerator.ContainerCodeTemplates.cs new file mode 100644 index 0000000..3b8802e --- /dev/null +++ b/FactoryGenerator/Generation/FactoryGenerator.ContainerCodeTemplates.cs @@ -0,0 +1,723 @@ +using System.Collections.Generic; +using System.Collections.Immutable; +using System.Linq; +using System.Text; +using Microsoft.CodeAnalysis; + +namespace FactoryGenerator +{ + /// + /// Renders the generated DependencyInjectionContainer/LifetimeScope partial-class + /// source: the dictionary-based lookup, constructors, member declarations, and collection + /// (IEnumerable<T>) declarations. + /// + /// LifetimeScope is a thin subclass of DependencyInjectionContainer, not a + /// hand-duplicated sibling: every declaration emitted here (lookup dictionary, factory members, + /// collection accessors) is written once and inherited by both. The only place root-vs-scope + /// behavior actually differs is singleton storage, which + /// resolves per-instance via m_singletonOwner rather than via virtual dispatch or a second + /// copy of every member. + /// + public partial class FactoryGenerator + { + private static void GenerateCode(InjectionAnalysis analysis, Compilation compilation, ILogger log, SourceProductionContext context) + { + var ordered = analysis.Ordered; + var interfaceInjectors = analysis.InterfaceInjectors; + var interfaceMemberNames = analysis.InterfaceMemberNames; + var availableInterfaceFullNames = analysis.AvailableInterfaceFullNames; + + // Computed once and reused by both cycle-detection and the declarations loop below, + // instead of each independently re-selecting a constructor/lambda member per injection. + var resolutions = ResolveInjections(ordered, availableInterfaceFullNames); + + CheckForCycles(interfaceInjectors, availableInterfaceFullNames, resolutions); + log.Log(LogLevel.Debug, "Starting Code Generation"); + var usingStatements = $@" +using System; +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; +#nullable enable"; + + var lookup = $@"{usingStatements} +[GeneratedCode(""{ToolName}"", ""{Version}"")] +#nullable enable +#pragma warning disable CS0169, CS0414 +// Not sealed: {LifetimeName} is a thin subclass (see Constructor.g.cs) that reuses every member +// declared here instead of duplicating them. Root-vs-scope singleton ownership is resolved via +// m_singletonOwner (see SymbolUtility.SingletonFactory), not virtual dispatch. +public partial class {ClassName} : IContainer, IContainerScopeFactory, IContainerRegistrationMetadata, IContainerCacheInvalidator, IAsyncDisposable, IContainerLocalCollectionResolver +{{ + +#pragma warning restore CS0169, CS0414 + private IContainer GetRoot() + {{ + IContainer root = this; + while(root.Base != null) + {{ + root = root.Base; + }} + return root; + }} + private IContainer GetTop() + {{ + IContainer top = this; + while(top.Inheritor != null) + {{ + top = top.Inheritor; + }} + 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(); + private readonly {ClassName} m_singletonOwner; + private Dictionary> m_lookup; + private Dictionary> m_localCollectionLookup; + private Dictionary m_booleans; + private readonly ResolvedInstanceTracker m_resolvedInstances = new(); + + internal void TrackResolvedInstance(object instance) => m_resolvedInstances.Track(instance); + + public bool TryResolveLocalCollection(Type type, out object? resolved) + {{ + if (m_localCollectionLookup.TryGetValue(type, out var factory)) + {{ + resolved = factory(); + return true; + }} + + resolved = default; + return false; + }} + + public T Resolve() + {{ + if (m_lookup.TryGetValue(typeof(T), out var factory)) + return (T)factory(); + if (Base is not null) + return Base.Resolve(); + throw new KeyNotFoundException($""The type {{typeof(T)}} has not been registered, and thus cannot be resolved""); + }} + + public object Resolve(Type type) + {{ + if (m_lookup.TryGetValue(type, out var factory)) + return factory(); + if (Base is not null) + return Base.Resolve(type); + throw new KeyNotFoundException($""The type {{type}} has not been registered, and thus cannot be resolved""); + }} + + public void Dispose() + {{ + DetachFromBase(); + m_resolvedInstances.Dispose(); + }} + + public ValueTask DisposeAsync() + {{ + DetachFromBase(); + return m_resolvedInstances.DisposeAsync(); + }} + + public bool TryResolve(Type type, out object? resolved) + {{ + if(m_lookup.TryGetValue(type, out var factory)) + {{ + resolved = factory(); + return true; + }} + if(Base is not null) + return Base.TryResolve(type, out resolved); + resolved = default; + return false; + }} + + public bool TryResolve(out T? resolved) + {{ + if(m_lookup.TryGetValue(typeof(T), out var factory)) + {{ + resolved = (T)factory(); + return true; + }} + if(Base is not null) + return Base.TryResolve(out resolved); + resolved = default; + return false; + }} + public bool IsRegistered(Type type) + {{ + return m_lookup.ContainsKey(type) || Base?.IsRegistered(type) == true; + }} + public bool IsRegistered() => IsRegistered(typeof(T)); + public bool GetBoolean(string key) + {{ + return m_booleans.TryGetValue(key, out var value) && value; + }} + public IEnumerable<(string Key, bool Value)> GetBooleans() + {{ + foreach(var pair in m_booleans) + {{ + yield return (pair.Key, pair.Value); + }} + }} +}}"; + context.AddSource($"FactoryGenerator.{ClassName}/Lookup.g.cs", lookup); + + var booleanKeys = analysis.RawInjections.Select(inj => inj.BooleanInjection) + .Where(b => b is not null) + .Select(b => b!.Key) + .Distinct() + .ToArray(); + + var declarations = new Dictionary(); + var constructorParameters = new List(); + var seenConstructorParameterKeys = new HashSet(); + + foreach (var injection in ordered) + { + var resolution = resolutions[injection]; + declarations[injection.Name] = Declaration(injection, resolution.Creation); + + foreach (var param in resolution.MissingParameters) + { + var key = param.TypeFullName + " " + param.Name; + if (seenConstructorParameterKeys.Add(key)) + constructorParameters.Add(param); + } + } + + // A single partitioning pass instead of two ToArray()+List.Remove() loops: each Remove() + // is an O(n) scan of its own, so removing k matches from an n-item list cost O(n*k) for + // no reason. Order matters here — a collection-typed parameter is classified as + // "localized" before the IContainer/self check even runs, matching the original two-pass + // precedence (a parameter can't reach the IContainer check once collection-classified). + var externalParameterCandidates = constructorParameters; + constructorParameters = new List(externalParameterCandidates.Count); + var localizedParameters = new List(); + foreach (var parameter in externalParameterCandidates) + { + if (parameter.IsCollection && parameter.CollectionElementFullName is not null) + { + localizedParameters.Add(parameter); + continue; + } + + if (parameter.TypeFullName.Contains("IContainer")) + { + log.Log(LogLevel.Debug, $"Registering {parameter.Name} as Self"); + declarations[parameter.Name] = $"private IContainer {parameter.Name} => this;"; + continue; + } + + constructorParameters.Add(parameter); + } + + ValidateExternalParameterTypes(constructorParameters); + + 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))) + .Concat(ordered.Select(injection => injection.LazyFieldName)) + .Concat(new[] + { + "Base", + "Inheritor", + "GetRoot", + "GetTop", + "Dispose", + "Resolve", + "TryResolve", + "TryResolveLocalCollection", + "IsRegistered", + "GetBoolean", + "GetBooleans", + "BeginLifetimeScope", + "DisposeAsync", + "TrackResolvedInstance", + "InitializeLookup", + "m_resolvedInstances", + "m_localCollectionLookup", + "m_lock", + "m_lookup", + "m_booleans", + "m_singletonOwner", + "singletonOwner", + "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]; + var ifaceMember = interfaceMemberNames[ifaceFull]; + var ifaceMethodName = ifaceMember + "()"; + + if (possibilities.All(i => i.BooleanInjection == null)) + { + var chosen = possibilities.Last(); + if (ifaceMethodName == chosen.Name) continue; + if (declarations.ContainsKey(ifaceMethodName)) continue; + log.Log(LogLevel.Information, $"Selecting {chosen.Name} for {ifaceFull}"); + declarations[ifaceMethodName] = $"internal {ifaceFull} {ifaceMethodName} => {chosen.Name};"; + } + else + { + var ternary = BuildBooleanSelectionExpression( + ifaceFull, + possibilities, + booleanIdentifiers, + possibility => possibility.Name); + + if (declarations.ContainsKey(ifaceMethodName)) continue; + log.Log(LogLevel.Information, $"Selecting {ternary} for {ifaceFull}"); + declarations[ifaceMethodName] = $"internal {ifaceFull} {ifaceMethodName} => {ternary};"; + } + } + + var arrayDeclarations = new Dictionary(); + foreach (var pair in interfaceInjectors) + { + 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); + } + + foreach (var parameter in localizedParameters) + { + 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); + } + + 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 constructor = "(" + string.Join(", ", allArguments) + ")"; + + 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", + 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 = 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 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; + context.AddSource($"FactoryGenerator.{ClassName}/Constructor.g.cs", Constructor(usingStatements, constructorFields, + constructor, constructorAssignments, + dictSize, interfacePairs, localizedPairs, enumerablePairs, constructorPairs, localCollectionPairs, + externalParameters, resolvedConstructorAssignments, booleanParameters, allArguments)); + Declarations(usingStatements, declarations, ClassName, context); + ArrayDeclarations(usingStatements, arrayDeclarations, ClassName, context); + + // LifetimeScope: a thin subclass of {ClassName} (see the Constructor above for the + // protected scope constructor it forwards to). It inherits the lookup dictionary, every + // factory member, and every collection accessor unchanged — no duplicate declarations. + var scopeConstructorParameters = string.Join(", ", + new[] {$"{ClassName} singletonOwner", "IContainer? baseContainer"}.Concat(allArguments)); + var scopeBaseArguments = string.Join(", ", + new[] {"singletonOwner", "baseContainer"}.Concat(allArguments.Select(arg => arg.Split(' ').Last()))); + context.AddSource($"FactoryGenerator.{LifetimeName}/{LifetimeName}.g.cs", $@"{usingStatements} +[GeneratedCode(""{ToolName}"", ""{Version}"")] +#nullable enable +internal sealed class {LifetimeName} : {ClassName} +{{ + internal {LifetimeName}({scopeConstructorParameters}) : base({scopeBaseArguments}) + {{ + }} +}} +"); + + // Emit the static factory + module initializer for plugin container registration + context.AddSource("FactoryGenerator.Helpers/Entrypoint.g.cs", $@" +using System; +using System.Runtime.CompilerServices; +using FactoryGenerator; + +#if !NET5_0_OR_GREATER +namespace System.Runtime.CompilerServices +{{ + [AttributeUsage(AttributeTargets.Method, AllowMultiple = false)] + internal sealed class ModuleInitializerAttribute : Attribute {{ }} +}} +#endif + +namespace {compilation.Assembly.Name}.Generated +{{ + /// + /// Provides a static factory for the generated container and auto-registers it in the ContainerRegistry on assembly load. + /// + public static class ContainerEntryPoint + {{ + /// + /// Creates a new DependencyInjectionContainer that chains on top of the given base container. + /// + public static IContainer Create(IContainer baseContainer) + {{ + return new {ClassName}(baseContainer); + }} + + /// + /// The assembly name this container was generated for. + /// + public static string AssemblyName => ""{compilation.Assembly.Name}""; + + [ModuleInitializer] + internal static void Register() + {{ + ContainerRegistry.Register(""{compilation.Assembly.Name}"", Create); + }} + }} +}} +"); + } + + /// + /// Builds the Constructor.g.cs fragment for : the constructor + /// fields, the three constructor overloads (root, cross-assembly "merging", and the + /// protected scope overload used only by ), and the shared + /// BeginLifetimeScope implementation. A single InitializeLookup helper builds + /// m_lookup/m_localCollectionLookup identically for all three, since the + /// factory closures they capture already resolve correctly against whichever instance + /// ( or ) constructed them. + /// + private static string Constructor(string usingStatements, string constructorFields, string constructor, string constructorAssignments, int dictSize, + List<(string TypeName, string MemberName)> interfaceTypePairs, List<(string TypeName, string Expression)> localizedParamPairs, + List<(string TypeName, string Expression)> enumerablePairs, List<(string TypeName, string Expression)> constructorParamPairs, + List<(string TypeName, string Expression)> localCollectionPairs, + List externalParameters, string resolvingConstructorAssignments, + IReadOnlyList<(string Key, string Identifier)> booleans, List allArguments) + { + var booleanDictionaryEntries = string.Join("\n", booleans.Select(boolean => $"\t\t{{ \"{boolean.Key}\", {boolean.Identifier} }},")); + var booleanFieldsFromBase = string.Join("\n", booleans.Select(boolean => $"\t this.{boolean.Identifier} = Base.GetBoolean(\"{boolean.Key}\");")); + + var scopeArgumentValues = new List {"m_singletonOwner", "baseContainer"}; + scopeArgumentValues.AddRange(booleans.Select(boolean => boolean.Identifier)); + scopeArgumentValues.AddRange(externalParameters.Select(parameter => $"baseContainer != null ? baseContainer.Resolve<{parameter.TypeFullName}>() : {parameter.Name}")); + + return $@"{usingStatements} +#pragma warning disable CS8618 // m_lookup/m_localCollectionLookup are always assigned by InitializeLookup(), called from every constructor. +public partial class {ClassName} +{{ + {constructorFields} + + public {ClassName}{constructor} + {{ + m_singletonOwner = this; + {constructorAssignments} + + m_booleans = new({booleans.Count}) {{ +{booleanDictionaryEntries} + }}; + InitializeLookup(); + }} + + /// + /// Cross-assembly composition constructor: chains this assembly's own container on top of + /// another (possibly different assembly's) via , + /// inheriting its booleans and resolving any of this assembly's external parameters from it. + /// Unrelated to lifetime scoping. + /// + public {ClassName}(IContainer Base) + {{ + m_singletonOwner = this; + this.Base = Base; + AttachToBase(Base); + {resolvingConstructorAssignments} + +{booleanFieldsFromBase} + + m_booleans = new(); + foreach(var (key, value) in Base.GetBooleans()) + {{ + m_booleans[key] = value; + }} + InitializeLookup(); + }} + + /// + /// Used only by . is always the + /// original root container (see ), so every scope — + /// however many are created, and regardless of nesting — shares exactly one singleton owner. + /// + protected {ClassName}({ClassName} singletonOwner, IContainer? baseContainer{(allArguments.Count > 0 ? ", " + string.Join(", ", allArguments) : string.Empty)}) + {{ + m_singletonOwner = singletonOwner; + this.Base = baseContainer; + if (baseContainer is not null) + {{ + AttachToBase(baseContainer); + TrackResolvedInstance(baseContainer); + }} + {constructorAssignments} + + m_booleans = new({booleans.Count}) {{ +{booleanDictionaryEntries} + }}; + InitializeLookup(); + }} + + public ILifetimeScope BeginLifetimeScope() + {{ + var baseContainer = Base?.BeginLifetimeScope() as IContainer; + return BeginLifetimeScope(baseContainer); + }} + + public ILifetimeScope BeginLifetimeScope(IContainer? baseContainer) + {{ + var scope = new {LifetimeName}({string.Join(", ", scopeArgumentValues)}); + TrackResolvedInstance(scope); + return scope; + }} + + private void InitializeLookup() + {{ + m_lookup = new({dictSize}) {{ +{MakeDictionaryFromTypes(interfaceTypePairs)} +{MakeDictionaryFromParams(localizedParamPairs)} +{MakeDictionaryFromParams(enumerablePairs)} +{MakeDictionaryFromParams(constructorParamPairs)} + }}; + m_localCollectionLookup = new({localCollectionPairs.Count}) {{ +{MakeDictionaryFromParams(localCollectionPairs)} + }}; + }} +}} +#pragma warning restore CS8618"; + } + + + private static void ArrayDeclarations(string usingStatements, Dictionary arrayDeclarations, string className, SourceProductionContext context) + { + var cacheInvalidations = string.Join("\n ", arrayDeclarations.Keys.Select(name => $"m_{name} = null;")); + foreach (var declaration in arrayDeclarations) + { + context.AddSource($"FactoryGenerator.Collection.Declarations/{declaration.Key}.g.cs", $@"{usingStatements} +public partial class {className} +{{ + {declaration.Value} +}}"); + } + + context.AddSource($"FactoryGenerator.{className}/Collection_Invalidation.g.cs", $@"{usingStatements} +public partial class {className} +{{ + public void InvalidateCollectionCaches() + {{ + {cacheInvalidations} + }} +}}"); + } + + private static void Declarations(string usingStatements, Dictionary declarations, string className, SourceProductionContext context) + { + foreach (var declaration in declarations) + { + context.AddSource($"FactoryGenerator.Declarations/{declaration.Key}.g.cs", + $$""" + {{usingStatements}} + public partial class {{className}} + { + {{declaration.Value}} + }; + """); + } + } + + private static void MakeArray(Dictionary declarations, string name, + string elementTypeFullName, Dictionary> interfaceInjectors, + IReadOnlyDictionary booleanIdentifiers) + { + 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} + {{ + 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(!(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(!(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} + {{ + get + {{ + var cached = m_{name}; + if (cached != null) + return cached; + + lock (m_lock) + {{ + cached = m_{name}; + if (cached != null) + return cached; + return m_{name} = {factoryName}; + }} + }} + }} + internal IEnumerable<{elementTypeFullName}> local_{name} => {localFactoryName}; + internal IEnumerable<{elementTypeFullName}>? m_{name};" + factory; + } + + private static string MakeDictionaryFromTypes(IEnumerable<(string TypeName, string MemberName)> pairs) + { + var builder = new StringBuilder(); + foreach (var (typeName, memberName) in pairs) + builder.AppendLine($"\t\t\t{{ typeof({typeName}),{memberName} }},"); + return builder.ToString(); + } + + private static string MakeDictionaryFromParams(IEnumerable<(string TypeName, string Expression)> pairs) + { + var builder = new StringBuilder(); + foreach (var (typeName, expression) in pairs) + builder.AppendLine($"\t\t\t{{ typeof({typeName}), () => {expression} }},"); + return builder.ToString(); + } + } +} \ No newline at end of file diff --git a/FactoryGenerator/Generation/FactoryGenerator.CycleDetection.cs b/FactoryGenerator/Generation/FactoryGenerator.CycleDetection.cs new file mode 100644 index 0000000..eef8b31 --- /dev/null +++ b/FactoryGenerator/Generation/FactoryGenerator.CycleDetection.cs @@ -0,0 +1,160 @@ +using System; +using System.Collections.Generic; +using System.Collections.Immutable; +using System.Linq; +using System.Text; +using System.Threading; +using System.Threading.Tasks; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using Microsoft.CodeAnalysis.Diagnostics; + +namespace FactoryGenerator +{ + /// + /// Detects cyclic dependencies between injections ahead of code generation, so a misconfigured + /// dependency graph fails fast with a readable diagnostic instead of producing code that + /// would recurse infinitely (or fail to compile) at runtime. + /// + public partial class FactoryGenerator + { + private static List GetReachableImplementations(List possibilities) + { + if (possibilities.Count == 0) + return new List(); + + 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, InjectionResolution resolution, HashSet availableInterfaceFullNames) + { + 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; + } + + if (resolution.ConstructorParameters is not { } parameters) + yield break; + + foreach (var dependency in GetParameterDependencies(parameters, availableInterfaceFullNames)) + yield return dependency; + } + + private static IEnumerable GetParameterDependencies( + ImmutableArray parameters, + HashSet availableInterfaceFullNames) + { + foreach (var parameter in 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 CheckForCycles( + Dictionary> interfaceInjectors, + HashSet availableInterfaceFullNames, + IReadOnlyDictionary resolutions) + { + 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, resolutions[injection], availableInterfaceFullNames)) + deps.Add(dep); + } + + graph[ifaceName] = deps; + nodeOwner[ifaceName] = reachable; + } + + 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. The owner description is only ever + // needed here, on the rare path where a cycle actually exists, so it's built + // lazily instead of unconditionally for every interface on every run. + var cycleStart = path.IndexOf(dep); + var cycle = path.GetRange(cycleStart, path.Count - cycleStart); + cycle.Add(dep); + var owner = nodeOwner.TryGetValue(node, out var reachable) + ? string.Join(", ", reachable.Select(injection => injection.TypeFullName).Distinct()) + : 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 + } + } +} diff --git a/FactoryGenerator/Generation/FactoryGenerator.IdentifierNaming.cs b/FactoryGenerator/Generation/FactoryGenerator.IdentifierNaming.cs new file mode 100644 index 0000000..2b03847 --- /dev/null +++ b/FactoryGenerator/Generation/FactoryGenerator.IdentifierNaming.cs @@ -0,0 +1,171 @@ +using System; +using System.Collections.Generic; +using System.Collections.Immutable; +using System.Linq; +using System.Text; +using System.Threading; +using System.Threading.Tasks; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using Microsoft.CodeAnalysis.Diagnostics; + +namespace FactoryGenerator +{ + /// + /// Helpers for producing valid, collision-free C# identifiers for generated boolean + /// parameters, externally-supplied constructor parameters, and boolean-gated selection + /// expressions shared between the container and static-extension code generators. + /// + public partial class FactoryGenerator + { + private static void ValidateExternalParameterTypes(List constructorParameters) + { + var ambiguousParameters = constructorParameters + .GroupBy(parameter => parameter.TypeFullName) + .Select(group => new + { + 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 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_"); + } + + 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 string BuildMissingImplementationExpression(string typeFullName) + { + return $"throw new global::System.InvalidOperationException(\"Cannot resolve {typeFullName} without a matching implementation\")"; + } + } +} diff --git a/FactoryGenerator/Generation/FactoryGenerator.InjectionOrdering.cs b/FactoryGenerator/Generation/FactoryGenerator.InjectionOrdering.cs new file mode 100644 index 0000000..10677b3 --- /dev/null +++ b/FactoryGenerator/Generation/FactoryGenerator.InjectionOrdering.cs @@ -0,0 +1,190 @@ +using System; +using System.Collections.Generic; +using System.Collections.Immutable; +using System.Linq; +using System.Text; +using System.Threading; +using System.Threading.Tasks; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using Microsoft.CodeAnalysis.Diagnostics; + +namespace FactoryGenerator +{ + /// + /// Orders discovered injections (by assembly priority/distance) and indexes them by the + /// interfaces they can satisfy, forming the basis for interface-to-implementation selection. + /// + public partial class FactoryGenerator + { + /// + /// The result of analyzing the full set of discovered injections once: assembly-priority + /// ordering plus the interface-to-implementation index derived from it. Both the + /// dictionary-based container generator and the static-extensions generator need an + /// identical copy of this; computing it once here — as a single incremental-pipeline stage + /// (see ) — means it is derived exactly once per compilation instead + /// of once per consumer. + /// + private sealed class InjectionAnalysis : IEquatable + { + public InjectionAnalysis( + ImmutableArray rawInjections, + ImmutableArray ordered, + Dictionary> interfaceInjectors, + Dictionary interfaceMemberNames, + HashSet availableInterfaceFullNames) + { + RawInjections = rawInjections; + Ordered = ordered; + InterfaceInjectors = interfaceInjectors; + InterfaceMemberNames = interfaceMemberNames; + AvailableInterfaceFullNames = availableInterfaceFullNames; + } + + /// + /// Injections in original discovery order, exactly as produced by FindMethods. + /// The dictionary-based container generator derives its (positional) boolean-parameter + /// order from the first-seen order in this sequence — must not be + /// substituted for it, since re-sorting would silently reorder generated constructor + /// parameters. + /// + public ImmutableArray RawInjections { get; } + + /// Injections in final assembly-priority/distance order. + public ImmutableArray Ordered { get; } + + /// Interface full name → the injections that can satisfy it, in priority order. + public Dictionary> InterfaceInjectors { get; } + + /// Interface full name → its generated member name (parallel to ). + public Dictionary InterfaceMemberNames { get; } + + /// Every interface full name any injection can satisfy — 's keys as a set. + public HashSet AvailableInterfaceFullNames { get; } + + // Equality (and hashing) is defined purely in terms of RawInjections: everything else + // (Ordered, InterfaceInjectors, InterfaceMemberNames, AvailableInterfaceFullNames) is a + // deterministic pure function of RawInjections plus the Compilation that BuildInjectionAnalysis + // was invoked with (already tracked separately by the incremental pipeline). Comparing the + // raw, order-sensitive sequence — rather than the re-sorted Ordered sequence — is required + // for correctness: two different discovery orders can sort into an identical Ordered + // sequence while still needing different generated boolean-parameter ordering. + public bool Equals(InjectionAnalysis? other) + { + if (other is null) return false; + if (ReferenceEquals(this, other)) return true; + return RawInjections.SequenceEqual(other.RawInjections); + } + + public override bool Equals(object? obj) => obj is InjectionAnalysis other && Equals(other); + + public override int GetHashCode() + { + var hash = RawInjections.Length; + foreach (var injection in RawInjections) + hash = (hash * 397) ^ injection.GetHashCode(); + return hash; + } + } + + /// + /// Builds the shared for a compilation's discovered + /// injections. This is the single incremental-pipeline stage both + /// and consume (see ) instead + /// of each independently calling /. + /// + private static InjectionAnalysis BuildInjectionAnalysis(ImmutableArray dataInjections, Compilation compilation, CancellationToken token) + { + var ordered = OrderInjections(dataInjections, compilation).ToImmutableArray(); + token.ThrowIfCancellationRequested(); + var (interfaceInjectors, interfaceMemberNames) = BuildInterfaceInjectors(ordered); + var availableInterfaceFullNames = new HashSet(interfaceInjectors.Keys); + return new InjectionAnalysis(dataInjections, ordered, interfaceInjectors, interfaceMemberNames, availableInterfaceFullNames); + } + + private static List OrderInjections(ImmutableArray dataInjections, Compilation compilation) + { + var ordered = dataInjections.Reverse().ToList(); + var assemblyDistances = BuildAssemblyDistances(compilation, ordered.Select(injection => injection.AssemblyName)); + + 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) + { + 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 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) + { + 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.TryGetValue(ifaceFull, out var injectors)) + { + injectors = new List(); + interfaceInjectors[ifaceFull] = injectors; + interfaceMemberNames[ifaceFull] = ifaceMember; + } + + injectors.Add(injection); + } + } + + return (interfaceInjectors, interfaceMemberNames); + } + } +} diff --git a/FactoryGenerator/Generation/FactoryGenerator.StaticExtensionsGenerator.cs b/FactoryGenerator/Generation/FactoryGenerator.StaticExtensionsGenerator.cs new file mode 100644 index 0000000..1c0a904 --- /dev/null +++ b/FactoryGenerator/Generation/FactoryGenerator.StaticExtensionsGenerator.cs @@ -0,0 +1,807 @@ +using System; +using System.Collections.Generic; +using System.Collections.Immutable; +using System.Linq; +using System.Text; +using System.Threading; +using System.Threading.Tasks; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using Microsoft.CodeAnalysis.Diagnostics; + +namespace FactoryGenerator +{ + /// + /// Generates the C# 14+ static-extension resolve API (T.Resolve(container)), an + /// alternative to the dictionary-based container that inlines the full construction chain + /// directly at each call site, avoiding dictionary lookups and factory-delegate indirection. + /// + public partial class FactoryGenerator + { + 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 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 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, + ((InjectionAnalysis Analysis, Compilation Compilation) Left, bool SupportsExtensions) data) + { + if (!data.SupportsExtensions) return; + GenerateStaticExtensions(data.Left.Analysis, data.Left.Compilation, context); + } + + private static void GenerateStaticExtensions(InjectionAnalysis analysis, Compilation compilation, SourceProductionContext context) + { + var interfaceInjectors = analysis.InterfaceInjectors; + var interfaceMemberNames = analysis.InterfaceMemberNames; + var availableInterfaces = analysis.AvailableInterfaceFullNames; + + // Shared across spec-building (PopulateStaticExtensionDirectRequirements) and code + // emission (BuildStaticCreateExpression) below, so each injection's constructor is + // selected once instead of twice — see GetCachedBestConstructor. + var constructorCache = new Dictionary(); + var specs = BuildStaticExtensionSpecs(interfaceInjectors, interfaceMemberNames, availableInterfaces, constructorCache); + + 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 usings = $""" + using System; + using System.CodeDom.Compiler; + using System.Collections.Generic; + using System.Collections.Immutable; + using System.Linq; + namespace {compilation.Assembly.Name}.Generated; + #nullable enable + """; + var state = $$""" + {{usings}} + [GeneratedCode("{{ToolName}}", "{{Version}}")] + 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); + } + } + """; + context.AddSource("FactoryGenerator.StaticExtensions/StaticResolveState.g.cs", state); + + foreach (var ifaceFull in interfaceMemberNames.Keys) + { + var spec = specs[ifaceFull]; + var helpers = BuildStaticExtensionClass(spec, specs, availableInterfaces, booleanIdentifiers, externalIdentifiers, constructorCache); + + context.AddSource($"FactoryGenerator.StaticExtensions/{spec.ExtensionClassName}.g.cs", + $$""" + {{usings}} + [GeneratedCode("{{ToolName}}", "{{Version}}")] + public static class {{spec.ExtensionClassName}} + { + {{helpers}} + extension({{spec.TypeFullName}}) + { + {{BuildStaticPublicResolveMethods(spec, booleanIdentifiers, externalIdentifiers)}} + } + } + """); + } + } + + private static Dictionary BuildStaticExtensionSpecs( + Dictionary> interfaceInjectors, + Dictionary interfaceMemberNames, + HashSet availableInterfaces, + Dictionary constructorCache) + { + var specs = interfaceInjectors.ToDictionary( + pair => pair.Key, + pair => CreateDirectStaticExtensionSpec(pair.Key, pair.Value, interfaceMemberNames[pair.Key], availableInterfaces, interfaceInjectors, constructorCache), + StringComparer.Ordinal); + + PropagateStaticExtensionRequirements(specs); + + foreach (var spec in specs.Values) + { + 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, + HashSet availableInterfaces, + Dictionary> interfaceInjectors, + Dictionary constructorCache) + { + var spec = new StaticExtensionSpec(typeFullName, typeMemberName, typeMemberName + "Extensions", possibilities); + PopulateStaticExtensionDirectRequirements(spec, availableInterfaces, interfaceInjectors, constructorCache); + return spec; + } + + private static void PopulateStaticExtensionDirectRequirements( + StaticExtensionSpec spec, + HashSet availableInterfaces, + Dictionary> interfaceInjectors, + Dictionary constructorCache) + { + 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 { } lambda) + { + if (interfaceInjectors.ContainsKey(lambda.ContainingTypeFullName)) + AddDistinctDependency(spec.Dependencies, new StaticDependencyReference(lambda.ContainingTypeFullName, false)); + + foreach (var parameter in lambda.MethodParameters) + AddStaticParameterRequirement(spec, parameter, interfaceInjectors); + + continue; + } + + var constructor = GetCachedBestConstructor(possibility, availableInterfaces, constructorCache); + 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) + { + 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); + } + } + } 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, + HashSet availableInterfaces, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers, + Dictionary constructorCache) + { + 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, constructorCache)); + } + + return string.Join("\n\n", parts); + } + + 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 + { + $@" 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 + {{ + var source = new List<{spec.TypeFullName}>({resolveInvocations.Length}) {{ + {string.Join(",\n ", resolveInvocations)} + }}; +{string.Join("\n", conditionalInvocations)} + if (container is not null) + {{ + var b = container.Base; + while (b is not null) + {{ + if (b.TryResolve>(out var additional)) + source.AddRange(additional!); + b = b.Base; + }} + + b = container.Inheritor; + while (b is not null) + {{ + if (b.TryResolve>(out var additional)) + source.AddRange(additional!); + b = b.Inheritor; + }} + }} + + 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); + var tracksResolvedInstance = injection.Disposable || injection.AsyncDisposable; + + if (injection.Singleton || injection.Scoped) + { + if (tracksResolvedInstance) + { + return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) + {{ + if (container is not null) + {{ + 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.TrackResolvedInstance(value); + container.{injection.LazyFieldName} = value; + return value; + }} + }} + + return Create_{helperName}({BuildStaticInternalInvocationArguments(spec, booleanIdentifiers, externalIdentifiers, "null", "state")}); + }}"; + } + + return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) + {{ + if (container is not null) + {{ + 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")}); + }}"; + } + + if (tracksResolvedInstance) + { + return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) + {{ + if (container is not null) + {{ + var value = {createInvocation}; + container.TrackResolvedInstance(value); + return value; + }} + + return Create_{helperName}({BuildStaticInternalInvocationArguments(spec, booleanIdentifiers, externalIdentifiers, "null", "state")}); + }}"; + } + + return $@" private static {injection.TypeFullName} Resolve_{helperName}({parameterList}) + {{ + return {createInvocation}; + }}"; + } + + private static string BuildStaticCreateInjectionMethod( + StaticExtensionSpec spec, + InjectionData injection, + IReadOnlyDictionary specs, + HashSet availableInterfaces, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers, + Dictionary constructorCache) + { + var helperName = GetStaticInjectionHelperName(injection); + var parameterList = BuildStaticMethodParameterList(spec, booleanIdentifiers, externalIdentifiers, includeContainer: true, includeState: true); + var createExpression = BuildStaticCreateExpression(injection, specs, availableInterfaces, booleanIdentifiers, externalIdentifiers, constructorCache); + var returnStatement = createExpression.StartsWith("throw ", StringComparison.Ordinal) + ? createExpression + ";" + : "return " + createExpression + ";"; + + return $@" private static {injection.TypeFullName} Create_{helperName}({parameterList}) + {{ + {returnStatement} + }}"; + } + + /// + /// Looks up (or computes and caches) the best constructor for an injection. The + /// static-extensions pipeline needs this exact same, deterministic choice twice — once while + /// building each interface's (to discover its + /// external-parameter/boolean/dependency requirements, see + /// ) and again while emitting the + /// Create_X method body () — caching it here + /// avoids re-running 's per-constructor parameter analysis a + /// second time for the same injection. missing/nullableDefaults are unused by + /// either static-extensions call site (both only need the chosen constructor's parameter + /// list), so they aren't cached. + /// + private static ConstructorData? GetCachedBestConstructor( + InjectionData injection, + HashSet availableInterfaces, + Dictionary constructorCache) + { + if (constructorCache.TryGetValue(injection, out var cached)) + return cached; + + HashSet? missing = null; + HashSet? nullableDefaults = null; + var chosen = GetBestConstructor(injection, availableInterfaces, ref missing, ref nullableDefaults); + constructorCache[injection] = chosen; + return chosen; + } + + private static string BuildStaticCreateExpression( + InjectionData injection, + IReadOnlyDictionary specs, + HashSet availableInterfaces, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers, + Dictionary constructorCache) + { + if (injection.Lambda is LambdaData lambda) + { + if (!specs.TryGetValue(lambda.ContainingTypeFullName, out var containingSpec)) + return BuildMissingImplementationExpression(lambda.ContainingTypeFullName); + + var containingInvocation = BuildStaticResolveInvocation(containingSpec, "ResolveCore", booleanIdentifiers, externalIdentifiers); + if (!lambda.IsMethod) + return $"{containingInvocation}.{lambda.MemberName}"; + + var lambdaArguments = BuildStaticArgumentList(lambda.MethodParameters, specs, booleanIdentifiers, externalIdentifiers); + return $"{containingInvocation}.{lambda.MemberName}({lambdaArguments})"; + } + + var constructor = GetCachedBestConstructor(injection, availableInterfaces, constructorCache); + if (constructor is null) + return BuildMissingImplementationExpression(injection.TypeFullName); + + var constructorArguments = BuildStaticArgumentList(constructor.Parameters, specs, booleanIdentifiers, externalIdentifiers); + return $"new {injection.TypeFullName}({constructorArguments})"; + } + + 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 null) + { + useNamedArguments = true; + continue; + } + + arguments.Add(useNamedArguments + ? $"{parameter.Name}: {argumentExpression}" + : argumentExpression); + } + + return string.Join(", ", arguments); + } + + 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)) + { + 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( + parameter, + BuildStaticResolveAllInvocation(collectionSpec, booleanIdentifiers, externalIdentifiers)); + } + + var typeLookup = parameter.IsNullable + ? parameter.TypeFullName.TrimEnd('?') + : parameter.TypeFullName; + + if (specs.TryGetValue(typeLookup, out var dependencySpec)) + return BuildStaticResolveInvocation(dependencySpec, "ResolveCore", booleanIdentifiers, externalIdentifiers); + + if (parameter.HasExplicitDefault || parameter.IsParams) + return null; + + if (parameter.IsNullable) + return "null"; + + if (externalIdentifiers.TryGetValue(parameter.TypeFullName, out var identifier)) + return identifier; + + return BuildMissingImplementationExpression(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")})"; + } + + private static string BuildStaticResolveSelectionExpression( + StaticExtensionSpec spec, + IReadOnlyDictionary booleanIdentifiers, + IReadOnlyDictionary externalIdentifiers) + { + return BuildBooleanSelectionExpression( + spec.TypeFullName, + spec.Possibilities, + booleanIdentifiers, + possibility => BuildStaticResolveInjectionInvocation(spec, possibility, booleanIdentifiers, externalIdentifiers)); + } + + 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 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/Generation/FactoryGenerator.TypeDiscovery.cs b/FactoryGenerator/Generation/FactoryGenerator.TypeDiscovery.cs new file mode 100644 index 0000000..02f729d --- /dev/null +++ b/FactoryGenerator/Generation/FactoryGenerator.TypeDiscovery.cs @@ -0,0 +1,219 @@ +using System; +using System.Collections.Generic; +using System.Collections.Immutable; +using System.Linq; +using System.Text; +using System.Threading; +using System.Threading.Tasks; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using Microsoft.CodeAnalysis.Diagnostics; + +namespace FactoryGenerator +{ + /// + /// Discovers candidate types for injection. Scopes the scan to the small set of assemblies + /// that can actually declare a FactoryGenerator attribute, instead of the compilation's full + /// merged global namespace (which otherwise includes the entire referenced BCL). + /// + public partial class FactoryGenerator + { + /// + /// The name of the assembly that declares the FactoryGenerator marker attributes + /// ([Inject], [Singleton], ...). Applying any of those attributes to a type, + /// method, or property requires the containing assembly to (directly or transitively) + /// reference this assembly, which lets us narrow the type-discovery scan below. + /// + private const string AttributesAssemblyName = "FactoryGenerator.Attributes"; + + /// + /// Scope used by to discover candidate types. When + /// is non-empty, only those assemblies' own type trees are + /// scanned. Otherwise (the assembly could not be located + /// in the reference graph) we conservatively fall back to , which + /// mirrors the original, unscoped behavior. + /// + private readonly struct InjectionScanScope + { + public InjectionScanScope(INamespaceSymbol globalNamespace, ImmutableArray relevantAssemblies) + { + GlobalNamespace = globalNamespace; + RelevantAssemblies = relevantAssemblies; + } + + public INamespaceSymbol GlobalNamespace { get; } + public ImmutableArray RelevantAssemblies { get; } + } + + private static InjectionScanScope GetInjectionScanScope(Compilation compilation, CancellationToken token) + { + return new InjectionScanScope(compilation.GlobalNamespace, GetRelevantAssemblies(compilation, token)); + } + + /// + /// A real-world compilation typically references the entire BCL (hundreds of assemblies, + /// tens of thousands of types) via compilation.GlobalNamespace, yet only assemblies that + /// can (transitively) reach are able to declare any + /// FactoryGenerator attribute at all — applying the attribute requires a reference to its + /// declaring assembly. We compute that small "relevant" subset up front so + /// can scan each relevant assembly's own (unmerged) type tree instead of the merged global one. + /// + private static ImmutableArray GetRelevantAssemblies(Compilation compilation, CancellationToken token) + { + var allReachableAssemblies = new HashSet(SymbolEqualityComparer.Default) { compilation.Assembly }; + // Captured once per assembly during the BFS below and reused by CanReachAttributesAssembly, + // instead of calling GetReferencedAssemblies a second time per assembly during that check. + var references = new Dictionary(SymbolEqualityComparer.Default); + var toVisit = new Queue(); + toVisit.Enqueue(compilation.Assembly); + while (toVisit.Count > 0) + { + token.ThrowIfCancellationRequested(); + var current = toVisit.Dequeue(); + var directReferences = GetReferencedAssemblies(current).ToArray(); + references[current] = directReferences; + foreach (var referenced in directReferences) + { + if (allReachableAssemblies.Add(referenced)) + toVisit.Enqueue(referenced); + } + } + + var attributesAssemblies = allReachableAssemblies.Where(assembly => assembly.Name == AttributesAssemblyName).ToArray(); + if (attributesAssemblies.Length == 0) + return ImmutableArray.Empty; // Signals FindMethods to fall back to the unscoped scan. + + var canReachAttributes = new Dictionary(SymbolEqualityComparer.Default); + + bool CanReachAttributesAssembly(IAssemblySymbol assembly) + { + if (canReachAttributes.TryGetValue(assembly, out var known)) + return known; + + canReachAttributes[assembly] = false; // Guards against re-entrancy; reference graphs are acyclic anyway. + if (Array.IndexOf(attributesAssemblies, assembly) >= 0) + return canReachAttributes[assembly] = true; + + foreach (var referenced in references[assembly]) + { + if (CanReachAttributesAssembly(referenced)) + return canReachAttributes[assembly] = true; + } + + return false; + } + + var relevant = ImmutableArray.CreateBuilder(); + foreach (var assembly in allReachableAssemblies) + { + token.ThrowIfCancellationRequested(); + if (CanReachAttributesAssembly(assembly)) + relevant.Add(assembly); + } + + return relevant.ToImmutable(); + } + + private static IEnumerable GetCandidateTypes(InjectionScanScope scope, CancellationToken token) + { + if (scope.RelevantAssemblies.IsDefaultOrEmpty) + return SymbolUtility.GetAllTypes(scope.GlobalNamespace); + + return scope.RelevantAssemblies.SelectMany(assembly => + { + token.ThrowIfCancellationRequested(); + return SymbolUtility.GetAllTypes(assembly.GlobalNamespace); + }); + } + + private static IEnumerable FindMethods(InjectionScanScope scope, CancellationToken token) + { + foreach (var type in GetCandidateTypes(scope, token)) + { + token.ThrowIfCancellationRequested(); + if (type.TypeKind != TypeKind.Class && type.TypeKind != TypeKind.Interface) continue; + + // Cheap, allocation-free match check first: the vast majority of scanned types have + // no FactoryGenerator attribute at all (directly or via an implemented interface), so + // avoid building the concatenated attribute array (Concat + ToImmutableArray) unless + // the type actually matches. + if (HasInjectionAttribute(type)) + { + var typeAttributes = type.GetAttributes().Concat(type.AllInterfaces.SelectMany(i => i.GetAttributes())) + .ToImmutableArray(); + var info = Injection.Create(type, typeAttributes, token); + if (info is not null) yield return info; + } + + // Single pass over the member list instead of two separate GetMembers().OfType<>().Where() + // LINQ chains. Matches are buffered per member kind and yielded methods-then-properties + // below to preserve today's exact discovery order (positional boolean constructor + // parameters depend on it). + List<(IMethodSymbol Method, ImmutableArray Attributes)>? methodMatches = null; + List<(IPropertySymbol Property, ImmutableArray Attributes)>? propertyMatches = null; + foreach (var member in type.GetMembers()) + { + if (member.DeclaredAccessibility != Accessibility.Public) continue; + switch (member) + { + case IMethodSymbol method: + { + var attributes = method.GetAttributes(); + if (attributes.Any(IsInjection)) + (methodMatches ??= new List<(IMethodSymbol, ImmutableArray)>()).Add((method, attributes)); + break; + } + case IPropertySymbol property: + { + var attributes = property.GetAttributes(); + if (attributes.Any(IsInjection)) + (propertyMatches ??= new List<(IPropertySymbol, ImmutableArray)>()).Add((property, attributes)); + break; + } + } + } + + if (methodMatches is not null) + { + foreach (var (method, attributes) in methodMatches) + { + var info = Injection.Create(method, attributes, token); + if (info is not null) yield return info; + } + } + + if (propertyMatches is not null) + { + foreach (var (property, attributes) in propertyMatches) + { + var info = Injection.Create(property, attributes, token); + if (info is not null) yield return info; + } + } + } + + bool IsInjection(AttributeData attribute) + { + return attribute.AttributeClass?.Name.Contains("Inject") == true && attribute.AttributeClass.ToString().StartsWith("FactoryGenerator.Attributes"); + } + + bool HasInjectionAttribute(INamedTypeSymbol type) + { + foreach (var attribute in type.GetAttributes()) + { + if (IsInjection(attribute)) return true; + } + + foreach (var iface in type.AllInterfaces) + { + foreach (var attribute in iface.GetAttributes()) + { + if (IsInjection(attribute)) return true; + } + } + + return false; + } + } + } +} diff --git a/FactoryGenerator/Injection.cs b/FactoryGenerator/Injection.cs index f9120fa..80d0477 100644 --- a/FactoryGenerator/Injection.cs +++ b/FactoryGenerator/Injection.cs @@ -92,24 +92,32 @@ public static class Injection } } - var interfaces = acquireChildInterfaces ? namedTypeSymbol.AllInterfaces : namedTypeSymbol.Interfaces; + var isDisposable = namedTypeSymbol.AllInterfaces.Any(i => i.SpecialType == SpecialType.System_IDisposable); + var isAsyncDisposable = namedTypeSymbol.AllInterfaces.Any(IsAsyncDisposableInterface); + + // Accumulated into a single mutable list rather than chaining Add/AddRange/Remove/ + // RemoveRange calls on an ImmutableArray, each of which would otherwise copy the whole + // backing array — a list gives amortized O(1) appends and the same O(n) removals + // (unavoidable either way) without the repeated whole-array copies in between. + var baseInterfaces = acquireChildInterfaces ? namedTypeSymbol.AllInterfaces : namedTypeSymbol.Interfaces; + var interfaceList = new List(baseInterfaces.Length + attributedInterfaces.Count + 1); + interfaceList.AddRange(baseInterfaces); if (asSelf) - interfaces = interfaces.Add(namedTypeSymbol); - interfaces = interfaces.AddRange(attributedInterfaces); + interfaceList.Add(namedTypeSymbol); + interfaceList.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); + var disposableIface = interfaceList.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"); + interfaceList.Remove(disposableIface); + var asyncDisposableIface = interfaceList.FirstOrDefault(IsAsyncDisposableInterface); if (asyncDisposableIface is not null) - interfaces = interfaces.Remove(asyncDisposableIface); + interfaceList.Remove(asyncDisposableIface); - interfaces = interfaces - .RemoveRange(preventedInterfaces) - .Distinct(SymbolEqualityComparer.Default) - .Cast() + foreach (var prevented in preventedInterfaces) + interfaceList.Remove(prevented); + + var interfaces = interfaceList + .Distinct((IEqualityComparer) SymbolEqualityComparer.Default) .ToImmutableArray(); var ifaceFullNames = interfaces.Select(i => i.ToString()!).ToImmutableArray(); @@ -178,6 +186,16 @@ private static ParameterData ExtractParameter(IParameterSymbol parameter) return null; } + /// + /// Checks whether an interface is System.IAsyncDisposable without paying for a full + /// display-string (ToString()) computation on every interface. Roslyn has no + /// entry for it (it postdates the "special type" list), so a name + /// check first is used to filter out the overwhelming majority of unrelated interfaces before + /// falling back to the (cheaper, non-generic) containing-namespace comparison. + /// + private static bool IsAsyncDisposableInterface(INamedTypeSymbol i) => + i.Name == "IAsyncDisposable" && i.ContainingNamespace?.ToDisplayString() == "System"; + private static int GetAssemblyPriority(IAssemblySymbol? assemblySymbol) { if (assemblySymbol is null) diff --git a/FactoryGenerator/SymbolUtility.cs b/FactoryGenerator/SymbolUtility.cs index 80ef8fb..20bef92 100644 --- a/FactoryGenerator/SymbolUtility.cs +++ b/FactoryGenerator/SymbolUtility.cs @@ -30,23 +30,14 @@ public static IEnumerable GetAllTypes(INamespaceSymbol root) public static IEnumerable GetAllTypes(INamedTypeSymbol root) { - foreach (var namespaceOrTypeSymbol in root.GetMembers()) + // GetTypeMembers() returns only nested types, unlike GetMembers() which would force + // materializing every field/method/property/event of `root` just to filter them back out. + // The vast majority of scanned types have zero nested types, so this avoids real work. + foreach (var type in root.GetTypeMembers()) { - switch (namespaceOrTypeSymbol) - { - case INamespaceSymbol @namespace: - { - foreach (var nested in GetAllTypes(@namespace)) - yield return nested; - break; - } - case INamedTypeSymbol type: - - foreach (var nested in GetAllTypes(type)) - yield return nested; - yield return type; - break; - } + foreach (var nested in GetAllTypes(type)) + yield return nested; + yield return type; } } @@ -101,13 +92,32 @@ public static string MemberName(ISymbol? type) return sb.ToString(); } - public static string SingletonFactory(string typeName, string name, string lazyName, string creation, bool disposable) + /// + /// Generates a lazily-initialized, double-checked-lock singleton/scoped factory member. + /// + /// + /// True only for genuine [Singleton] injections (never [Scoped]). When set, the + /// member first checks m_singletonOwner — which is this for a root container and + /// the original root for a LifetimeScope (see ) + /// — and forwards to it if it isn't this. This is what makes a single generated member + /// declaration (shared by DependencyInjectionContainer and its LifetimeScope + /// subclass) resolve to one shared singleton instance across every scope, without needing a + /// separate forwarding declaration duplicated into the scope class. + /// + public static string SingletonFactory(string typeName, string name, string lazyName, string creation, bool disposable, bool forwardToOwner) { + var ownerForward = forwardToOwner + ? $@" + if (m_singletonOwner != this) + return m_singletonOwner.{name}; + " + : string.Empty; + if (disposable) { return $@" internal {typeName} {name} - {{ + {{{ownerForward} var cached = {lazyName}; if (cached != null) return cached; @@ -128,7 +138,7 @@ public static string SingletonFactory(string typeName, string name, string lazyN return $@" internal {typeName} {name} - {{ + {{{ownerForward} var cached = {lazyName}; if (cached != null) return cached; diff --git a/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs b/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs index f08435d..16c3943 100644 --- a/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs +++ b/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs @@ -40,7 +40,7 @@ public void BuildChainCreatesWorkingContainerPipeline() var baseContainer = new DependencyInjectionContainer(default, default, new NonInjectedClass()); // Build a chain using the registry - var final = ContainerRegistry.BuildChain(baseContainer, new[] { "Inheritor" }); + var final = ContainerRegistry.BuildChain(baseContainer, ["Inheritor"]); // The final container should be able to resolve types from the base final.Resolve().ShouldNotBeNull(); @@ -65,7 +65,7 @@ public void BuildChainWithExplicitAssemblyListSkipsCurrentContainerAssembly() EnsureContainerEntryPointModuleInitialized(); var baseContainer = new DependencyInjectionContainer(default, default, new NonInjectedClass()); - var final = ContainerRegistry.BuildChain(baseContainer, new[] { "Inheritor" }); + var final = ContainerRegistry.BuildChain(baseContainer, ["Inheritor"]); ReferenceEquals(final, baseContainer).ShouldBeTrue(); final.Inheritor.ShouldBeNull(); diff --git a/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs b/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs index a641c79..ed521d0 100644 --- a/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs +++ b/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs @@ -15,35 +15,35 @@ public class GeneratorBehaviorTests 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) - { - } -} -} -"""; + 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); @@ -56,28 +56,65 @@ public SecondConsumer(ExternalValue second) } [Test] - public void GeneratorSupportsBooleanKeysThatAreNotIdentifiers() + public void GeneratorPicksUpClassesViaAttributeAppliedToImplementedInterface() { + // [Inject] is declared on the interface, not the implementing class. FindMethods must + // still discover ImplicitlyInjected by checking type.AllInterfaces for attributes, not just + // the type's own attributes — pins down this (previously untested) discovery path so it + // can't silently regress in a future refactor of the type-scanning pipeline. const string source = """ -using FactoryGenerator.Attributes; + using FactoryGenerator.Attributes; -namespace Sample -{ -public interface IService -{ -} + namespace Sample + { + [Inject] + public interface IMarker + { + } -[Inject, Boolean("feature-flag")] -public class EnabledService : IService -{ -} + public class ImplicitlyInjected : IMarker + { + } + } + """; -[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("ImplicitlyInjected"); + } + + [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); @@ -85,9 +122,9 @@ public class FallbackService : IService generatorResult.Exception.ShouldBeNull(); outputCompilation.GetDiagnostics() - .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) - .ToArray() - .ShouldBeEmpty(); + .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\""); @@ -100,20 +137,20 @@ public void BooleanOnlyImplementationsThrowInsteadOfResolvingNull() { var assemblyName = "BooleanOnly" + Guid.NewGuid().ToString("N"); var source = $$""" -using FactoryGenerator.Attributes; + using FactoryGenerator.Attributes; -namespace {{assemblyName}} -{ -public interface IService -{ -} + namespace {{assemblyName}} + { + public interface IService + { + } -[Inject, Boolean("enabled")] -public class EnabledService : IService -{ -} -} -"""; + [Inject, Boolean("enabled")] + public class EnabledService : IService + { + } + } + """; var compilation = CreateCompilation(assemblyName, source); var (runResult, outputCompilation) = RunGenerator(compilation); @@ -121,22 +158,22 @@ public class EnabledService : IService generatorResult.Exception.ShouldBeNull(); outputCompilation.GetDiagnostics() - .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) - .ToArray() - .ShouldBeEmpty(); + .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(); + .Single(sourceResult => sourceResult.HintName == $"FactoryGenerator.Declarations/{assemblyName}_IService().g.cs") + .SourceText + .ToString(); declarations.ShouldContain(expectedMessage); declarations.ShouldNotContain("null!"); var staticExtensions = generatorResult.GeneratedSources - .Single(sourceResult => sourceResult.HintName == "DependencyInjectionContainer.StaticExtensions.g.cs") - .SourceText - .ToString(); + .Single(sourceResult => sourceResult.HintName == $"FactoryGenerator.StaticExtensions/{assemblyName}_IServiceExtensions.g.cs") + .SourceText + .ToString(); staticExtensions.ShouldContain(expectedMessage); var assembly = System.Reflection.Assembly.Load(EmitAssembly(outputCompilation)); @@ -146,7 +183,7 @@ public class EnabledService : IService containerType.ShouldNotBeNull(); serviceType.ShouldNotBeNull(); - var container = (IContainer)Activator.CreateInstance(containerType!, new object[] { false })!; + var container = (IContainer) Activator.CreateInstance(containerType!, [false])!; var exception = Should.Throw(() => container.Resolve(serviceType!)); exception.Message.ShouldContain(expectedMessage); } @@ -155,35 +192,35 @@ public class EnabledService : IService 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(); -} -} -"""; + 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); @@ -199,35 +236,35 @@ public Factory(IResult result) 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(); -} -} -"""; + 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); @@ -244,66 +281,66 @@ public void InjectedMethodsSurfaceExternalParametersAndHonorOptionalAndParamsArg { 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); - } -} -} -"""; + 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); @@ -311,9 +348,9 @@ public IResult Create(ExternalInput input, string label = "default", params IPar generatorResult.Exception.ShouldBeNull(); outputCompilation.GetDiagnostics() - .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) - .ToArray() - .ShouldBeEmpty(); + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .ToArray() + .ShouldBeEmpty(); var assembly = System.Reflection.Assembly.Load(EmitAssembly(outputCompilation)); var containerType = assembly.GetType($"{assemblyName}.Generated.DependencyInjectionContainer"); @@ -325,14 +362,14 @@ public IResult Create(ExternalInput input, string label = "default", params IPar serviceType.ShouldNotBeNull(); var constructor = containerType!.GetConstructors() - .Single(ctor => - { - var parameters = ctor.GetParameters(); - return parameters.Length == 1 && parameters[0].ParameterType == externalType; - }); + .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 container = (IContainer) constructor.Invoke([external!]); var resolved = container.Resolve(serviceType!); resolved.GetType().GetProperty("Summary")!.GetValue(resolved).ShouldBe("runtime:default"); @@ -344,38 +381,38 @@ 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; } -} -} -"""; + 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); @@ -383,9 +420,9 @@ public Consumer(string label = "default", params IPart[] parts) generatorResult.Exception.ShouldBeNull(); outputCompilation.GetDiagnostics() - .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) - .ToArray() - .ShouldBeEmpty(); + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .ToArray() + .ShouldBeEmpty(); var assembly = System.Reflection.Assembly.Load(EmitAssembly(outputCompilation)); var containerType = assembly.GetType($"{assemblyName}.Generated.DependencyInjectionContainer"); @@ -394,7 +431,7 @@ public Consumer(string label = "default", params IPart[] parts) containerType.ShouldNotBeNull(); consumerType.ShouldNotBeNull(); - var container = (IContainer)Activator.CreateInstance(containerType!)!; + var container = (IContainer) Activator.CreateInstance(containerType!)!; var resolved = container.Resolve(consumerType!); resolved.GetType().GetProperty("Summary")!.GetValue(resolved).ShouldBe("default"); @@ -408,38 +445,38 @@ public void AssemblyPriorityCanOverrideProjectGraphPrecedence() var derivedAssemblyName = "PriorityDerived" + Guid.NewGuid().ToString("N"); var baseSource = $$""" -using FactoryGenerator.Attributes; + using FactoryGenerator.Attributes; -[assembly: InjectionPriority(9)] + [assembly: InjectionPriority(9)] -namespace {{baseAssemblyName}} -{ -public interface IService -{ -} + namespace {{baseAssemblyName}} + { + public interface IService + { + } -[Inject] -public class BaseService : IService -{ -} -} -"""; + [Inject] + public class BaseService : IService + { + } + } + """; var derivedSource = $$""" -using FactoryGenerator.Attributes; -using {{baseAssemblyName}}; - -namespace {{derivedAssemblyName}} -{ -[Inject] -public class DerivedService : IService -{ -} -} -"""; + using FactoryGenerator.Attributes; + using {{baseAssemblyName}}; + + namespace {{derivedAssemblyName}} + { + [Inject] + public class DerivedService : IService + { + } + } + """; var baseCompilation = CreateCompilation(baseAssemblyName, baseSource); -var (baseReference, _) = EmitReference(baseCompilation); + var (baseReference, _) = EmitReference(baseCompilation); var derivedCompilation = CreateCompilation(derivedAssemblyName, derivedSource, baseReference); var (runResult, outputCompilation) = RunGenerator(derivedCompilation); @@ -447,9 +484,9 @@ public class DerivedService : IService generatorResult.Exception.ShouldBeNull(); outputCompilation.GetDiagnostics() - .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) - .ToArray() - .ShouldBeEmpty(); + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .ToArray() + .ShouldBeEmpty(); var serviceMemberName = baseAssemblyName + "_IService()"; var prioritizedImplementationMemberName = baseAssemblyName + "_BaseService()"; @@ -474,18 +511,18 @@ private static CSharpCompilation CreateCompilation(string assemblyName, string s "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(); + 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)); references.AddRange(additionalReferences); return CSharpCompilation.Create( assemblyName: assemblyName, - syntaxTrees: new[] { syntaxTree }, + syntaxTrees: [syntaxTree], references: references, options: new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary)); } @@ -497,7 +534,7 @@ private static CSharpCompilation CreateCompilation(string source) private static (GeneratorDriverRunResult RunResult, Compilation OutputCompilation) RunGenerator(CSharpCompilation compilation) { - var parseOptions = (CSharpParseOptions)compilation.SyntaxTrees.First().Options; + var parseOptions = (CSharpParseOptions) compilation.SyntaxTrees.First().Options; GeneratorDriver driver = CSharpGeneratorDriver.Create( [new global::FactoryGenerator.FactoryGenerator().AsSourceGenerator()], parseOptions: parseOptions); @@ -518,4 +555,4 @@ private static byte[] EmitAssembly(Compilation compilation) result.Success.ShouldBeTrue(string.Join(Environment.NewLine, result.Diagnostics)); return stream.ToArray(); } -} +} \ No newline at end of file diff --git a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs index 4332eb4..88288b0 100644 --- a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs +++ b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs @@ -2,7 +2,6 @@ using Inheritor; using Inheritor.Generated; using Shouldly; -using System.Threading.Tasks; using Type = Inherited.Type; namespace FactoryGenerator.Tests; @@ -191,6 +190,14 @@ public void InheritorsOverride() m_container.Resolve().ShouldBeOfType(); } + [Test] + public void OverridenCanStillBeResolvedAsIEnumerable() + { + var result = m_container.Resolve>().ToArray(); + result.Length.ShouldBe(2); + result.Any(x => x is Overriden).ShouldBeTrue(); + } + [Test] public void OverrideImplementationsPreventFalsePositiveCycleDetection() { @@ -361,6 +368,7 @@ public void ClassesInsideOtherClassesCanBeInjected() { m_container.Resolve(); } + [Test] public void ContainerMayCreateItself() { @@ -369,12 +377,14 @@ public void ContainerMayCreateItself() resolved.Count().ShouldBe(6); var nonInjected = m_container.Resolve(); } + [Test] public void HierarchicalContainersResolveArraysProperly() { var newContainer = new DependencyInjectionContainer(m_container); newContainer.Resolve().Arrays.Count().ShouldBe(6); } + [Test] public void HierarchicalContainersResolveUsesFallBackIfItCannotFindImplementation() { @@ -415,6 +425,7 @@ public void ContainerPropgatesRelevantBooleansCreateItself() newContainer.GetBoolean("A").ShouldBeFalse(); newContainer.GetBoolean("TestBool").ShouldBeTrue(); } + [Test] public void HierarchicalContainersPropgatesBooleansUnknownToIt() { @@ -538,9 +549,51 @@ public void BaseContainerSeesInheritorArraysAfterLinking() parent.Resolve().Items.Count().ShouldBe(10); } + // ── Large dependency tree (1-3-9-27) tests ───────────────────────────────── + + [Test] + public void LargeTreeRootResolves() + { + var root = m_container.Resolve(); + root.ShouldNotBeNull(); + } + + [Test] + public void LargeTreeBranchesAreDistinct() + { + var root = m_container.Resolve(); + new object[] {root.Branch1, root.Branch2, root.Branch3} + .ShouldAllBe(b => b != null); + } + + [Test] + public void LargeTreeLeavesResolveThrough() + { + var root = m_container.Resolve(); + // Spot-check: walk root → Branch1 → Mid1 → Leaf01 + root.Branch1.Mid1.Leaf01.ShouldBeOfType(); + // Walk root → Branch3 → Mid3 → Leaf27 + root.Branch3.Mid9.Leaf27.ShouldBeOfType(); + } + + [Test] + public void LargeTreeAllMidNodesResolve() + { + var root = m_container.Resolve(); + var mids = new object[] + { + root.Branch1.Mid1, root.Branch1.Mid2, root.Branch1.Mid3, + root.Branch2.Mid4, root.Branch2.Mid5, root.Branch2.Mid6, + root.Branch3.Mid7, root.Branch3.Mid8, root.Branch3.Mid9, + }; + mids.ShouldAllBe(m => m != null); + mids.Select(m => m.GetType()).Distinct().Count().ShouldBe(9); + } + private class DummyContainer : IContainer { public const string DummyText = "I am a bit of text"; + private static readonly IFallbackCollectionItem[] s_fallbackCollectionItems = [ new DummyFallbackCollectionItem(), @@ -605,11 +658,11 @@ public bool TryResolve(out T? resolved) if (typeof(T) == typeof(IEnumerable)) resolved = (T) (object) s_fallbackCollectionItems; return resolved != null; } + public IEnumerable<(string Key, bool Value)> GetBooleans() { return [("B", true), ("C", false)]; } - } private sealed class DummyFallbackCollectionItem : IFallbackCollectionItem; diff --git a/Tests/TestData/Inheritor/Types.cs b/Tests/TestData/Inheritor/Types.cs index 1163552..7527efe 100644 --- a/Tests/TestData/Inheritor/Types.cs +++ b/Tests/TestData/Inheritor/Types.cs @@ -73,4 +73,55 @@ public class SplitInheritor1 : ISplitArray; public class SplitInheritor2 : ISplitArray; [Inject] -public class SplitInheritor3 : ISplitArray; \ No newline at end of file +public class SplitInheritor3 : ISplitArray; + +// ── Large dependency tree (1-3-9-27) ───────────────────────────────────────── +// Root → 3 branches → 9 branches → 27 leaves. Exercises deep, wide resolution. + +// Leaves (27) +[Inject] public class Leaf01; +[Inject] public class Leaf02; +[Inject] public class Leaf03; +[Inject] public class Leaf04; +[Inject] public class Leaf05; +[Inject] public class Leaf06; +[Inject] public class Leaf07; +[Inject] public class Leaf08; +[Inject] public class Leaf09; +[Inject] public class Leaf10; +[Inject] public class Leaf11; +[Inject] public class Leaf12; +[Inject] public class Leaf13; +[Inject] public class Leaf14; +[Inject] public class Leaf15; +[Inject] public class Leaf16; +[Inject] public class Leaf17; +[Inject] public class Leaf18; +[Inject] public class Leaf19; +[Inject] public class Leaf20; +[Inject] public class Leaf21; +[Inject, Singleton] public class Leaf22; +[Inject] public class Leaf23; +[Inject, Singleton] public class Leaf24; +[Inject] public class Leaf25; +[Inject] public class Leaf26; +[Inject] public class Leaf27; + +// Mid-level (9) — each depends on 3 leaves +[Inject] public class Mid1(Leaf01 leaf01, Leaf02 leaf02, Leaf03 leaf03) { public Leaf01 Leaf01 => leaf01; public Leaf02 Leaf02 => leaf02; public Leaf03 Leaf03 => leaf03; } +[Inject] public class Mid2(Leaf04 leaf04, Leaf05 leaf05, Leaf06 leaf06) { public Leaf04 Leaf04 => leaf04; public Leaf05 Leaf05 => leaf05; public Leaf06 Leaf06 => leaf06; } +[Inject] public class Mid3(Leaf07 leaf07, Leaf08 leaf08, Leaf09 leaf09) { public Leaf07 Leaf07 => leaf07; public Leaf08 Leaf08 => leaf08; public Leaf09 Leaf09 => leaf09; } +[Inject] public class Mid4(Leaf10 leaf10, Leaf11 leaf11, Leaf12 leaf12) { public Leaf10 Leaf10 => leaf10; public Leaf11 Leaf11 => leaf11; public Leaf12 Leaf12 => leaf12; } +[Inject] public class Mid5(Leaf13 leaf13, Leaf14 leaf14, Leaf15 leaf15) { public Leaf13 Leaf13 => leaf13; public Leaf14 Leaf14 => leaf14; public Leaf15 Leaf15 => leaf15; } +[Inject] public class Mid6(Leaf16 leaf16, Leaf17 leaf17, Leaf18 leaf18) { public Leaf16 Leaf16 => leaf16; public Leaf17 Leaf17 => leaf17; public Leaf18 Leaf18 => leaf18; } +[Inject] public class Mid7(Leaf19 leaf19, Leaf20 leaf20, Leaf21 leaf21) { public Leaf19 Leaf19 => leaf19; public Leaf20 Leaf20 => leaf20; public Leaf21 Leaf21 => leaf21; } +[Inject] public class Mid8(Leaf22 leaf22, Leaf23 leaf23, Leaf24 leaf24) { public Leaf22 Leaf22 => leaf22; public Leaf23 Leaf23 => leaf23; public Leaf24 Leaf24 => leaf24; } +[Inject] public class Mid9(Leaf25 leaf25, Leaf26 leaf26, Leaf27 leaf27) { public Leaf25 Leaf25 => leaf25; public Leaf26 Leaf26 => leaf26; public Leaf27 Leaf27 => leaf27; } + +// Branches (3) — each depends on 3 mid-level nodes +[Inject] public class Branch1(Mid1 mid1, Mid2 mid2, Mid3 mid3) { public Mid1 Mid1 => mid1; public Mid2 Mid2 => mid2; public Mid3 Mid3 => mid3; } +[Inject] public class Branch2(Mid4 mid4, Mid5 mid5, Mid6 mid6) { public Mid4 Mid4 => mid4; public Mid5 Mid5 => mid5; public Mid6 Mid6 => mid6; } +[Inject] public class Branch3(Mid7 mid7, Mid8 mid8, Mid9 mid9) { public Mid7 Mid7 => mid7; public Mid8 Mid8 => mid8; public Mid9 Mid9 => mid9; } + +// Root (1) — depends on 3 branches +[Inject, Self] public class TreeRoot(Branch1 branch1, Branch2 branch2, Branch3 branch3) { public Branch1 Branch1 => branch1; public Branch2 Branch2 => branch2; public Branch3 Branch3 => branch3; } \ No newline at end of file