diff --git a/.github/workflows/benchmark.yml b/.github/workflows/benchmark.yml index a48ddd7..74f3b7c 100644 --- a/.github/workflows/benchmark.yml +++ b/.github/workflows/benchmark.yml @@ -23,11 +23,11 @@ jobs: - name: Setup dotnet ${{ matrix.dotnet-version }} uses: actions/setup-dotnet@v4 with: - dotnet-version: "9.0.x" + dotnet-version: "10.0.x" - name: Run benchmark run: cd Benchmarking/Benchmarks && dotnet run -c Release --exporters json --filter '*' - - name: Store benchmark result + - name: Store runtime benchmark result uses: rhysd/github-action-benchmark@v1 with: name: Benchmark.Net Benchmark @@ -39,3 +39,15 @@ jobs: # Show alert with commit comment on detecting possible performance regression alert-threshold: '200%' comment-on-alert: true + + - name: Store generator benchmark result + uses: rhysd/github-action-benchmark@v1 + with: + name: Generator Benchmark.Net Benchmark + tool: 'benchmarkdotnet' + output-file-path: Benchmarking/Benchmarks/BenchmarkDotNet.Artifacts/results/Benchmarks.GeneratorBenchmarks-report-full-compressed.json + github-token: ${{ secrets.GITHUB_TOKEN }} + auto-push: true + summary-always: true + alert-threshold: '200%' + comment-on-alert: true diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index 02c182e..cb76aff 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -28,13 +28,13 @@ jobs: - name: Setup dotnet ${{ matrix.dotnet-version }} uses: actions/setup-dotnet@v4 with: - dotnet-version: "9.0.x" + dotnet-version: "10.0.x" - name: Restore run: dotnet restore FactoryGenerator.sln - name: Build run: dotnet build FactoryGenerator.sln --no-restore - name: Test - run: dotnet test FactoryGenerator.sln --no-build --no-restore + run: dotnet test --no-build --no-restore --solution FactoryGenerator.sln pack: runs-on: ubuntu-latest @@ -45,7 +45,7 @@ jobs: - name: Setup dotnet ${{ matrix.dotnet-version }} uses: actions/setup-dotnet@v4 with: - dotnet-version: "9.0.x" + dotnet-version: "10.0.x" - name: Pack Generator run: dotnet pack FactoryGenerator/FactoryGenerator.csproj -o "${{ env.NuGetDirectory }}" --property:RepositoryCommit="${{ env.COMMIT_SHA }}" --property:InformationalVersion="UNRELEASED" --property:AssemblyVersion="0.0.0" --property:FileVersion="0.0.0" --property:Version="0.0.0" - name: Pack Attributes @@ -68,7 +68,7 @@ jobs: - name: Setup dotnet ${{ matrix.dotnet-version }} uses: actions/setup-dotnet@v4 with: - dotnet-version: "9.0.x" + dotnet-version: "10.0.x" - name: Pack Generator run: dotnet pack FactoryGenerator/FactoryGenerator.csproj -o "${{ env.NuGetDirectory }}" --property:RepositoryCommit="${{ env.COMMIT_SHA }}" --property:InformationalVersion="${{ github.ref_name }}" --property:AssemblyVersion="${{ github.ref_name }}" --property:FileVersion="${{ github.ref_name }}" --property:Version="${{ github.ref_name }}" - name: Pack Attributes @@ -94,7 +94,7 @@ jobs: - name: Setup dotnet ${{ matrix.dotnet-version }} uses: actions/setup-dotnet@v4 with: - dotnet-version: "9.0.x" + dotnet-version: "10.0.x" - name: Publish Nuget packages run: | for file in $(find "${{ env.NuGetDirectory }}" -type f -name "*.nupkg"); do @@ -110,11 +110,11 @@ jobs: - name: Setup dotnet ${{ matrix.dotnet-version }} uses: actions/setup-dotnet@v4 with: - dotnet-version: "9.0.x" + dotnet-version: "10.0.x" - name: Run benchmark run: cd Benchmarking/Benchmarks && dotnet run -c Release --exporters json --filter '*' - - name: Store benchmark result + - name: Store runtime benchmark result uses: rhysd/github-action-benchmark@v1 with: name: Benchmark.Net Benchmark @@ -124,5 +124,17 @@ jobs: summary-always: true # Show alert with commit comment on detecting possible performance regression alert-threshold: '200%' - comment-on-alert: true - fail-on-alert: true + 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/.gitignore b/.gitignore index 83371ec..45b9c53 100644 --- a/.gitignore +++ b/.gitignore @@ -21,7 +21,6 @@ [Rr]elease-x86/ [Dd]ebug-x86/ x64/ -build/ bld/ [Bb]in/ [Oo]bj/ diff --git a/Benchmarking/Benchmarks/Benchmarks.csproj b/Benchmarking/Benchmarks/Benchmarks.csproj index 7da39a4..78c4269 100644 --- a/Benchmarking/Benchmarks/Benchmarks.csproj +++ b/Benchmarking/Benchmarks/Benchmarks.csproj @@ -2,17 +2,20 @@ Exe - net9.0 + net10.0 + preview enable enable - true + + + diff --git a/Benchmarking/Benchmarks/GeneratorBenchmarks.cs b/Benchmarking/Benchmarks/GeneratorBenchmarks.cs new file mode 100644 index 0000000..14a136a --- /dev/null +++ b/Benchmarking/Benchmarks/GeneratorBenchmarks.cs @@ -0,0 +1,751 @@ +using System.Collections.Immutable; +using System.IO; +using System.Text; +using BenchmarkDotNet.Attributes; +using FactoryGenerator.Attributes; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using Microsoft.CodeAnalysis.Diagnostics; + +namespace Benchmarks; + +[MemoryDiagnoser] +[ShortRunJob] +[JsonExporterAttribute.Full] +[JsonExporterAttribute.FullCompressed] +public class GeneratorBenchmarks +{ + private ColdGeneratorScenario m_constructorGraph = null!; + private ColdGeneratorScenario m_noiseHeavyProject = null!; + private ColdGeneratorScenario m_featureRichStaticExtensionsDisabled = null!; + private ColdGeneratorScenario m_featureRichStaticExtensionsEnabled = null!; + private ColdGeneratorScenario m_multiAssemblyOverrideGraph = null!; + private FeatureRichIncrementalScenario m_featureRichIncremental = null!; + private IncrementalGeneratorScenario m_referenceAssemblyIncremental = null!; + + [GlobalSetup] + public void Setup() + { + m_constructorGraph = GeneratorBenchmarkScenarioFactory.CreateConstructorGraph(serviceCount: 250); + m_noiseHeavyProject = GeneratorBenchmarkScenarioFactory.CreateNoiseHeavyProject(serviceCount: 64, noiseTypeCount: 2000); + m_featureRichStaticExtensionsDisabled = GeneratorBenchmarkScenarioFactory.CreateFeatureRichGraph(emitStaticExtensions: false); + m_featureRichStaticExtensionsEnabled = GeneratorBenchmarkScenarioFactory.CreateFeatureRichGraph(emitStaticExtensions: true); + m_multiAssemblyOverrideGraph = GeneratorBenchmarkScenarioFactory.CreateMultiAssemblyOverrideGraph(baseServiceCount: 128, overrideCount: 16); + m_featureRichIncremental = GeneratorBenchmarkScenarioFactory.CreateFeatureRichIncrementalScenario(); + m_referenceAssemblyIncremental = GeneratorBenchmarkScenarioFactory.CreateReferenceAssemblyIncrementalScenario(); + + GeneratorBenchmarkHarness.Validate(m_constructorGraph); + GeneratorBenchmarkHarness.Validate(m_noiseHeavyProject); + GeneratorBenchmarkHarness.Validate(m_featureRichStaticExtensionsDisabled); + GeneratorBenchmarkHarness.Validate(m_featureRichStaticExtensionsEnabled); + GeneratorBenchmarkHarness.Validate(m_multiAssemblyOverrideGraph); + } + + [Benchmark] + public int Cold_ConstructorGraph() => GeneratorBenchmarkHarness.RunCold(m_constructorGraph); + + [Benchmark] + public int Cold_NoiseHeavyProject() => GeneratorBenchmarkHarness.RunCold(m_noiseHeavyProject); + + [Benchmark] + public int Cold_FeatureRichGraph_StaticExtensionsDisabled() => GeneratorBenchmarkHarness.RunCold(m_featureRichStaticExtensionsDisabled); + + [Benchmark] + public int Cold_FeatureRichGraph_StaticExtensionsEnabled() => GeneratorBenchmarkHarness.RunCold(m_featureRichStaticExtensionsEnabled); + + [Benchmark] + public int Cold_MultiAssemblyOverrideGraph() => GeneratorBenchmarkHarness.RunCold(m_multiAssemblyOverrideGraph); +} + +internal sealed class ColdGeneratorScenario(CSharpCompilation compilation, AnalyzerConfigOptionsProvider optionsProvider) +{ + public CSharpCompilation Compilation { get; } = compilation; + public AnalyzerConfigOptionsProvider OptionsProvider { get; } = optionsProvider; +} + +internal sealed class FeatureRichIncrementalScenario( + GeneratorDriver warmDriver, + CSharpCompilation baselineCompilation, + CSharpCompilation unrelatedEditCompilation, + CSharpCompilation injectedSignatureEditCompilation, + CSharpCompilation addInjectCompilation) +{ + public GeneratorDriver WarmDriver { get; } = warmDriver; + public CSharpCompilation BaselineCompilation { get; } = baselineCompilation; + public CSharpCompilation UnrelatedEditCompilation { get; } = unrelatedEditCompilation; + public CSharpCompilation InjectedSignatureEditCompilation { get; } = injectedSignatureEditCompilation; + public CSharpCompilation AddInjectCompilation { get; } = addInjectCompilation; +} + +internal sealed class IncrementalGeneratorScenario(GeneratorDriver warmDriver, CSharpCompilation changedCompilation) +{ + public GeneratorDriver WarmDriver { get; } = warmDriver; + public CSharpCompilation ChangedCompilation { get; } = changedCompilation; +} + +internal static class GeneratorBenchmarkHarness +{ + public static int RunCold(ColdGeneratorScenario scenario) + { + var driver = CreateDriver(scenario.Compilation, scenario.OptionsProvider); + return RunAndSummarize(driver, scenario.Compilation); + } + + public static int RunIncremental(GeneratorDriver warmDriver, CSharpCompilation compilation) + { + return RunAndSummarize(warmDriver, compilation); + } + + public static GeneratorDriver WarmAndValidate(ColdGeneratorScenario scenario) + { + var driver = CreateDriver(scenario.Compilation, scenario.OptionsProvider); + return RunAndValidate(driver, scenario.Compilation); + } + + public static void Validate(ColdGeneratorScenario scenario) + { + _ = WarmAndValidate(scenario); + } + + public static void Validate(GeneratorDriver warmDriver, CSharpCompilation compilation) + { + _ = RunAndValidate(warmDriver, compilation); + } + + private static GeneratorDriver CreateDriver(CSharpCompilation compilation, AnalyzerConfigOptionsProvider optionsProvider) + { + return CSharpGeneratorDriver.Create( + [new global::FactoryGenerator.FactoryGenerator().AsSourceGenerator()], + parseOptions: (CSharpParseOptions) compilation.SyntaxTrees.First().Options, + optionsProvider: optionsProvider); + } + + private static GeneratorDriver RunAndValidate(GeneratorDriver driver, CSharpCompilation compilation) + { + driver = driver.RunGeneratorsAndUpdateCompilation(compilation, out var outputCompilation, out _); + var runResult = driver.GetRunResult(); + var exception = runResult.Results + .Select(result => result.Exception) + .FirstOrDefault(resultException => resultException is not null); + + if (exception is not null) + throw exception; + + var errors = outputCompilation.GetDiagnostics() + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .Select(diagnostic => diagnostic.ToString()) + .ToArray(); + + if (errors.Length != 0) + throw new InvalidOperationException(string.Join(Environment.NewLine, errors)); + + return driver; + } + + private static int RunAndSummarize(GeneratorDriver driver, CSharpCompilation compilation) + { + driver = driver.RunGeneratorsAndUpdateCompilation(compilation, out _, out _); + var runResult = driver.GetRunResult(); + + return runResult.Results.Sum(result => result.GeneratedSources.Sum(source => source.SourceText.Length)); + } +} + +internal static class GeneratorBenchmarkScenarioFactory +{ + private static readonly ImmutableArray s_metadataReferences = CreateMetadataReferences(); + private static readonly AnalyzerConfigOptionsProvider s_staticExtensionsEnabledOptions = new BenchmarkAnalyzerConfigOptionsProvider(true); + private static readonly AnalyzerConfigOptionsProvider s_staticExtensionsDisabledOptions = new BenchmarkAnalyzerConfigOptionsProvider(false); + + public static ColdGeneratorScenario CreateConstructorGraph(int serviceCount) + { + var compilation = CreateCompilation( + "GeneratorConstructorGraphBenchmarks", + new BenchmarkSourceDocument("ConstructorGraph.cs", BuildConstructorGraphSource("GeneratorConstructorGraphInput", serviceCount))); + + return new ColdGeneratorScenario(compilation, s_staticExtensionsDisabledOptions); + } + + public static ColdGeneratorScenario CreateNoiseHeavyProject(int serviceCount, int noiseTypeCount) + { + var compilation = CreateCompilation( + "GeneratorNoiseHeavyBenchmarks", + new BenchmarkSourceDocument("ConstructorGraph.cs", BuildConstructorGraphSource("GeneratorNoiseInput", serviceCount)), + new BenchmarkSourceDocument("Noise.cs", BuildNoiseSource("GeneratorNoiseInput", noiseTypeCount))); + + return new ColdGeneratorScenario(compilation, s_staticExtensionsDisabledOptions); + } + + public static ColdGeneratorScenario CreateFeatureRichGraph(bool emitStaticExtensions) + { + var compilation = CreateCompilation( + emitStaticExtensions ? "GeneratorFeatureRichStaticExtensionsBenchmarks" : "GeneratorFeatureRichBenchmarks", + new BenchmarkSourceDocument( + "FeatureGraph.cs", + BuildFeatureRichSource("GeneratorFeatureRichInput", includeAdditionalExternalParameter: false, includeExtraWidgetInjection: false, labelDefault: "default", retryCountDefault: 3)), + new BenchmarkSourceDocument("Utilities.cs", BuildUtilitySource("GeneratorFeatureRichInput", utilitySuffix: "Baseline"))); + + return new ColdGeneratorScenario(compilation, emitStaticExtensions ? s_staticExtensionsEnabledOptions : s_staticExtensionsDisabledOptions); + } + + public static ColdGeneratorScenario CreateMultiAssemblyOverrideGraph(int baseServiceCount, int overrideCount) + { + const string baseAssemblyName = "GeneratorOverrideBase"; + const string derivedAssemblyName = "GeneratorOverrideDerived"; + + var baseCompilation = CreateCompilation( + baseAssemblyName, + new BenchmarkSourceDocument("BaseServices.cs", BuildOverrideBaseSource(baseAssemblyName, baseServiceCount))); + var baseReference = EmitReference(baseCompilation); + var derivedCompilation = CreateCompilation( + derivedAssemblyName, + baseReference, + new BenchmarkSourceDocument("DerivedServices.cs", BuildOverrideDerivedSource(baseAssemblyName, derivedAssemblyName, baseServiceCount, overrideCount))); + + return new ColdGeneratorScenario(derivedCompilation, s_staticExtensionsDisabledOptions); + } + + public static FeatureRichIncrementalScenario CreateFeatureRichIncrementalScenario() + { + const string assemblyName = "GeneratorFeatureRichIncremental"; + + var baselineCompilation = CreateCompilation( + assemblyName, + new BenchmarkSourceDocument( + "FeatureGraph.cs", BuildFeatureRichSource(assemblyName, includeAdditionalExternalParameter: false, includeExtraWidgetInjection: false, labelDefault: "default", retryCountDefault: 3)), + new BenchmarkSourceDocument("Utilities.cs", BuildUtilitySource(assemblyName, utilitySuffix: "Baseline"))); + var unrelatedEditCompilation = CreateCompilation( + assemblyName, + new BenchmarkSourceDocument( + "FeatureGraph.cs", BuildFeatureRichSource(assemblyName, includeAdditionalExternalParameter: false, includeExtraWidgetInjection: false, labelDefault: "default", retryCountDefault: 3)), + new BenchmarkSourceDocument("Utilities.cs", BuildUtilitySource(assemblyName, utilitySuffix: "Edited"))); + var injectedSignatureEditCompilation = CreateCompilation( + assemblyName, + new BenchmarkSourceDocument( + "FeatureGraph.cs", BuildFeatureRichSource(assemblyName, includeAdditionalExternalParameter: true, includeExtraWidgetInjection: false, labelDefault: "edited", retryCountDefault: 5)), + new BenchmarkSourceDocument("Utilities.cs", BuildUtilitySource(assemblyName, utilitySuffix: "Baseline"))); + var addInjectCompilation = CreateCompilation( + assemblyName, + new BenchmarkSourceDocument( + "FeatureGraph.cs", BuildFeatureRichSource(assemblyName, includeAdditionalExternalParameter: false, includeExtraWidgetInjection: true, labelDefault: "default", retryCountDefault: 3)), + new BenchmarkSourceDocument("Utilities.cs", BuildUtilitySource(assemblyName, utilitySuffix: "Baseline"))); + + var baselineScenario = new ColdGeneratorScenario(baselineCompilation, s_staticExtensionsEnabledOptions); + var warmDriver = GeneratorBenchmarkHarness.WarmAndValidate(baselineScenario); + GeneratorBenchmarkHarness.Validate(GeneratorBenchmarkHarness.WarmAndValidate(baselineScenario), unrelatedEditCompilation); + GeneratorBenchmarkHarness.Validate(GeneratorBenchmarkHarness.WarmAndValidate(baselineScenario), injectedSignatureEditCompilation); + GeneratorBenchmarkHarness.Validate(GeneratorBenchmarkHarness.WarmAndValidate(baselineScenario), addInjectCompilation); + + return new FeatureRichIncrementalScenario( + warmDriver, + baselineCompilation, + unrelatedEditCompilation, + injectedSignatureEditCompilation, + addInjectCompilation); + } + + public static IncrementalGeneratorScenario CreateReferenceAssemblyIncrementalScenario() + { + const string baseAssemblyName = "GeneratorReferenceBase"; + const string derivedAssemblyName = "GeneratorReferenceDerived"; + + var baseCompilation = CreateCompilation( + baseAssemblyName, + new BenchmarkSourceDocument("BaseServices.cs", BuildReferenceBaseSource(baseAssemblyName, includeSecondBasePart: false))); + var changedBaseCompilation = CreateCompilation( + baseAssemblyName, + new BenchmarkSourceDocument("BaseServices.cs", BuildReferenceBaseSource(baseAssemblyName, includeSecondBasePart: true))); + + var baselineCompilation = CreateCompilation( + derivedAssemblyName, + EmitReference(baseCompilation), + new BenchmarkSourceDocument("DerivedServices.cs", BuildReferenceDerivedSource(baseAssemblyName, derivedAssemblyName))); + var changedCompilation = CreateCompilation( + derivedAssemblyName, + EmitReference(changedBaseCompilation), + new BenchmarkSourceDocument("DerivedServices.cs", BuildReferenceDerivedSource(baseAssemblyName, derivedAssemblyName))); + + var baselineScenario = new ColdGeneratorScenario(baselineCompilation, s_staticExtensionsEnabledOptions); + var warmDriver = GeneratorBenchmarkHarness.WarmAndValidate(baselineScenario); + GeneratorBenchmarkHarness.Validate(GeneratorBenchmarkHarness.WarmAndValidate(baselineScenario), changedCompilation); + + return new IncrementalGeneratorScenario(warmDriver, changedCompilation); + } + + private static CSharpCompilation CreateCompilation(string assemblyName, params BenchmarkSourceDocument[] documents) + { + return CreateCompilation(assemblyName, s_metadataReferences, documents); + } + + private static CSharpCompilation CreateCompilation( + string assemblyName, + MetadataReference additionalReference, + params BenchmarkSourceDocument[] documents) + { + return CreateCompilation(assemblyName, s_metadataReferences.Add(additionalReference), documents); + } + + private static CSharpCompilation CreateCompilation( + string assemblyName, + ImmutableArray references, + params BenchmarkSourceDocument[] documents) + { + var syntaxTrees = documents + .Select(document => CSharpSyntaxTree.ParseText( + document.Source, + new CSharpParseOptions(LanguageVersion.Preview), + path: document.FileName)) + .ToArray(); + + return CSharpCompilation.Create( + assemblyName, + syntaxTrees, + references, + new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary)); + } + + private static MetadataReference EmitReference(Compilation compilation) + { + using var stream = new MemoryStream(); + var result = compilation.Emit(stream); + if (!result.Success) + { + throw new InvalidOperationException( + string.Join(Environment.NewLine, result.Diagnostics.Select(diagnostic => diagnostic.ToString()))); + } + + return MetadataReference.CreateFromImage(stream.ToArray()); + } + + private static ImmutableArray CreateMetadataReferences() + { + var excludedAssemblies = new HashSet(StringComparer.Ordinal) + { + "Benchmarks", + "FactoryGenerator", + "FactoryGenerator.Attributes", + "FactoryGenerator.Extensions.AspNetCore", + "FactoryGenerator.Extensions.AspNetCore.Tests", + "FactoryGenerator.Tests", + "Inherited", + "Inheritor", + "TestWebApp" + }; + + return + [ + .. ((string?) AppContext.GetData("TRUSTED_PLATFORM_ASSEMBLIES"))! + .Split(Path.PathSeparator) + .Where(path => !excludedAssemblies.Contains(Path.GetFileNameWithoutExtension(path))) + .Select(path => (MetadataReference) MetadataReference.CreateFromFile(path)), + + MetadataReference.CreateFromFile(typeof(InjectAttribute).Assembly.Location) + ]; + } + + private static string BuildConstructorGraphSource(string namespaceName, int serviceCount) + { + var sb = new StringBuilder(); + sb.AppendLine("using FactoryGenerator.Attributes;"); + sb.AppendLine(); + sb.AppendLine($"namespace {namespaceName}"); + sb.AppendLine("{"); + + for (var i = 0; i < serviceCount; i++) + { + sb.AppendLine($"public interface IService{i}"); + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine(); + sb.AppendLine("[Inject]"); + if (i == 0) + { + sb.AppendLine($"public sealed class Service{i} : IService{i}"); + sb.AppendLine("{"); + sb.AppendLine("}"); + } + else + { + sb.AppendLine($"public sealed class Service{i}(IService{i - 1} previous) : IService{i}"); + sb.AppendLine("{"); + sb.AppendLine("}"); + } + + sb.AppendLine(); + } + + sb.AppendLine("[Inject]"); + sb.AppendLine($"public sealed class RootConsumer(IService{serviceCount - 1} root)"); + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine("}"); + + return sb.ToString(); + } + + private static string BuildNoiseSource(string namespaceName, int noiseTypeCount) + { + var sb = new StringBuilder(); + sb.AppendLine("using System;"); + sb.AppendLine(); + sb.AppendLine($"namespace {namespaceName}"); + sb.AppendLine("{"); + + for (var i = 0; i < noiseTypeCount; i++) + { + sb.AppendLine($"public sealed class NoiseType{i}"); + sb.AppendLine("{"); + sb.AppendLine($" public int Compute(int value) => value + {i};"); + sb.AppendLine($" public string Name => \"NoiseType{i}\";"); + sb.AppendLine(" public DateTime Timestamp => DateTime.UnixEpoch;"); + sb.AppendLine("}"); + sb.AppendLine(); + } + + sb.AppendLine("}"); + return sb.ToString(); + } + + private static string BuildFeatureRichSource( + string namespaceName, + bool includeAdditionalExternalParameter, + bool includeExtraWidgetInjection, + string labelDefault, + int retryCountDefault) + { + var additionalExternalParameter = includeAdditionalExternalParameter ? ", AdditionalExternalDependency additional" : string.Empty; + var widgetCAttribute = includeExtraWidgetInjection ? "[Inject]\n" : string.Empty; + + return $$""" + using System.Collections.Generic; + using FactoryGenerator.Attributes; + + namespace {{namespaceName}} + { + public sealed class ExternalDependency + { + } + + public sealed class AdditionalExternalDependency + { + } + + public interface IFlaggedFeature + { + } + + [Inject, Boolean("feature_enabled")] + public sealed class EnabledFeature : IFlaggedFeature + { + } + + [Inject] + public sealed class FallbackFeature : IFlaggedFeature + { + } + + public interface IWidget + { + } + + [Inject] + public sealed class WidgetA : IWidget + { + } + + [Inject] + public sealed class WidgetB : IWidget + { + } + + {{widgetCAttribute}}public sealed class WidgetC : IWidget + { + } + + public interface IPropertyResult + { + } + + public sealed class PropertyResult(IFlaggedFeature feature) : IPropertyResult + { + public IFlaggedFeature Feature { get; } = feature; + } + + public interface IPropertyFactory + { + [Inject] + IPropertyResult Value { get; } + } + + [Inject] + public sealed class PropertyFactory(IFlaggedFeature feature) : IPropertyFactory + { + public IPropertyResult Value => new PropertyResult(feature); + } + + public interface IMethodResult + { + } + + public sealed class MethodResult( + IFlaggedFeature feature, + IEnumerable widgets, + ExternalDependency external, + string label, + int retryCount) : IMethodResult + { + public IFlaggedFeature Feature { get; } = feature; + public IEnumerable Widgets { get; } = widgets; + public ExternalDependency External { get; } = external; + public string Label { get; } = label; + public int RetryCount { get; } = retryCount; + } + + public interface IFeatureFactory + { + [Inject] + IMethodResult Create( + ExternalDependency external{{additionalExternalParameter}}, + string label = "{{labelDefault}}", + int retryCount = {{retryCountDefault}}, + params IWidget[] widgets); + } + + [Inject] + public sealed class FeatureFactory(IFlaggedFeature feature) : IFeatureFactory + { + public IMethodResult Create( + ExternalDependency external{{additionalExternalParameter}}, + string label = "{{labelDefault}}", + int retryCount = {{retryCountDefault}}, + params IWidget[] widgets) + { + return new MethodResult(feature, widgets, external, label, retryCount); + } + } + + [Inject] + public sealed class FeatureGraphConsumer( + IMethodResult methodResult, + IPropertyResult propertyResult, + IEnumerable widgets, + IFlaggedFeature feature) + { + public IMethodResult MethodResult { get; } = methodResult; + public IPropertyResult PropertyResult { get; } = propertyResult; + public IEnumerable Widgets { get; } = widgets; + public IFlaggedFeature Feature { get; } = feature; + } + } + """; + } + + private static string BuildUtilitySource(string namespaceName, string utilitySuffix) + { + return $$""" + namespace {{namespaceName}} + { + public static class UtilityValues + { + public const string Marker = "{{utilitySuffix}}"; + + public static string Combine(string prefix) + { + return prefix + Marker; + } + } + } + """; + } + + private static string BuildOverrideBaseSource(string assemblyName, int baseServiceCount) + { + var sb = new StringBuilder(); + sb.AppendLine("using FactoryGenerator.Attributes;"); + sb.AppendLine(); + sb.AppendLine("[assembly: InjectionPriority(9)]"); + sb.AppendLine(); + sb.AppendLine($"namespace {assemblyName}"); + sb.AppendLine("{"); + sb.AppendLine("public interface ISharedService"); + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine(); + sb.AppendLine("[Inject]"); + sb.AppendLine("public sealed class BaseSharedService : ISharedService"); + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine(); + + for (var i = 0; i < baseServiceCount; i++) + { + sb.AppendLine($"public interface INode{i}"); + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine(); + sb.AppendLine("[Inject]"); + if (i == 0) + { + sb.AppendLine($"public sealed class BaseNode{i}(ISharedService sharedService) : INode{i}"); + } + else + { + sb.AppendLine($"public sealed class BaseNode{i}(INode{i - 1} previous) : INode{i}"); + } + + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine(); + } + + sb.AppendLine("}"); + return sb.ToString(); + } + + private static string BuildOverrideDerivedSource(string baseAssemblyName, string derivedAssemblyName, int baseServiceCount, int overrideCount) + { + var sb = new StringBuilder(); + sb.AppendLine($"using {baseAssemblyName};"); + sb.AppendLine("using FactoryGenerator.Attributes;"); + sb.AppendLine(); + sb.AppendLine($"namespace {derivedAssemblyName}"); + sb.AppendLine("{"); + + for (var i = 0; i < overrideCount; i++) + { + var serviceIndex = i * Math.Max(1, baseServiceCount / overrideCount); + sb.AppendLine("[Inject]"); + if (serviceIndex == 0) + { + sb.AppendLine($"public sealed class DerivedNode{serviceIndex}(ISharedService sharedService) : INode{serviceIndex}"); + } + else + { + sb.AppendLine($"public sealed class DerivedNode{serviceIndex}(INode{serviceIndex - 1} previous) : INode{serviceIndex}"); + } + + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine(); + } + + sb.AppendLine("[Inject]"); + sb.AppendLine($"public sealed class DerivedRoot(ISharedService sharedService, INode0 firstNode, INode{baseServiceCount - 1} lastNode)"); + sb.AppendLine("{"); + sb.AppendLine("}"); + sb.AppendLine("}"); + return sb.ToString(); + } + + private static string BuildReferenceBaseSource(string assemblyName, bool includeSecondBasePart) + { + var secondPart = includeSecondBasePart + ? """ + + [Inject] + public sealed class BasePartTwo : IBasePart + { + } + """ + : string.Empty; + + return $$""" + using System.Collections.Generic; + using FactoryGenerator.Attributes; + + namespace {{assemblyName}} + { + public interface IBasePart + { + } + + [Inject] + public sealed class BasePartOne : IBasePart + { + } + {{secondPart}} + + public interface IBaseService + { + } + + [Inject] + public sealed class BaseService(IEnumerable parts) : IBaseService + { + public IEnumerable Parts { get; } = parts; + } + } + """; + } + + private static string BuildReferenceDerivedSource(string baseAssemblyName, string derivedAssemblyName) + { + return $$""" + using System.Collections.Generic; + using FactoryGenerator.Attributes; + using {{baseAssemblyName}}; + + namespace {{derivedAssemblyName}} + { + [Inject] + public sealed class DerivedPart : IBasePart + { + } + + [Inject] + public sealed class DerivedConsumer(IBaseService service, IEnumerable parts) + { + public IBaseService Service { get; } = service; + public IEnumerable Parts { get; } = parts; + } + } + """; + } +} + +internal sealed class BenchmarkSourceDocument(string fileName, string source) +{ + public string FileName { get; } = fileName; + public string Source { get; } = source; +} + +internal sealed class BenchmarkAnalyzerConfigOptionsProvider(bool emitStaticExtensions) : AnalyzerConfigOptionsProvider +{ + private readonly AnalyzerConfigOptions m_globalOptions = new DictionaryAnalyzerConfigOptions( + new Dictionary(StringComparer.OrdinalIgnoreCase) + { + ["build_property.FactoryGenerator_EmitStaticExtensions"] = emitStaticExtensions ? "true" : "false" + }); + + public override AnalyzerConfigOptions GlobalOptions => m_globalOptions; + + public override AnalyzerConfigOptions GetOptions(SyntaxTree tree) => EmptyAnalyzerConfigOptions.Instance; + + public override AnalyzerConfigOptions GetOptions(AdditionalText textFile) => EmptyAnalyzerConfigOptions.Instance; +} + +internal sealed class DictionaryAnalyzerConfigOptions(IReadOnlyDictionary values) : AnalyzerConfigOptions +{ + public override bool TryGetValue(string key, out string value) + { + if (values.TryGetValue(key, out var foundValue)) + { + value = foundValue; + return true; + } + + value = string.Empty; + return false; + } +} + +internal sealed class EmptyAnalyzerConfigOptions : AnalyzerConfigOptions +{ + public static EmptyAnalyzerConfigOptions Instance { get; } = new(); + + public override bool TryGetValue(string key, out string value) + { + value = string.Empty; + return false; + } +} \ No newline at end of file diff --git a/Benchmarking/Benchmarks/Program.cs b/Benchmarking/Benchmarks/Program.cs index ed35ed6..5d93083 100644 --- a/Benchmarking/Benchmarks/Program.cs +++ b/Benchmarking/Benchmarks/Program.cs @@ -7,6 +7,8 @@ namespace Benchmarks; +// ── Dictionary-based resolution (existing path) ────────────────────────────── + [MemoryDiagnoser] [JsonExporterAttribute.Full] [JsonExporterAttribute.FullCompressed] @@ -33,13 +35,44 @@ public class ResolveBenchmarks public IContainer Create() => new DependencyInjectionContainer(default, default, default!); [Benchmark] - public IContainer CreateFromSelf() => new DependencyInjectionContainer(m_container); + public void CreateFromSelf() + { + // Child containers attach to their base until disposed, so each benchmark + // invocation must clean up or the inheritor chain grows across operations. + using var child = new DependencyInjectionContainer(m_container); + } + + // ── Static-extension resolution (C# 14 / .NET 10+ path) ───────────────────── + // + // Each Resolve(container?) call inlines the full construction chain directly — + // no dictionary lookup, no factory-method indirection. + // + // Null-container variants bypass the singleton cache entirely and perform a + // fresh allocation on every call, exposing the raw construction cost. + [Benchmark] + public ISingleton ExtensionResolveSingleton() => ISingleton.Resolve(m_container); + + [Benchmark] + public ISingleton ExtensionResolveSingletonNullContainer() => ISingleton.Resolve(null); + + [Benchmark] + public IOverridable ExtensionResolveTransient() => IOverridable.Resolve(m_container); + + [Benchmark] + public ChainA ExtensionResolveChain() => ChainA.Resolve(m_container); + + [Benchmark] + public ChainA ExtensionResolveChainNullContainer() => ChainA.Resolve(null); + + [Benchmark] + public ArrayConsumer ExtensionResolveWithCollection() => ArrayConsumer.Resolve(m_container); + + [Benchmark] + public ArrayConsumer ExtensionResolveWithCollectionNullContainer() => ArrayConsumer.Resolve(null); } internal static class Program { - private static void Main(string[] args) - { - var summary = BenchmarkRunner.Run(); - } + private static void Main(string[] args) => + BenchmarkSwitcher.FromAssembly(typeof(Program).Assembly).Run(args); } \ No newline at end of file diff --git a/Directory.Packages.props b/Directory.Packages.props index 70bdb68..0291d13 100644 --- a/Directory.Packages.props +++ b/Directory.Packages.props @@ -1,33 +1,26 @@ - - true - true - - - - - - - - - - - - - - - - - all - runtime; build; native; contentfiles; analyzers; buildtransitive - - - - all - runtime; build; native; contentfiles; analyzers; buildtransitive - - - - - + + true + true + + + + + + + + + + + + + + + all + runtime; build; native; contentfiles; analyzers; buildtransitive + + + + + \ No newline at end of file diff --git a/FactoryGenerator.Attributes/Attributes/InjectionPriorityAttribute.cs b/FactoryGenerator.Attributes/Attributes/InjectionPriorityAttribute.cs new file mode 100644 index 0000000..5c5d606 --- /dev/null +++ b/FactoryGenerator.Attributes/Attributes/InjectionPriorityAttribute.cs @@ -0,0 +1,9 @@ +using System; + +namespace FactoryGenerator.Attributes; + +[AttributeUsage(AttributeTargets.Assembly)] +public class InjectionPriorityAttribute(int priority) : Attribute +{ + public int Priority { get; } = priority; +} diff --git a/FactoryGenerator.Attributes/ContainerRegistry.cs b/FactoryGenerator.Attributes/ContainerRegistry.cs index 65e43c6..6031c61 100644 --- a/FactoryGenerator.Attributes/ContainerRegistry.cs +++ b/FactoryGenerator.Attributes/ContainerRegistry.cs @@ -42,9 +42,12 @@ public static IContainer BuildChain(IContainer baseContainer) snapshot = s_registrations.OrderBy(r => r.Priority).ToList(); } + var existingAssemblies = GetAssemblyNames(baseContainer); var current = baseContainer; foreach (var registration in snapshot) { + if (!existingAssemblies.Add(registration.AssemblyName)) + continue; current = registration.Factory(current); } @@ -65,9 +68,13 @@ public static IContainer BuildChain(IContainer baseContainer, IEnumerable(s_registrations); } + var existingAssemblies = GetAssemblyNames(baseContainer); var current = baseContainer; foreach (var name in assemblyNames) { + if (!existingAssemblies.Add(name)) + continue; + var registration = snapshot.Find(r => r.AssemblyName == name); if (registration == null) { @@ -82,6 +89,25 @@ public static IContainer BuildChain(IContainer baseContainer, IEnumerable GetAssemblyNames(IContainer container) + { + var names = new HashSet(StringComparer.Ordinal); + + for (var current = container; current is not null; current = current.Base) + { + if (current is IContainerRegistrationMetadata metadata) + names.Add(metadata.AssemblyName); + } + + for (var current = container.Inheritor; current is not null; current = current.Inheritor) + { + if (current is IContainerRegistrationMetadata metadata) + names.Add(metadata.AssemblyName); + } + + return names; + } + /// /// Returns the names of all currently registered container assemblies. /// diff --git a/FactoryGenerator.Attributes/FactoryGenerator.Attributes.csproj b/FactoryGenerator.Attributes/FactoryGenerator.Attributes.csproj index 4c79949..b0de705 100644 --- a/FactoryGenerator.Attributes/FactoryGenerator.Attributes.csproj +++ b/FactoryGenerator.Attributes/FactoryGenerator.Attributes.csproj @@ -6,4 +6,8 @@ latest + + + + diff --git a/FactoryGenerator.Attributes/IContainer.cs b/FactoryGenerator.Attributes/IContainer.cs index b367ce9..b0a8f33 100644 --- a/FactoryGenerator.Attributes/IContainer.cs +++ b/FactoryGenerator.Attributes/IContainer.cs @@ -1,5 +1,4 @@ using System; -using System.Collections; using System.Collections.Generic; namespace FactoryGenerator; @@ -25,4 +24,28 @@ public interface IContainer : ILifetimeScope { IContainer? Base { get; } IContainer? Inheritor { get; set; } +} + +public interface IContainerScopeFactory +{ + ILifetimeScope BeginLifetimeScope(IContainer? baseContainer); +} + +public interface IContainerRegistrationMetadata +{ + string AssemblyName { get; } +} + +public interface IContainerCacheInvalidator +{ + void InvalidateCollectionCaches(); +} + +public interface IContainerLocalCollectionResolver +{ + bool TryResolveLocalCollection(Type type, out object? resolved); +} + +public interface IServiceProviderBackedContainer +{ } \ No newline at end of file diff --git a/FactoryGenerator.Attributes/LifetimeScopeDisposalExtensions.cs b/FactoryGenerator.Attributes/LifetimeScopeDisposalExtensions.cs new file mode 100644 index 0000000..3945d42 --- /dev/null +++ b/FactoryGenerator.Attributes/LifetimeScopeDisposalExtensions.cs @@ -0,0 +1,19 @@ +using System; +using System.Threading.Tasks; + +namespace FactoryGenerator; + +public static class LifetimeScopeDisposalExtensions +{ + public static ValueTask DisposeAsync(this ILifetimeScope scope) + { + if (scope is null) + throw new ArgumentNullException(nameof(scope)); + + if (scope is IAsyncDisposable asyncDisposable) + return asyncDisposable.DisposeAsync(); + + scope.Dispose(); + return default; + } +} diff --git a/FactoryGenerator.Attributes/ResolvedInstanceTracker.cs b/FactoryGenerator.Attributes/ResolvedInstanceTracker.cs new file mode 100644 index 0000000..d5d65ea --- /dev/null +++ b/FactoryGenerator.Attributes/ResolvedInstanceTracker.cs @@ -0,0 +1,126 @@ +using System; +using System.Collections.Generic; +using System.Threading.Tasks; + +namespace FactoryGenerator; + +#nullable enable + +public sealed class ResolvedInstanceTracker : IDisposable, IAsyncDisposable +{ + private enum DisposalMode + { + Active = 0, + Synchronous = 1, + Asynchronous = 2 + } + + private readonly object m_lock = new object(); + private List>? m_instances = new List>(); + private DisposalMode m_disposalMode; + + public void Track(object? instance) + { + if (instance is null) + return; + + if (instance is not IDisposable && instance is not IAsyncDisposable) + return; + + DisposalMode disposalMode; + lock (m_lock) + { + disposalMode = m_disposalMode; + if (disposalMode == DisposalMode.Active) + { + m_instances!.Add(new WeakReference(instance)); + return; + } + } + + if (disposalMode == DisposalMode.Asynchronous) + { + DisposeAsynchronously(instance).AsTask().GetAwaiter().GetResult(); + return; + } + + DisposeSynchronously(instance); + } + + public void Dispose() + { + var trackedInstances = BeginSynchronousDisposal(); + if (trackedInstances is null) + return; + + for (var index = trackedInstances.Count - 1; index >= 0; index--) + { + if (trackedInstances[index].TryGetTarget(out var instance)) + DisposeSynchronously(instance); + } + } + + public async ValueTask DisposeAsync() + { + var trackedInstances = BeginAsynchronousDisposal(); + if (trackedInstances is null) + return; + + for (var index = trackedInstances.Count - 1; index >= 0; index--) + { + if (trackedInstances[index].TryGetTarget(out var instance)) + await DisposeAsynchronously(instance).ConfigureAwait(false); + } + } + + private List>? BeginSynchronousDisposal() + { + lock (m_lock) + { + if (m_disposalMode != DisposalMode.Active) + return null; + + m_disposalMode = DisposalMode.Synchronous; + var instances = m_instances; + m_instances = null; + return instances; + } + } + + private List>? BeginAsynchronousDisposal() + { + lock (m_lock) + { + if (m_disposalMode != DisposalMode.Active) + return null; + + m_disposalMode = DisposalMode.Asynchronous; + var instances = m_instances; + m_instances = null; + return instances; + } + } + + private static void DisposeSynchronously(object instance) + { + if (instance is IDisposable disposable) + { + disposable.Dispose(); + return; + } + + if (instance is IAsyncDisposable asyncDisposable) + asyncDisposable.DisposeAsync().AsTask().GetAwaiter().GetResult(); + } + + private static ValueTask DisposeAsynchronously(object instance) + { + if (instance is IAsyncDisposable asyncDisposable) + return asyncDisposable.DisposeAsync(); + + if (instance is IDisposable disposable) + disposable.Dispose(); + + return default; + } +} diff --git a/FactoryGenerator.Extensions.AspNetCore/AppBuilderExtensions.cs b/FactoryGenerator.Extensions.AspNetCore/AppBuilderExtensions.cs index 59260ef..5009be6 100644 --- a/FactoryGenerator.Extensions.AspNetCore/AppBuilderExtensions.cs +++ b/FactoryGenerator.Extensions.AspNetCore/AppBuilderExtensions.cs @@ -1,6 +1,4 @@ -using System; -using FactoryGenerator; -using Microsoft.AspNetCore.Builder; +using Microsoft.AspNetCore.Builder; namespace FactoryGenerator.Extensions.AspNetCore; diff --git a/FactoryGenerator.Extensions.AspNetCore/FactoryGenerator.Extensions.AspNetCore.csproj b/FactoryGenerator.Extensions.AspNetCore/FactoryGenerator.Extensions.AspNetCore.csproj index 79f3779..95b0faa 100644 --- a/FactoryGenerator.Extensions.AspNetCore/FactoryGenerator.Extensions.AspNetCore.csproj +++ b/FactoryGenerator.Extensions.AspNetCore/FactoryGenerator.Extensions.AspNetCore.csproj @@ -1,7 +1,7 @@  - net8.0 + net10.0 FactoryGenerator.Extensions.AspNetCore latest diff --git a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs index 94a7e24..fe33c4d 100644 --- a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs +++ b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorMiddleware.cs @@ -1,7 +1,5 @@ using System.Threading.Tasks; -using FactoryGenerator; using Microsoft.AspNetCore.Http; -using Microsoft.Extensions.DependencyInjection; namespace FactoryGenerator.Extensions.AspNetCore; @@ -18,8 +16,11 @@ public FactoryGeneratorMiddleware(RequestDelegate next, IContainer container) public async Task Invoke(HttpContext context) { - var scope = _container.BeginLifetimeScope(); var originalProvider = context.RequestServices; + var requestContainer = new ServiceProviderAdapter(originalProvider, baseContainer: _container); + var scope = _container is IContainerScopeFactory scopeFactory + ? scopeFactory.BeginLifetimeScope(requestContainer) + : _container.BeginLifetimeScope(); context.RequestServices = new FactoryGeneratorServiceProvider(originalProvider, scope); try @@ -28,10 +29,9 @@ public async Task Invoke(HttpContext context) } finally { - // The FactoryGeneratorServiceProvider.Dispose will dispose the scope if (context.RequestServices is FactoryGeneratorServiceProvider wrapper) { - wrapper.Dispose(); + await wrapper.DisposeAsync(); } context.RequestServices = originalProvider; } diff --git a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorServiceProvider.cs b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorServiceProvider.cs index fdfdcad..407ab6b 100644 --- a/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorServiceProvider.cs +++ b/FactoryGenerator.Extensions.AspNetCore/FactoryGeneratorServiceProvider.cs @@ -1,11 +1,11 @@ using System; -using FactoryGenerator; +using System.Threading.Tasks; using Microsoft.Extensions.DependencyInjection; namespace FactoryGenerator.Extensions.AspNetCore; #nullable enable -internal sealed class FactoryGeneratorServiceProvider : IServiceProvider, ISupportRequiredService, IDisposable +internal sealed class FactoryGeneratorServiceProvider : IServiceProvider, ISupportRequiredService, IDisposable, IAsyncDisposable { private readonly IServiceProvider _baseProvider; private readonly ILifetimeScope _scope; @@ -51,4 +51,13 @@ public void Dispose() { _scope.Dispose(); } + + public ValueTask DisposeAsync() + { + if (_scope is IAsyncDisposable asyncDisposable) + return asyncDisposable.DisposeAsync(); + + _scope.Dispose(); + return default; + } } diff --git a/FactoryGenerator.Extensions.AspNetCore/ServiceProviderAdapter.cs b/FactoryGenerator.Extensions.AspNetCore/ServiceProviderAdapter.cs index a2dfbc2..174168a 100644 --- a/FactoryGenerator.Extensions.AspNetCore/ServiceProviderAdapter.cs +++ b/FactoryGenerator.Extensions.AspNetCore/ServiceProviderAdapter.cs @@ -1,22 +1,25 @@ using System; using System.Collections.Generic; +using System.Threading.Tasks; using Microsoft.Extensions.DependencyInjection; namespace FactoryGenerator.Extensions.AspNetCore; #nullable enable -internal sealed class ServiceProviderAdapter : IContainer, IDisposable +internal sealed class ServiceProviderAdapter : IContainer, IDisposable, IAsyncDisposable, IContainerLocalCollectionResolver, IServiceProviderBackedContainer { private readonly IServiceProvider _serviceProvider; private readonly IServiceScope? _serviceScope; + private readonly IContainer? _baseContainer; - public ServiceProviderAdapter(IServiceProvider serviceProvider, IServiceScope? serviceScope = null) + public ServiceProviderAdapter(IServiceProvider serviceProvider, IServiceScope? serviceScope = null, IContainer? baseContainer = null) { _serviceProvider = serviceProvider; _serviceScope = serviceScope; + _baseContainer = baseContainer; } - public IContainer? Base => null; + public IContainer? Base => _baseContainer; public IContainer? Inheritor { get; set; } public void Dispose() @@ -24,10 +27,20 @@ public void Dispose() _serviceScope?.Dispose(); } + public ValueTask DisposeAsync() + { + if (_serviceScope is IAsyncDisposable asyncDisposable) + return asyncDisposable.DisposeAsync(); + + _serviceScope?.Dispose(); + return default; + } + public T Resolve() { var service = _serviceProvider.GetService(); if (service != null) return service; + if (_baseContainer is not null) return _baseContainer.Resolve(); throw new KeyNotFoundException($"The type {typeof(T)} has not been registered in the IServiceProvider."); } @@ -35,40 +48,55 @@ public object Resolve(Type type) { var service = _serviceProvider.GetService(type); if (service != null) return service; + if (_baseContainer is not null) return _baseContainer.Resolve(type); throw new KeyNotFoundException($"The type {type} has not been registered in the IServiceProvider."); } public bool TryResolve(Type type, out object? resolved) { resolved = _serviceProvider.GetService(type); - return resolved != null; + if (resolved is not null) return true; + if (_baseContainer is not null) return _baseContainer.TryResolve(type, out resolved); + return false; } public bool TryResolve(out T? resolved) { resolved = _serviceProvider.GetService(); - return resolved != null; + if (resolved is not null) return true; + if (_baseContainer is not null) return _baseContainer.TryResolve(out resolved); + return false; + } + + public bool TryResolveLocalCollection(Type type, out object? resolved) + { + resolved = _serviceProvider.GetService(type); + return resolved is not null; } public bool IsRegistered(Type type) { // IServiceProvider doesn't have a reliable IsRegistered method without resolution. // We return true if it can be resolved. - return _serviceProvider.GetService(type) != null; + return _serviceProvider.GetService(type) != null || _baseContainer?.IsRegistered(type) == true; } public bool IsRegistered() => IsRegistered(typeof(T)); - public bool GetBoolean(string key) => false; + public bool GetBoolean(string key) => _baseContainer?.GetBoolean(key) == true; public IEnumerable<(string Key, bool Value)> GetBooleans() { - yield break; + if (_baseContainer is null) + yield break; + + foreach (var boolean in _baseContainer.GetBooleans()) + yield return boolean; } public ILifetimeScope BeginLifetimeScope() { - var scope = _serviceProvider.CreateScope(); - return new ServiceProviderAdapter(scope.ServiceProvider, scope); + var scope = _serviceProvider.CreateAsyncScope(); + return new ServiceProviderAdapter(scope.ServiceProvider, scope, _baseContainer); } } diff --git a/FactoryGenerator.Extensions.AspNetCore/ServiceProviderExtensions.cs b/FactoryGenerator.Extensions.AspNetCore/ServiceProviderExtensions.cs index 1bf54f0..d045e3e 100644 --- a/FactoryGenerator.Extensions.AspNetCore/ServiceProviderExtensions.cs +++ b/FactoryGenerator.Extensions.AspNetCore/ServiceProviderExtensions.cs @@ -1,5 +1,4 @@ using System; -using FactoryGenerator; namespace FactoryGenerator.Extensions.AspNetCore; diff --git a/FactoryGenerator/FactoryGenerator.cs b/FactoryGenerator/FactoryGenerator.cs index be14e7f..16ab857 100644 --- a/FactoryGenerator/FactoryGenerator.cs +++ b/FactoryGenerator/FactoryGenerator.cs @@ -4,8 +4,9 @@ using System.Linq; using System.Text; using System.Threading; +using System.Threading.Tasks; using Microsoft.CodeAnalysis; -using Microsoft.CodeAnalysis.CSharp.Syntax; +using Microsoft.CodeAnalysis.CSharp; using Microsoft.CodeAnalysis.Diagnostics; namespace FactoryGenerator @@ -29,10 +30,15 @@ public void Initialize(IncrementalGeneratorInitializationContext context) var rest = references.SelectMany(FindMethods); var attributes = rest.Collect(); var compilation = context.CompilationProvider; - var syntaxUsages = context.SyntaxProvider.CreateSyntaxProvider(ResolveSymbols, ResolveTransformations) - .Collect(); - var combined = attributes.Combine(compilation).Combine(syntaxUsages).Combine(logProvider); + var combined = attributes.Combine(compilation).Combine(logProvider); context.RegisterSourceOutput(combined, MakeAutofacModule); + + var supportsStaticExtensions = context.ParseOptionsProvider.Select(IsAtLeastCSharp14); + var emitStaticExtensions = context.AnalyzerConfigOptionsProvider.Select(GetEmitStaticExtensions); + var staticExtensionsEnabled = supportsStaticExtensions.Combine(emitStaticExtensions) + .Select(static (pair, _) => pair.Left && pair.Right); + var extensionData = attributes.Combine(compilation).Combine(staticExtensionsEnabled); + context.RegisterSourceOutput(extensionData, MakeStaticExtensions); } private IncrementalValueProvider SetupLog(IncrementalGeneratorInitializationContext context) @@ -53,15 +59,13 @@ public void Initialize(IncrementalGeneratorInitializationContext context) } private void MakeAutofacModule(SourceProductionContext context, - (((ImmutableArray Injections, Compilation Compilation) Left, ImmutableArray CompileTimeResolvedTypes) Left, LoggingOptions? log) - data) + ((ImmutableArray Injections, Compilation Compilation) Left, LoggingOptions? log) data) { - var injections = data.Left.Left.Injections; - var compilation = data.Left.Left.Compilation; - var usages = data.Left.CompileTimeResolvedTypes; + var injections = data.Left.Injections; + var compilation = data.Left.Compilation; var log = data.log?.FileName == null ? NullLogger.Instance : new Logger(data.log.FileName, data.log.LogLevel); - var source = GenerateCode(injections, compilation, usages, log).ToArray(); + var source = GenerateCode(injections, compilation, log).ToArray(); context.AddSource("DependencyInjectionContainer.Lookup.g.cs", source[0]); context.AddSource("DependencyInjectionContainer.Constructor.g.cs", source[1]); context.AddSource("DependencyInjectionContainer.Declarations.g.cs", source[2]); @@ -73,30 +77,6 @@ private void MakeAutofacModule(SourceProductionContext context, context.AddSource("ContainerEntryPoint.g.cs", source[8]); } - private UsageData? ResolveTransformations(GeneratorSyntaxContext context, CancellationToken token) - { - var typeArguments = context.Node.DescendantNodes().OfType().FirstOrDefault(); - if (typeArguments is null) return null; - var identifier = typeArguments.DescendantNodes().FirstOrDefault(); - if (identifier is null) return null; - var info = context.SemanticModel.GetSymbolInfo(identifier, token); - if (info.Symbol is not INamedTypeSymbol symbol) return null; - if (!SymbolUtility.IsEnumerable(symbol)) return null; - if (symbol.TypeArguments.Length != 1) return null; - var elemType = symbol.TypeArguments[0]; - return new UsageData( - fullName: symbol.ToString()!, - memberName: SymbolUtility.MemberName(symbol).Replace("()", ""), - elementTypeFullName: elemType.ToString()!, - elementTypeMemberName: SymbolUtility.MemberName(elemType).Replace("()", "")); - } - - private bool ResolveSymbols(SyntaxNode node, CancellationToken token) - { - if (node is not MemberAccessExpressionSyntax invocation) return false; - return invocation.ToString().Contains("Resolve"); - } - private static IEnumerable FindMethods(INamespaceSymbol namespaceSymbol, CancellationToken token) { foreach (var type in SymbolUtility.GetAllTypes(namespaceSymbol)) @@ -147,15 +127,16 @@ private static INamespaceSymbol GetGlobalNamespace(Compilation compilation, Canc private const string LifetimeName = "LifetimeScope"; private static IEnumerable GenerateCode(ImmutableArray dataInjections, - Compilation compilation, ImmutableArray usages, ILogger log) + Compilation compilation, ILogger log) { - CheckForCycles(dataInjections); + 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; @@ -165,7 +146,7 @@ namespace {compilation.Assembly.Name}.Generated; [GeneratedCode(""{ToolName}"", ""{Version}"")] #nullable enable #pragma warning disable CS0169, CS0414 -public sealed partial class {ClassName} : IContainer +public sealed partial class {ClassName} : IContainer, IContainerScopeFactory, IContainerRegistrationMetadata, IContainerCacheInvalidator, IAsyncDisposable, IContainerLocalCollectionResolver {{ #pragma warning restore CS0169, CS0414 @@ -187,19 +168,82 @@ private IContainer GetTop() }} return top; }} + private void AttachToBase(IContainer baseContainer) + {{ + if (baseContainer.Inheritor is null) + {{ + baseContainer.Inheritor = this; + InvalidateCollectionCachesInChain(); + return; + }} + + var current = baseContainer.Inheritor; + while (current!.Inheritor is not null) + {{ + current = current.Inheritor; + }} + + current.Inheritor = this; + InvalidateCollectionCachesInChain(); + }} + private void DetachFromBase() + {{ + if (Base is null) + return; + + if (Base.Inheritor == this) + {{ + Base.Inheritor = Inheritor; + Inheritor = null; + InvalidateCollectionCachesInChain(); + return; + }} + + var current = Base.Inheritor; + while (current is not null && current.Inheritor != this) + {{ + current = current.Inheritor; + }} + + if (current is null) + return; + + current.Inheritor = Inheritor; + Inheritor = null; + InvalidateCollectionCachesInChain(); + }} + private void InvalidateCollectionCachesInChain() + {{ + var current = GetRoot(); + while (current is not null) + {{ + if (current is IContainerCacheInvalidator invalidator) + invalidator.InvalidateCollectionCaches(); + + current = current.Inheritor; + }} + }} + public string AssemblyName => ""{compilation.Assembly.Name}""; public IContainer? Base {{ get; }} public IContainer? Inheritor {{ get; set; }} - private readonly object m_lock = new(); + internal readonly object m_lock = new(); private Dictionary> m_lookup; + private Dictionary> m_localCollectionLookup; private Dictionary m_booleans; - private List>? resolvedInstances; + private readonly ResolvedInstanceTracker m_resolvedInstances = new(); - private List> GetResolvedInstances() + internal void TrackResolvedInstance(object instance) => m_resolvedInstances.Track(instance); + + public bool TryResolveLocalCollection(Type type, out object? resolved) {{ - if (resolvedInstances is null) - lock (m_lock) - resolvedInstances ??= new List>(); - return resolvedInstances; + if (m_localCollectionLookup.TryGetValue(type, out var factory)) + {{ + resolved = factory(); + return true; + }} + + resolved = default; + return false; }} public T Resolve() @@ -222,18 +266,14 @@ public object Resolve(Type type) public void Dispose() {{ - if (resolvedInstances is not null) - {{ - foreach (var weakReference in resolvedInstances) - {{ - if(weakReference.TryGetTarget(out var disposable)) - {{ - disposable.Dispose(); - }} - }} - resolvedInstances.Clear(); - }} - Base?.Dispose(); + DetachFromBase(); + m_resolvedInstances.Dispose(); + }} + + public ValueTask DisposeAsync() + {{ + DetachFromBase(); + return m_resolvedInstances.DisposeAsync(); }} public bool TryResolve(Type type, out object? resolved) @@ -279,38 +319,10 @@ public bool GetBoolean(string key) }} }}"; - var booleans = dataInjections.Select(inj => inj.BooleanInjection).Where(b => b is not null) - .Select(b => b!.Key).Distinct().ToArray(); - var allArguments = booleans.Select(b => $"bool {b}").ToList(); - var justBooleans = allArguments.ToList(); - var allParameters = booleans.Select(b => $"{b}").ToList(); - var ordered = dataInjections.Reverse().ToList(); - - foreach (var injection in ordered.ToArray()) - { - log.Log(LogLevel.Debug, $"Traversing {injection.Name}"); - if (!injection.IsTestType) continue; - ordered.Remove(injection); - ordered.Add(injection); - } - - var interfaceInjectors = new Dictionary>(); - var interfaceMemberNames = new Dictionary(); - - foreach (var injection in ordered) - { - for (int i = 0; i < injection.InterfaceFullNames.Length; i++) - { - var ifaceFull = injection.InterfaceFullNames[i]; - var ifaceMember = injection.InterfaceMemberNames[i]; - if (!interfaceInjectors.ContainsKey(ifaceFull)) - { - interfaceInjectors[ifaceFull] = new List(); - interfaceMemberNames[ifaceFull] = ifaceMember; - } - interfaceInjectors[ifaceFull].Add(injection); - } - } + var booleanKeys = dataInjections.Select(inj => inj.BooleanInjection).Where(b => b is not null) + .Select(b => b!.Key).Distinct().ToArray(); + var ordered = OrderInjections(dataInjections, compilation, log); + var (interfaceInjectors, interfaceMemberNames) = BuildInterfaceInjectors(ordered); var declarations = new Dictionary(); var scopedDeclarations = new Dictionary(); @@ -322,7 +334,7 @@ public bool GetBoolean(string key) declarations[injection.Name] = Declaration(injection, availableInterfaceFullNames, false); scopedDeclarations[injection.Name] = Declaration(injection, availableInterfaceFullNames, true); - var missing = GetBestConstructorMissing(injection, availableInterfaceFullNames); + var missing = GetInjectionMissingParameters(injection, availableInterfaceFullNames); foreach (var param in missing) { var key = param.TypeFullName + " " + param.Name; @@ -331,6 +343,61 @@ public bool GetBoolean(string key) } } + var localizedParameters = new List(); + foreach (var parameter in constructorParameters.ToArray()) + { + if (!parameter.IsCollection) continue; + if (parameter.CollectionElementFullName is null) continue; + constructorParameters.Remove(parameter); + localizedParameters.Add(parameter); + } + + foreach (var parameter in constructorParameters.ToArray()) + { + if (!parameter.TypeFullName.Contains("IContainer")) continue; + log.Log(LogLevel.Debug, $"Registering {parameter.Name} as Self"); + declarations[parameter.Name] = $"private IContainer {parameter.Name} => this;"; + scopedDeclarations[parameter.Name] = $"private IContainer {parameter.Name} => this;"; + constructorParameters.Remove(parameter); + } + + ValidateExternalParameterTypes(constructorParameters); + + var booleanReservedNames = constructorParameters.Select(parameter => parameter.Name) + .Concat(localizedParameters.Select(parameter => "coll_" + parameter.CollectionElementMemberName!)) + .Concat(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]; @@ -352,19 +419,11 @@ public bool GetBoolean(string key) } else { - var keys = possibilities.Select(i => i.BooleanInjection?.Key).OfType().Distinct().Reverse().ToArray(); - var fallback = possibilities.LastOrDefault(p => p.BooleanInjection == null); - var last = keys.Last(); - var ternary = new StringBuilder(); - foreach (var key in keys) - { - var trueValue = possibilities.LastOrDefault(p => - p.BooleanInjection?.Value == true && p.BooleanInjection?.Key == key); - trueValue ??= fallback; - ternary.Append(key == last - ? $"{key} ? {trueValue?.Name ?? "null!"} : {fallback?.Name ?? "null!"}" - : $"{key} ? {trueValue?.Name ?? "null!"} : "); - } + var ternary = BuildBooleanSelectionExpression( + ifaceFull, + possibilities, + booleanIdentifiers, + possibility => possibility.Name); if (!declarations.ContainsKey(ifaceMethodName)) { @@ -375,83 +434,76 @@ public bool GetBoolean(string key) } } - var localizedParameters = new List(); var arrayDeclarations = new Dictionary(); - - foreach (var parameter in constructorParameters.ToArray()) + foreach (var pair in interfaceInjectors) { - if (!parameter.IsCollection) continue; - if (parameter.CollectionElementFullName is null) continue; - var name = parameter.Name; - log.Log(LogLevel.Debug, $"Creating Collection: {name} of element type {parameter.CollectionElementFullName}"); - MakeArray(arrayDeclarations, name, parameter.CollectionElementFullName, parameter.CollectionElementMemberName!, interfaceInjectors); - constructorParameters.Remove(parameter); - localizedParameters.Add(parameter); + var name = "coll_" + interfaceMemberNames[pair.Key]; + if (arrayDeclarations.ContainsKey(name)) + continue; + log.Log(LogLevel.Debug, $"Creating Collection: {name} of element type {pair.Key}"); + MakeArray(arrayDeclarations, name, pair.Key, interfaceInjectors, booleanIdentifiers); } - var requestedUsages = new List(); - - foreach (var request in usages) + foreach (var parameter in localizedParameters) { - if (request is null) continue; - if (localizedParameters.Any(p => p.TypeFullName == request.FullName)) continue; - log.Log(LogLevel.Information, $"Creating Requested: {request.FullName}"); - log.Log(LogLevel.Debug, $"Creating Array: {request.MemberName} of type {request.ElementTypeFullName}[]"); - MakeArray(arrayDeclarations, request.MemberName, request.ElementTypeFullName, request.ElementTypeMemberName, interfaceInjectors, true); - requestedUsages.Add(request); + var name = "coll_" + parameter.CollectionElementMemberName!; + if (arrayDeclarations.ContainsKey(name)) + continue; + log.Log(LogLevel.Debug, $"Creating Collection: {name} of element type {parameter.CollectionElementFullName}"); + MakeArray(arrayDeclarations, name, parameter.CollectionElementFullName!, interfaceInjectors, booleanIdentifiers); } - foreach (var parameter in constructorParameters.ToArray()) - { - if (!parameter.TypeFullName.Contains("IContainer")) continue; - log.Log(LogLevel.Debug, $"Registering {parameter.Name} as Self"); - declarations[parameter.Name] = $"private IContainer {parameter.Name} => this;"; - scopedDeclarations[parameter.Name] = $"private IContainer {parameter.Name} => this;"; - constructorParameters.Remove(parameter); - } + var externalParameters = constructorParameters.OrderBy(parameter => parameter.TypeFullName).ToList(); + var allArguments = booleanParameters.Select(parameter => $"bool {parameter.Identifier}").ToList(); + allArguments.AddRange(externalParameters.Select(parameter => $"{parameter.TypeFullName} {parameter.Name}").Distinct()); - var arguments = constructorParameters.OrderBy(p => p.TypeFullName).Select(p => $"{p.TypeFullName} {p.Name}").Distinct(); - var parameters = constructorParameters.OrderBy(p => p.TypeFullName).Select(p => p.Name).Distinct(); - allArguments.AddRange(arguments); var lifetimeArguments = allArguments.ToList(); - allParameters.AddRange(parameters); + lifetimeArguments.Insert(0, "IContainer? baseContainer"); + lifetimeArguments.Insert(0, $"{ClassName} fallback"); + var lifetimeParameters = new List { "this", "baseContainer" }; + lifetimeParameters.AddRange(booleanParameters.Select(parameter => parameter.Identifier)); + lifetimeParameters.AddRange(externalParameters.Select(parameter => $"baseContainer != null ? baseContainer.Resolve<{parameter.TypeFullName}>() : {parameter.Name}")); var constructor = "(" + string.Join(", ", allArguments) + ")"; - lifetimeArguments.Insert(0, $"{ClassName} fallback"); - allParameters.Insert(0, "this"); var lifetimeConstructor = "(" + string.Join(", ", lifetimeArguments) + ")"; - var lifetimeParameters = string.Join(", ", allParameters); + var lifetimeParameterValues = string.Join(", ", lifetimeParameters); log.Log(LogLevel.Debug, $"Resulting Constructor: {constructor}"); - var constructorFields = string.Join("\n\t", allArguments.Select(arg => arg + ";")); + var constructorFields = string.Join("\n\t", allArguments.Select(arg => "internal " + arg + ";")); var constructorAssignments = string.Join("\n\t\t", allArguments.Select(arg => arg.Split(' ').Last()).Select(arg => $"this.{arg} = {arg};")); var resolvedConstructorAssignments = string.Join("\n\t\t", - allArguments.Select(a => a.Split(' ')).Where(a => a[0] != "bool") - .Select(a => $"this.{a[1]} = Base.Resolve<{a[0]}>();")); + externalParameters.Select(parameter => $"this.{parameter.Name} = Base.Resolve<{parameter.TypeFullName}>();")); var interfacePairs = interfaceInjectors.Keys.Select(k => (TypeName: k, MemberName: interfaceMemberNames[k])).ToList(); // ReadOnlySpan is a ref struct and cannot be placed in the lookup dictionary var localizedForDict = localizedParameters.Where(p => p.CollectionKind != CollectionKind.ReadOnlySpan).ToList(); - var localizedPairs = localizedForDict - .Select(p => (TypeName: p.TypeFullName, Expression: CollectionDictExpression(p.CollectionKind, p.Name))) + var localizedPairs = DistinctByTypeName(localizedForDict + .Select(p => (TypeName: p.TypeFullName, Expression: CollectionDictExpression(p.CollectionKind, "coll_" + p.CollectionElementMemberName!))) + .ToList(), pair => pair.TypeName); + var localizedTypes = new HashSet(localizedPairs.Select(pair => pair.TypeName)); + var enumerablePairs = interfaceInjectors.Keys + .Select(key => (TypeName: $"System.Collections.Generic.IEnumerable<{key}>", Expression: "coll_" + interfaceMemberNames[key])) + .Where(pair => !localizedTypes.Contains(pair.TypeName)) .ToList(); - var requestedPairs = requestedUsages.Select(u => (TypeName: u.FullName, MemberName: u.MemberName)).ToList(); - var constructorPairs = constructorParameters.Select(p => (TypeName: p.TypeFullName, Expression: p.Name)).ToList(); + var 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 = interfaceInjectors.Count + localizedForDict.Count + requestedUsages.Count + constructorParameters.Count; + var dictSize = interfacePairs.Count + localizedPairs.Count + enumerablePairs.Count + constructorPairs.Count; yield return Constructor(usingStatements, constructorFields, constructor, constructorAssignments, - dictSize, interfacePairs, localizedPairs, requestedPairs, constructorPairs, - true, ClassName, lifetimeParameters, - resolvingConstructorAssignments: resolvedConstructorAssignments, booleans: justBooleans); + dictSize, interfacePairs, localizedPairs, enumerablePairs, constructorPairs, 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 +public sealed partial class LifetimeScope : IContainer, IContainerScopeFactory, IContainerRegistrationMetadata, IContainerCacheInvalidator, IAsyncDisposable, IContainerLocalCollectionResolver {{ #pragma warning restore CS0169, CS0414 private IContainer GetRoot() @@ -472,27 +524,95 @@ private IContainer GetTop() }} return top; }} + private void AttachToBase(IContainer baseContainer) + {{ + if (baseContainer.Inheritor is null) + {{ + baseContainer.Inheritor = this; + InvalidateCollectionCachesInChain(); + return; + }} + + var current = baseContainer.Inheritor; + while (current!.Inheritor is not null) + {{ + current = current.Inheritor; + }} + + current.Inheritor = this; + InvalidateCollectionCachesInChain(); + }} + private void DetachFromBase() + {{ + if (Base is null) + return; + + if (Base.Inheritor == this) + {{ + Base.Inheritor = Inheritor; + Inheritor = null; + InvalidateCollectionCachesInChain(); + return; + }} + + var current = Base.Inheritor; + while (current is not null && current.Inheritor != this) + {{ + current = current.Inheritor; + }} + + if (current is null) + return; + + current.Inheritor = Inheritor; + Inheritor = null; + InvalidateCollectionCachesInChain(); + }} + private void InvalidateCollectionCachesInChain() + {{ + var current = GetRoot(); + while (current is not null) + {{ + if (current is IContainerCacheInvalidator invalidator) + invalidator.InvalidateCollectionCaches(); + + current = current.Inheritor; + }} + }} + public string AssemblyName => ""{compilation.Assembly.Name}""; public IContainer? Base {{ get; }} public IContainer? Inheritor {{ get; set; }} public ILifetimeScope BeginLifetimeScope() {{ - var scope = m_fallback.BeginLifetimeScope(); - GetResolvedInstances().Add(new WeakReference(scope)); + var baseContainer = Base?.BeginLifetimeScope() as IContainer; + return BeginLifetimeScope(baseContainer); + }} + public ILifetimeScope BeginLifetimeScope(IContainer? baseContainer) + {{ + var scope = m_fallback.BeginLifetimeScope(baseContainer); + TrackResolvedInstance(scope); return scope; }} - private readonly object m_lock = new(); + internal readonly object m_lock = new(); private {ClassName} m_fallback; private Dictionary> m_lookup; + private Dictionary> m_localCollectionLookup; private Dictionary m_booleans; - private List>? resolvedInstances; + private readonly ResolvedInstanceTracker m_resolvedInstances = new(); - private List> GetResolvedInstances() + internal void TrackResolvedInstance(object instance) => m_resolvedInstances.Track(instance); + + public bool TryResolveLocalCollection(Type type, out object? resolved) {{ - if (resolvedInstances is null) - lock (m_lock) - resolvedInstances ??= new List>(); - return resolvedInstances; + if (m_localCollectionLookup.TryGetValue(type, out var factory)) + {{ + resolved = factory(); + return true; + }} + + resolved = default; + return false; }} public T Resolve() @@ -515,18 +635,14 @@ public object Resolve(Type type) public void Dispose() {{ - if (resolvedInstances is not null) - {{ - foreach (var weakReference in resolvedInstances) - {{ - if(weakReference.TryGetTarget(out var disposable)) - {{ - disposable.Dispose(); - }} - }} - resolvedInstances.Clear(); - }} - Base?.Dispose(); + DetachFromBase(); + m_resolvedInstances.Dispose(); + }} + + public ValueTask DisposeAsync() + {{ + DetachFromBase(); + return m_resolvedInstances.DisposeAsync(); }} public bool TryResolve(Type type, out object? resolved) @@ -575,9 +691,9 @@ public bool GetBoolean(string key) "; yield return Constructor(usingStatements, constructorFields, lifetimeConstructor, constructorAssignments, - dictSize, interfacePairs, localizedPairs, requestedPairs, constructorPairs, + dictSize, interfacePairs, localizedPairs, enumerablePairs, constructorPairs, localCollectionPairs, false, LifetimeName, - resolvingConstructorAssignments: resolvedConstructorAssignments, addMergingConstructor: false, booleans: justBooleans); + resolvingConstructorAssignments: resolvedConstructorAssignments, addMergingConstructor: false, booleans: booleanParameters); yield return Declarations(usingStatements, scopedDeclarations, LifetimeName); yield return ArrayDeclarations(usingStatements, arrayDeclarations, LifetimeName); @@ -625,129 +741,481 @@ internal static void Register() "; } - private static void CheckForCycles(ImmutableArray dataInjections) + 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 tree = new Dictionary>(); - foreach (var injection in dataInjections) + 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) { - if (injection.Lambda != null) continue; - var node = new List(); - foreach (var ifaceName in injection.InterfaceFullNames) + var current = queue.Dequeue(); + if (relevantAssemblyNames.Contains(current.Assembly.Name) + && (!distances.TryGetValue(current.Assembly.Name, out var existingDistance) + || current.Distance < existingDistance)) { - if (!tree.ContainsKey(ifaceName)) - tree[ifaceName] = node; + distances[current.Assembly.Name] = current.Distance; } - foreach (var ctor in injection.Constructors) + foreach (var referencedAssembly in GetReferencedAssemblies(current.Assembly)) { - foreach (var parameter in ctor.Parameters) - { - string? depName; - if (parameter.IsCollection) - { - if (parameter.CollectionElementFullName is null) continue; - depName = parameter.CollectionElementFullName; - } - else - { - // Strip ? so nullable params resolve to their underlying type in the cycle graph - depName = parameter.IsNullable - ? parameter.TypeFullName.TrimEnd('?') - : parameter.TypeFullName; - } + if (!visited.Add(referencedAssembly)) + continue; - node.Add(depName); - if (tree.TryGetValue(depName, out var list)) - { - foreach (var ifaceName in injection.InterfaceFullNames) - { - if (list.Contains(ifaceName)) - throw new InvalidOperationException( - $"Cyclic Dependency Detected between {injection.TypeFullName} and {ifaceName}"); - } - } - } + queue.Enqueue((referencedAssembly, current.Distance + 1)); } } + + return distances; } - 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 MemberName)> requestedPairs, IEnumerable<(string TypeName, string Expression)> constructorParamPairs, - bool addLifetimeScopeFunction, string className, string? lifetimeParameters = null, - string? fromConstructor = null, string? resolvingConstructorAssignments = null, bool addMergingConstructor = true, List booleans = null!) + private static IEnumerable GetReferencedAssemblies(IAssemblySymbol assembly) { - var lifetimeScopeFunction = addLifetimeScopeFunction - ? $@" -public ILifetimeScope BeginLifetimeScope() -{{ - var scope = new {LifetimeName}({lifetimeParameters}); - GetResolvedInstances().Add(new WeakReference(scope)); - return scope; -}}" : string.Empty; + foreach (var module in assembly.Modules) + { + foreach (var referencedAssembly in module.ReferencedAssemblySymbols) + yield return referencedAssembly; + } + } - var mergingConstructor = addMergingConstructor ? $@" -public {className}(IContainer Base{fromConstructor}) -{{ - this.Base = Base; - Base.Inheritor = this; - {resolvingConstructorAssignments} - -{string.Join("\n", booleans.Select(b => b.Split(' ').Last()).Select(b => $"\t this.{b} = Base.GetBoolean(\"{b}\");"))} - - m_lookup = new({dictSize}) {{ -{MakeDictionaryFromTypes(interfaceTypePairs)} -{MakeDictionaryFromParams(localizedParamPairs)} -{MakeDictionaryFromTypes(requestedPairs)} -{MakeDictionaryFromParams(constructorParamPairs)} - }}; - m_booleans = new(); - foreach(var (key, value) in Base.GetBooleans()) - {{ - m_booleans[key] = value; - }} -}}" : string.Empty; + 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(); - var extraConstruction = addLifetimeScopeFunction ? string.Empty : "m_fallback = fallback;"; - return $@"{usingStatements} -public partial class {className} -{{ - {constructorFields} - public {className}{constructor} - {{ - {extraConstruction} - {constructorAssignments} - - m_lookup = new({dictSize}) {{ -{MakeDictionaryFromTypes(interfaceTypePairs)} -{MakeDictionaryFromParams(localizedParamPairs)} -{MakeDictionaryFromTypes(requestedPairs)} -{MakeDictionaryFromParams(constructorParamPairs)} - }}; - - m_booleans = new({booleans.Count}) {{ -{string.Join("\n", booleans.Select(b => b.Split(' ').Last()).Select(b => $"\t\t{{ \"{b}\", {b} }},"))} - }}; - }} - {mergingConstructor} - {lifetimeScopeFunction} + 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); + } + } - private static string ArrayDeclarations(string usingStatements, Dictionary arrayDeclarations, string className) - { - return $@"{usingStatements} -public partial class {className} -{{ - {string.Join("\n\t", arrayDeclarations.Values)} -}}"; + return (interfaceInjectors, interfaceMemberNames); } - private static string Declarations(string usingStatements, Dictionary declarations, string className) + private static List GetReachableImplementations(List possibilities) { - return $@"{usingStatements} + 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)} @@ -755,62 +1223,95 @@ public partial class {className} } private static void MakeArray(Dictionary declarations, string name, - string elementTypeFullName, string elementTypeMemberName, - Dictionary> interfaceInjectors, bool function = false) - { - var factoryName = $"new {elementTypeFullName}[0]"; - var factory = string.Empty; - var functionString = function ? "()" : string.Empty; - var starter = function ? string.Empty : "get {"; - var ender = function ? string.Empty : "}"; - if (interfaceInjectors.TryGetValue(elementTypeFullName, out var injections)) - { - factoryName = $"Create{name}()".Replace("_", ""); - var nonBooleanInjections = injections.Where(i => i.BooleanInjection == null).ToList(); - var booleanInjections = injections.Where(b => b.BooleanInjection != null).ToList(); - factory = @$" + 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}>({nonBooleanInjections.Count}) {{ - {string.Join(",\n\t\t\t", nonBooleanInjections.Select(i => i.Name))} - }}; - {string.Join("\n\t\t\t", booleanInjections.Select(i => $"if({i.BooleanInjection!.Key}) source.Add({i.Name});"))} + var source = new List<{elementTypeFullName}>({localFactoryName}); var b = Base; + var frameworkCollectionSourceSeen = false; while(b is not null) {{ - if(b.TryResolve>(out var additional)) source.AddRange(additional!); + if(!(frameworkCollectionSourceSeen && b is IServiceProviderBackedContainer)) + {{ + if (b is IContainerLocalCollectionResolver localResolver) + {{ + if(localResolver.TryResolveLocalCollection(typeof(IEnumerable<{elementTypeFullName}>), out var localAdditional)) + source.AddRange((IEnumerable<{elementTypeFullName}>)localAdditional!); + }} + else if(b.TryResolve>(out var additional)) + {{ + source.AddRange(additional!); + }} + + if (b is IServiceProviderBackedContainer) + frameworkCollectionSourceSeen = true; + }} b = b.Base; }} + var inheritorFrameworkCollectionSourceSeen = false; b = Inheritor; while(b is not null) {{ - if(b.TryResolve>(out var additional)) source.AddRange(additional!); + if(!(inheritorFrameworkCollectionSourceSeen && b is IServiceProviderBackedContainer)) + {{ + if (b is IContainerLocalCollectionResolver localResolver) + {{ + if(localResolver.TryResolveLocalCollection(typeof(IEnumerable<{elementTypeFullName}>), out var localAdditional)) + source.AddRange((IEnumerable<{elementTypeFullName}>)localAdditional!); + }} + else if(b.TryResolve>(out var additional)) + {{ + source.AddRange(additional!); + }} + + if (b is IServiceProviderBackedContainer) + inheritorFrameworkCollectionSourceSeen = true; + }} b = b.Inheritor; }} Reentrant_{name} = false; return source; }}"; - } declarations[name] = $@" - internal IEnumerable<{elementTypeFullName}> {name}{functionString} + internal IEnumerable<{elementTypeFullName}> {name} {{ - {starter} - var cached = m_{name}; - if (cached != null) - return cached; - - lock (m_lock) + get {{ - cached = m_{name}; + var cached = m_{name}; if (cached != null) return cached; - return m_{name} = {factoryName}; + + lock (m_lock) + {{ + cached = m_{name}; + if (cached != null) + return cached; + return m_{name} = {factoryName}; + }} }} - {ender} }} + internal IEnumerable<{elementTypeFullName}> local_{name} => {localFactoryName}; internal IEnumerable<{elementTypeFullName}>? m_{name};" + factory; } @@ -840,9 +1341,9 @@ private static string Declaration(InjectionData injection, ImmutableArray m_fallback.{name};"; if (injection.Singleton || injection.Scoped) - return SymbolUtility.SingletonFactory(injection.TypeFullName, name, lazyName, creation, injection.Disposable); + return SymbolUtility.SingletonFactory(injection.TypeFullName, name, lazyName, creation, injection.Disposable || injection.AsyncDisposable); - if (injection.Disposable) + if (injection.Disposable || injection.AsyncDisposable) return SymbolUtility.DisposableFactory(injection.TypeFullName, name, creation); return $"internal {injection.TypeFullName} {name} => {creation};"; @@ -856,22 +1357,13 @@ private static string CreationCall(InjectionData injection, ImmutableArray? lambdaNullableDefaults = null; - if (lambda.IsMethod && lambda.MethodParameters.Length > 0) + if (lambda.IsMethod) { - foreach (var p in lambda.MethodParameters) - { - if (!p.IsNullable) continue; - var baseType = p.TypeFullName.TrimEnd('?'); - if (availableInterfaceFullNames.Contains(baseType)) continue; - lambdaNullableDefaults ??= new HashSet(); - lambdaNullableDefaults.Add(p); - } + HashSet? lambdaMissing = null; + HashSet? lambdaNullableDefaults = null; + AnalyzeParameters(lambda.MethodParameters, availableInterfaceFullNames, ref lambdaMissing, ref lambdaNullableDefaults); + return $"{lambda.ContainingTypeMemberName}.{lambda.MemberName}{MakeMethodCall(lambda.MethodParameters, lambdaMissing, lambdaNullableDefaults)}"; } - - if (lambda.IsMethod) - return $"{lambda.ContainingTypeMemberName}.{lambda.MemberName}{MakeMethodCall(lambda.MethodParameters, null, lambdaNullableDefaults)}"; else return $"{lambda.ContainingTypeMemberName}.{lambda.MemberName}"; } @@ -894,96 +1386,137 @@ private static string CreationCall(InjectionData injection, ImmutableArray(); - var localNullableDefaults = new HashSet(); - foreach (var parameter in ctor.Parameters) - { - // For nullable params, check availability of the underlying non-nullable type - var typeLookup = parameter.IsNullable - ? parameter.TypeFullName.TrimEnd('?') - : parameter.TypeFullName; - if (availableInterfaceFullNames.Contains(typeLookup)) continue; - if (parameter.HasExplicitDefault) continue; - if (parameter.IsParams) continue; - // Collection params (IEnumerable, T[], List, etc.) are always satisfiable - // via MakeArray – add to missing for factory generation but keep constructor valid - if (parameter.IsCollection) - { - localMissing.Add(parameter); - continue; - } - // Nullable reference/value params that aren't registered default to null - if (parameter.IsNullable) - { - localNullableDefaults.Add(parameter); - continue; - } - valid = false; - localMissing.Add(parameter); - } + HashSet? localMissing = null; + HashSet? localNullableDefaults = null; + AnalyzeParameters(ctor.Parameters, availableInterfaceFullNames, ref localMissing, ref localNullableDefaults, out var valid); if (valid) { chosen = ctor; - missing = localMissing.Count > 0 ? localMissing : null; - nullableDefaults = localNullableDefaults.Count > 0 ? localNullableDefaults : null; + missing = localMissing; + nullableDefaults = localNullableDefaults; break; } - if ((missing?.Count ?? int.MaxValue) <= localMissing.Count) continue; + if ((missing?.Count ?? int.MaxValue) <= (localMissing?.Count ?? 0)) continue; chosen = ctor; missing = localMissing; - nullableDefaults = localNullableDefaults.Count > 0 ? localNullableDefaults : null; + nullableDefaults = localNullableDefaults; } return chosen; } - private static IEnumerable GetBestConstructorMissing(InjectionData injection, + private static IEnumerable GetInjectionMissingParameters(InjectionData injection, ImmutableArray availableInterfaceFullNames) { + if (injection.Lambda is LambdaData lambda) + { + if (!lambda.IsMethod) + return Enumerable.Empty(); + + HashSet? lambdaMissing = null; + HashSet? lambdaNullableDefaults = null; + AnalyzeParameters(lambda.MethodParameters, availableInterfaceFullNames, ref lambdaMissing, ref lambdaNullableDefaults); + return lambdaMissing ?? Enumerable.Empty(); + } + HashSet? missing = null; HashSet? nullableDefaults = null; GetBestConstructor(injection, availableInterfaceFullNames, ref missing, ref nullableDefaults); return missing ?? Enumerable.Empty(); } - private static string MakeConstructorCall(ConstructorData ctor, HashSet? missing, HashSet? nullableDefaults) + private static void AnalyzeParameters( + ImmutableArray parameters, + ImmutableArray availableInterfaceFullNames, + ref HashSet? missing, + ref HashSet? nullableDefaults) { - var args = new List(); - foreach (var parameter in ctor.Parameters) + AnalyzeParameters(parameters, availableInterfaceFullNames, ref missing, ref nullableDefaults, out _); + } + + private static void AnalyzeParameters( + ImmutableArray parameters, + ImmutableArray availableInterfaceFullNames, + ref HashSet? missing, + ref HashSet? nullableDefaults, + out bool valid) + { + var localMissing = new HashSet(); + var localNullableDefaults = new HashSet(); + valid = true; + + foreach (var parameter in parameters) { - if (nullableDefaults?.Contains(parameter) == true) + var typeLookup = parameter.IsNullable + ? parameter.TypeFullName.TrimEnd('?') + : parameter.TypeFullName; + + if (parameter.IsCollection) { - args.Add("null"); + localMissing.Add(parameter); continue; } - if (missing?.Contains(parameter) == true) + + if (availableInterfaceFullNames.Contains(typeLookup)) + continue; + + if (parameter.HasExplicitDefault || parameter.IsParams) + continue; + + if (parameter.IsNullable) { - args.Add(CollectionConstructorArg(parameter)); + localNullableDefaults.Add(parameter); continue; } - args.Add(parameter.TypeMemberName + "()"); + + valid = false; + localMissing.Add(parameter); } - return $"({string.Join(", ", args)})"; + + missing = localMissing.Count > 0 ? localMissing : null; + nullableDefaults = localNullableDefaults.Count > 0 ? localNullableDefaults : null; + } + + private static string MakeConstructorCall(ConstructorData ctor, HashSet? missing, HashSet? nullableDefaults) + { + return MakeInvocationCall(ctor.Parameters, missing, nullableDefaults); } private static string MakeMethodCall(ImmutableArray parameters, HashSet? missing, HashSet? nullableDefaults = null) + { + return MakeInvocationCall(parameters, missing, nullableDefaults); + } + + private static string MakeInvocationCall( + ImmutableArray parameters, + HashSet? missing, + HashSet? nullableDefaults) { var args = new List(); + var useNamedArguments = false; foreach (var parameter in parameters) { if (nullableDefaults?.Contains(parameter) == true) { - args.Add("null"); + args.Add(useNamedArguments ? $"{parameter.Name}: null" : "null"); continue; } if (missing?.Contains(parameter) == true) { - args.Add(CollectionConstructorArg(parameter)); + var argument = CollectionConstructorArg(parameter); + args.Add(useNamedArguments ? $"{parameter.Name}: {argument}" : argument); continue; } - args.Add(parameter.TypeMemberName + "()"); + + if (parameter.HasExplicitDefault || parameter.IsParams) + { + useNamedArguments = true; + continue; + } + + var resolvedArgument = parameter.TypeMemberName + "()"; + args.Add(useNamedArguments ? $"{parameter.Name}: {resolvedArgument}" : resolvedArgument); } return $"({string.Join(", ", args)})"; } @@ -993,15 +1526,20 @@ private static string MakeMethodCall(ImmutableArray parameters, H /// in a generated constructor call. Collection params are converted from the cached /// IEnumerable<T> factory to the exact type requested. /// - private static string CollectionConstructorArg(ParameterData parameter) => - parameter.CollectionKind switch - { - CollectionKind.Array => $"{parameter.Name}.ToArray()", - CollectionKind.List => $"{parameter.Name}.ToList()", - CollectionKind.ImmutableArray => $"ImmutableArray.CreateRange({parameter.Name})", - CollectionKind.ReadOnlySpan => $"new global::System.ReadOnlySpan<{parameter.CollectionElementFullName}>({parameter.Name}.ToArray())", - _ => parameter.Name, // Enumerable or plain missing → use name directly + 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 @@ -1015,5 +1553,757 @@ private static string CollectionDictExpression(CollectionKind kind, string facto CollectionKind.ImmutableArray => $"ImmutableArray.CreateRange({factoryName})", _ => factoryName, // Enumerable → direct }; + + // ── C# 14 static-extension generation ──────────────────────────────────── + + private static bool IsAtLeastCSharp14(ParseOptions options, CancellationToken _) + { + if (options is not CSharpParseOptions csOptions) return false; + // C# 14 = 1400 in Roslyn's LanguageVersion enum. + // LanguageVersion.Preview == int.MaxValue, which is also >= 1400. + const int CSharp14 = 1400; + return (int)csOptions.LanguageVersion >= CSharp14; + } + + private static 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 2be0a66..a27e13a 100644 --- a/FactoryGenerator/FactoryGenerator.csproj +++ b/FactoryGenerator/FactoryGenerator.csproj @@ -12,9 +12,10 @@ - + + diff --git a/FactoryGenerator/Injection.cs b/FactoryGenerator/Injection.cs index 890bf13..f9120fa 100644 --- a/FactoryGenerator/Injection.cs +++ b/FactoryGenerator/Injection.cs @@ -41,6 +41,9 @@ public static class Injection } if (namedTypeSymbol is null) return null; + var assembly = symbol.ContainingAssembly ?? namedTypeSymbol.ContainingAssembly; + var assemblyName = assembly?.Name ?? string.Empty; + var assemblyPriority = GetAssemblyPriority(assembly); var singleInstance = false; var acquireChildInterfaces = false; @@ -94,10 +97,14 @@ public static class Injection interfaces = interfaces.Add(namedTypeSymbol); interfaces = interfaces.AddRange(attributedInterfaces); - var isDisposable = namedTypeSymbol.AllInterfaces.Any(i => i.Name.Equals("IDisposable")); - var disposableIface = interfaces.FirstOrDefault(i => i.Name.Contains("IDisposable")); + var isDisposable = namedTypeSymbol.AllInterfaces.Any(i => i.SpecialType == SpecialType.System_IDisposable); + var isAsyncDisposable = namedTypeSymbol.AllInterfaces.Any(i => i.ToString() == "System.IAsyncDisposable"); + var disposableIface = interfaces.FirstOrDefault(i => i.SpecialType == SpecialType.System_IDisposable); if (disposableIface is not null) interfaces = interfaces.Remove(disposableIface); + var asyncDisposableIface = interfaces.FirstOrDefault(i => i.ToString() == "System.IAsyncDisposable"); + if (asyncDisposableIface is not null) + interfaces = interfaces.Remove(asyncDisposableIface); interfaces = interfaces .RemoveRange(preventedInterfaces) @@ -121,12 +128,14 @@ public static class Injection return new InjectionData( typeFullName: namedTypeSymbol.ToString()!, typeMemberName: typeMemberName, - isTestType: namedTypeSymbol.ToString()!.Contains("Test"), + assemblyName: assemblyName, + assemblyPriority: assemblyPriority, interfaceFullNames: ifaceFullNames, interfaceMemberNames: ifaceMemberNames, singleton: singleInstance, scoped: scoped, disposable: isDisposable, + asyncDisposable: isAsyncDisposable, booleanInjection: boolean, constructors: constructors, lambda: lambdaData); @@ -168,5 +177,23 @@ private static ParameterData ExtractParameter(IParameterSymbol parameter) return new BooleanInjection(true, key); return null; } + + private static int GetAssemblyPriority(IAssemblySymbol? assemblySymbol) + { + if (assemblySymbol is null) + return 0; + + foreach (var attributeData in assemblySymbol.GetAttributes()) + { + if (attributeData.AttributeClass?.ToString() != "FactoryGenerator.Attributes.InjectionPriorityAttribute") + continue; + + if (attributeData.ConstructorArguments.Length == 1 + && attributeData.ConstructorArguments[0].Value is int priority) + return priority; + } + + return 0; + } } } diff --git a/FactoryGenerator/InjectionData.cs b/FactoryGenerator/InjectionData.cs index 2ad2cc1..d03ef48 100644 --- a/FactoryGenerator/InjectionData.cs +++ b/FactoryGenerator/InjectionData.cs @@ -1,5 +1,4 @@ using System; -using System.Collections.Generic; using System.Collections.Immutable; using System.Linq; @@ -11,12 +10,14 @@ public sealed class InjectionData : IEquatable { public string TypeFullName { get; } public string TypeMemberName { get; } // MemberName(type) without "()" - public bool IsTestType { get; } + public string AssemblyName { get; } + public int AssemblyPriority { get; } public ImmutableArray InterfaceFullNames { get; } public ImmutableArray InterfaceMemberNames { get; } // parallel to InterfaceFullNames, without "()" public bool Singleton { get; } public bool Scoped { get; } public bool Disposable { get; } + public bool AsyncDisposable { get; } public BooleanInjection? BooleanInjection { get; } public ImmutableArray Constructors { get; } public LambdaData? Lambda { get; } @@ -26,20 +27,22 @@ public sealed class InjectionData : IEquatable public string LazyFieldName => "m_" + TypeMemberName + (Lambda?.MemberName ?? string.Empty); public InjectionData( - string typeFullName, string typeMemberName, bool isTestType, + string typeFullName, string typeMemberName, string assemblyName, int assemblyPriority, ImmutableArray interfaceFullNames, ImmutableArray interfaceMemberNames, - bool singleton, bool scoped, bool disposable, + bool singleton, bool scoped, bool disposable, bool asyncDisposable, BooleanInjection? booleanInjection, ImmutableArray constructors, LambdaData? lambda) { TypeFullName = typeFullName; TypeMemberName = typeMemberName; - IsTestType = isTestType; + AssemblyName = assemblyName; + AssemblyPriority = assemblyPriority; InterfaceFullNames = interfaceFullNames; InterfaceMemberNames = interfaceMemberNames; Singleton = singleton; Scoped = scoped; Disposable = disposable; + AsyncDisposable = asyncDisposable; BooleanInjection = booleanInjection; Constructors = constructors; Lambda = lambda; @@ -51,12 +54,14 @@ public bool Equals(InjectionData? other) if (ReferenceEquals(this, other)) return true; return TypeFullName == other.TypeFullName && TypeMemberName == other.TypeMemberName - && IsTestType == other.IsTestType + && AssemblyName == other.AssemblyName + && AssemblyPriority == other.AssemblyPriority && InterfaceFullNames.SequenceEqual(other.InterfaceFullNames) && InterfaceMemberNames.SequenceEqual(other.InterfaceMemberNames) && Singleton == other.Singleton && Scoped == other.Scoped && Disposable == other.Disposable + && AsyncDisposable == other.AsyncDisposable && Equals(BooleanInjection, other.BooleanInjection) && Constructors.SequenceEqual(other.Constructors) && Equals(Lambda, other.Lambda); @@ -173,32 +178,4 @@ public bool Equals(LambdaData? other) public override int GetHashCode() => ContainingTypeFullName.GetHashCode(); } - public sealed class UsageData : IEquatable - { - public string FullName { get; } - public string MemberName { get; } // SymbolUtility.MemberName(type) without "()" - public string ElementTypeFullName { get; } - public string ElementTypeMemberName { get; } // without "()" - - public UsageData(string fullName, string memberName, string elementTypeFullName, string elementTypeMemberName) - { - FullName = fullName; - MemberName = memberName; - ElementTypeFullName = elementTypeFullName; - ElementTypeMemberName = elementTypeMemberName; - } - - public bool Equals(UsageData? other) - { - if (other is null) return false; - if (ReferenceEquals(this, other)) return true; - return FullName == other.FullName - && MemberName == other.MemberName - && ElementTypeFullName == other.ElementTypeFullName - && ElementTypeMemberName == other.ElementTypeMemberName; - } - - public override bool Equals(object? obj) => obj is UsageData other && Equals(other); - public override int GetHashCode() => FullName.GetHashCode(); - } } diff --git a/FactoryGenerator/Logger.cs b/FactoryGenerator/Logger.cs index 7c95cce..618c82f 100644 --- a/FactoryGenerator/Logger.cs +++ b/FactoryGenerator/Logger.cs @@ -1,5 +1,4 @@ using System; -using System.Diagnostics; using System.IO; namespace FactoryGenerator diff --git a/FactoryGenerator/SymbolUtility.cs b/FactoryGenerator/SymbolUtility.cs index f28762d..80ef8fb 100644 --- a/FactoryGenerator/SymbolUtility.cs +++ b/FactoryGenerator/SymbolUtility.cs @@ -1,4 +1,3 @@ -using System; using System.Collections.Generic; using System.Text; using Microsoft.CodeAnalysis; @@ -119,7 +118,7 @@ public static string SingletonFactory(string typeName, string name, string lazyN if (cached != null) return cached; var value = {creation}; - GetResolvedInstances().Add(new WeakReference(value)); + TrackResolvedInstance(value); {lazyName} = value; return value; }} @@ -151,7 +150,7 @@ internal static string DisposableFactory(string typeName, string name, string cr internal {typeName} {name} {{ var value = {creationCall}; - GetResolvedInstances().Add(new WeakReference(value)); + TrackResolvedInstance(value); return value; }}"; } diff --git a/FactoryGenerator/build/FactoryGenerator.props b/FactoryGenerator/build/FactoryGenerator.props new file mode 100644 index 0000000..5e84143 --- /dev/null +++ b/FactoryGenerator/build/FactoryGenerator.props @@ -0,0 +1,11 @@ + + + + true + + + + + + + diff --git a/README.md b/README.md index fe2fb0b..d7bdd82 100644 --- a/README.md +++ b/README.md @@ -11,6 +11,8 @@ with [Autofac](https://autofac.org/) beyond syntax choices. - **Attribute-based Generation:** Simply decorate your code with attributes like ```[Inject]```,```[Singleton]```,```[Self]``` and more and your IoC container will be woven together. - **Test-Overridability:** Need to swap out one injection for another to test something? Simply ```[Inject]``` a replacement inside your test project for a new container. +- **Static Extensions (C# 14+):** On .NET 10 and later, every registered interface gains a static ```Resolve``` method that inlines the full construction chain — no dictionary, no virtual dispatch. +- **Plugin Architecture:** Load AOT-compiled plugin assemblies at runtime and chain their containers together without reflection. ## Documentation @@ -101,6 +103,8 @@ public class Program ``` Of note is perhaps `Generated.DependencyInjectionContainer`, this is the Compile-time created implementation of our IoC container, it implements the interface `FactoryGenerator.IContainer`. +Generated containers also implement `IAsyncDisposable`. If you are working with an `ILifetimeScope` or `IContainer` reference, `await scope.DisposeAsync()` is the preferred path, but a synchronous `Dispose()` will also block until async-only services finish disposing. + ### Attributes | Attribute | Description | Requires | @@ -114,6 +118,10 @@ Of note is perhaps `Generated.DependencyInjectionContainer`, this is the Compile | ```Scoped``` | Ensures that this type will be resolved once per created scope, if you do not use IContainer.BeginLifetimeScope(), this behaves like a singleton | ```Inject``` | | ```Boolean(string key)``` | Creates a Runtime switch to decide whether this type should be the one that gets resolved
(or the best fitting fallback option, otherwise) | ```Inject``` | +`InjectionPriority(int)` is an **assembly-level** attribute rather than an injection attribute. + +If every implementation of a service is guarded by `[Boolean(...)]` and no ungated fallback exists, resolving that service throws `InvalidOperationException` when none of the booleans select an implementation. + ### Overriding Overriding Injections, i.e if in the graph above _Dependency C_ injected an ```ISomething``` instance and that specific implementation of ```ISomething``` will not work for anything that uses _Dependency B_, then _Dependency B_ can substitute that injection by providing it's own injection of ```ISomething```. This overriding generally follows the project dependency tree, so if _Project A_ depends on _Project B_ which depends on _Project C_, A can override both B and C, but B cannot override A. @@ -121,6 +129,18 @@ Overriding Injections, i.e if in the graph above _Dependency C_ injected an ```I **Note** Overriding Injections will not work if you resolve an ```IEnumerable```, as that will net your a collection of all ```ISomething``` that have been injected. +If you need to override the normal project-graph precedence, you can assign an assembly-level priority: + +```csharp +using FactoryGenerator.Attributes; + +[assembly: InjectionPriority(9)] +``` + +Higher priority values win over lower ones, and assemblies default to priority `0`. If two assemblies have the same priority, FactoryGenerator falls back to the normal dependency graph ordering, where the current project overrides its references and direct references override deeper transitive ones. Assemblies in the same graph tier are then ordered deterministically by assembly name. + +This assembly-level `InjectionPriority` only affects **which implementation wins during generated service resolution** when multiple assemblies provide the same service. It does **not** control plugin/container chaining order in `ContainerRegistry`. + ### Unprovided Values What happens if there are some constructor values needed by certain injected implementations, such as command line arguments, that cannot be known at compile time? @@ -167,6 +187,45 @@ public class Provider : IProvider ``` With this code, it is now possible to do `container.Resolve()`, which will effectively return the result of `new Provider().Method()`, although, since `Method` is `[Inject]`ed as a `[Singleton]`, the result will be cached and the same instance will be returned at every call to `Resolve` as well as shared between all Injected implementations that require a `IResultType`. +Injected method parameters follow the same rules as constructor parameters: unresolved external values are surfaced on the generated container constructor, optional parameters keep their declared defaults, and `params` collections are supplied from the container when possible. + +### Static Extensions (C# 14 / .NET 10+) + +When targeting C# 14 or later, FactoryGenerator automatically emits [static extension methods](https://learn.microsoft.com/en-us/dotnet/csharp/whats-new/csharp-14#extension-members) for every registered interface. This provides a dictionary-free, inline resolution path that the JIT can aggressively optimize. + +Instead of `container.Resolve()`, you can call: +```csharp +var singleton = ISingleton.Resolve(container); +var transient = IOverridable.Resolve(container); +var chain = ChainA.Resolve(container); +``` + +Each generated `Resolve` method inlines the full construction chain directly — no dictionary lookup, no factory-method indirection. Singletons use double-checked locking against the container's cache field, while transients emit a pure `new` expression. + +If a resolution graph depends on runtime booleans or external constructor values, the generated static method carries those inputs explicitly: +```csharp +var switched = ISwitchableInterface.Resolve(container, testBool: true); +var built = Constructed.Resolve(container, nonInjectedClassArgument: options); +``` + +**Null-container mode:** Passing `null` instead of a container instance bypasses the singleton/scoped cache entirely and performs a fresh allocation on every call, while still requiring any runtime inputs needed by the graph: +```csharp +// Fresh allocation every time — no singleton cache +var fresh = ISingleton.Resolve(null); +var switched = ISwitchableInterface.Resolve(testBool: true); +``` + +Collection dependencies (`IEnumerable`, arrays, `List`, `ImmutableArray`, `ReadOnlySpan`) are resolved through the same generated static pipeline, so the extension path and the normal container path construct equivalent object graphs. Direct `Resolve>()` calls are generated for every registered service type, even when no constructor or source usage referenced that collection shape ahead of time. If the current generated container has no local implementations for a collection element type, collection resolution still falls back to base containers and ASP.NET Core `IServiceProvider` sources. + +The static extensions are generated alongside the standard dictionary-based container and require no additional configuration. If the consuming project's language version is below C# 14, the extensions are simply not emitted. + +**Opting out:** If you are on C# 14+ but do not want the static extensions (for example, to reduce generated code size or avoid conflicts), set the following property in your `.csproj`: +```xml + + false + +``` + ### ASP.NET Core Integration For web applications, you can integrate FactoryGenerator with the standard `IServiceProvider`. @@ -201,6 +260,8 @@ app.Run(); This integration allows you to inject standard framework services (like `IConfiguration` or `ILogger`) into your `[Inject]`ed classes, and ensures that `[Scoped]` services are correctly disposed of at the end of each HTTP request. +If a request-scoped FactoryGenerator service only implements `IAsyncDisposable`, the middleware disposes it asynchronously at the end of the request. + ### Plugin / AOT Container Loading FactoryGenerator supports a plugin architecture where multiple assemblies each generate their own container, and these containers are chained together at runtime. This works seamlessly with AOT-compiled assemblies since no reflection is involved — discovery is entirely push-based via ```[ModuleInitializer]```. @@ -283,5 +344,7 @@ Alternatively, plugins can be registered with a priority for automatic ordering. ContainerRegistry.Register("MyPlugin", ContainerEntryPoint.Create, priority: 10); ``` +This `ContainerRegistry` priority is separate from `[assembly: InjectionPriority(...)]`: it only determines **where a plugin container is inserted in the runtime chain**, not which competing implementation inside a generated container wins for a service. + **Note** This system is fully AOT-compatible. No reflection is used for container discovery — ```[ModuleInitializer]``` methods run automatically when an assembly is loaded. Each generated ```ContainerEntryPoint.Create``` directly instantiates the concrete generated container without ```Activator.CreateInstance``` or type scanning. diff --git a/README_NUGET.md b/README_NUGET.md index 8dd5b04..08b748b 100644 --- a/README_NUGET.md +++ b/README_NUGET.md @@ -5,3 +5,5 @@ with [Autofac](https://autofac.org/) beyond syntax choices. - **Attribute-based Generation:** Simply decorate your code with attributes like ```[Inject]```,```[Singleton]```,```[Self]``` and more and your IoC container will be woven together. - **Test-Overridability:** Need to swap out one injection for another to test something? Simply ```[Inject]``` a replacement inside your test project for a new container. - **ASP.NET Core Integration:** Seamlessly integrate your source-generated container with the standard ASP.NET Core web pipeline. +- **Static Extensions (C# 14+):** On .NET 10 and later, every registered interface gains a static ```Resolve``` method that inlines the full construction chain — no dictionary, no virtual dispatch. +- **Plugin Architecture:** Load AOT-compiled plugin assemblies at runtime and chain their containers together without reflection. diff --git a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj index fc8a819..45c1d4a 100644 --- a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj +++ b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/FactoryGenerator.Extensions.AspNetCore.Tests.csproj @@ -1,26 +1,22 @@  - net8.0 + net10.0 false true - + - - - runtime; build; native; contentfiles; analyzers; buildtransitive - all - + diff --git a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs index a559d3c..84cc1f0 100644 --- a/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs +++ b/Tests/FactoryGenerator.Extensions.AspNetCore.Tests/IntegrationTests.cs @@ -1,4 +1,7 @@ -using System.Net; +using System; +using System.Collections.Generic; +using System.Net; +using System.Linq; using System.Threading.Tasks; using Microsoft.AspNetCore.Builder; using Microsoft.AspNetCore.Hosting; @@ -7,10 +10,7 @@ using Microsoft.Extensions.DependencyInjection; using Microsoft.Extensions.Hosting; using Shouldly; -using Xunit; -using FactoryGenerator; using FactoryGenerator.Attributes; -using FactoryGenerator.Extensions.AspNetCore; namespace FactoryGenerator.Extensions.AspNetCore.Tests; @@ -35,9 +35,65 @@ public class OtherService : IOtherService public string GetValue() => "Hello from IServiceProvider"; } +public interface IRequestScopedDependency +{ + Guid Id { get; } +} + +public class RequestScopedDependency : IRequestScopedDependency +{ + public Guid Id { get; } = Guid.NewGuid(); +} + +public interface IRequestScopedFactoryService +{ + Guid GetRequestId(); +} + +[Inject] +public class RequestScopedFactoryService(IRequestScopedDependency dependency) : IRequestScopedFactoryService +{ + public Guid GetRequestId() => dependency.Id; +} + +public interface IAsyncRequestScopedService +{ +} + +[Inject, Scoped] +public class AsyncRequestScopedService : IAsyncRequestScopedService, IAsyncDisposable +{ + public static int DisposeAsyncCount { get; private set; } + + public static void Reset() + { + DisposeAsyncCount = 0; + } + + public ValueTask DisposeAsync() + { + DisposeAsyncCount++; + return default; + } +} + +public interface IFrameworkOnlyCollectionItem +{ +} + +public sealed class FrameworkOnlyCollectionItem : IFrameworkOnlyCollectionItem +{ +} + +[Inject, Self] +public class FrameworkOnlyCollectionConsumer(IEnumerable items) +{ + public IEnumerable Items { get; } = items; +} + public class IntegrationTests { - [Fact] + [Test] public async Task Middleware_Integrates_FactoryGenerator_With_RequestServices() { // Setup @@ -49,6 +105,7 @@ public async Task Middleware_Integrates_FactoryGenerator_With_RequestServices() .ConfigureServices(services => { services.AddSingleton(); + services.AddScoped(); }) .Configure(app => { @@ -66,8 +123,10 @@ public async Task Middleware_Integrates_FactoryGenerator_With_RequestServices() { var myService = context.RequestServices.GetRequiredService(); var otherService = context.RequestServices.GetRequiredService(); + var requestScopedFactoryService = context.RequestServices.GetRequiredService(); + var requestScopedDependency = context.RequestServices.GetRequiredService(); - await context.Response.WriteAsync($"{myService.GetValue()} | {otherService.GetValue()}"); + await context.Response.WriteAsync($"{myService.GetValue()} | {otherService.GetValue()} | {requestScopedFactoryService.GetRequestId()} | {requestScopedDependency.Id}"); }); }); }) @@ -81,6 +140,128 @@ public async Task Middleware_Integrates_FactoryGenerator_With_RequestServices() // Assert response.StatusCode.ShouldBe(HttpStatusCode.OK); var content = await response.Content.ReadAsStringAsync(); - content.ShouldBe("Hello from FactoryGenerator | Hello from IServiceProvider"); + var parts = content.Split(" | "); + parts.Length.ShouldBe(4); + parts[0].ShouldBe("Hello from FactoryGenerator"); + parts[1].ShouldBe("Hello from IServiceProvider"); + parts[2].ShouldBe(parts[3]); + } + + [Test] + public async Task Middleware_Uses_Current_RequestScope_For_FrameworkScopedDependencies() + { + var host = await new HostBuilder() + .ConfigureWebHost(webBuilder => + { + webBuilder + .UseTestServer() + .ConfigureServices(services => + { + services.AddScoped(); + }) + .Configure(app => + { + var adapter = app.ApplicationServices.ToContainer(); + var container = new Generated.DependencyInjectionContainer(adapter); + + app.UseFactoryGenerator(container); + + app.Run(async context => + { + var requestScopedFactoryService = context.RequestServices.GetRequiredService(); + var requestScopedDependency = context.RequestServices.GetRequiredService(); + + await context.Response.WriteAsync($"{requestScopedFactoryService.GetRequestId()}|{requestScopedDependency.Id}"); + }); + }); + }) + .StartAsync(); + + var client = host.GetTestClient(); + + var firstResponse = await client.GetStringAsync("/"); + var secondResponse = await client.GetStringAsync("/"); + + var firstParts = firstResponse.Split('|'); + var secondParts = secondResponse.Split('|'); + + firstParts.Length.ShouldBe(2); + secondParts.Length.ShouldBe(2); + firstParts[0].ShouldBe(firstParts[1]); + secondParts[0].ShouldBe(secondParts[1]); + firstParts[0].ShouldNotBe(secondParts[0]); + } + + [Test] + public async Task Middleware_Disposes_AsyncOnly_FactoryScopedServices() + { + AsyncRequestScopedService.Reset(); + + var host = await new HostBuilder() + .ConfigureWebHost(webBuilder => + { + webBuilder + .UseTestServer() + .ConfigureServices(services => + { + services.AddScoped(); + }) + .Configure(app => + { + var adapter = app.ApplicationServices.ToContainer(); + var container = new Generated.DependencyInjectionContainer(adapter); + + app.UseFactoryGenerator(container); + + app.Run(context => + { + _ = context.RequestServices.GetRequiredService(); + return context.Response.WriteAsync("ok"); + }); + }); + }) + .StartAsync(); + + var client = host.GetTestClient(); + var response = await client.GetAsync("/"); + + response.StatusCode.ShouldBe(HttpStatusCode.OK); + AsyncRequestScopedService.DisposeAsyncCount.ShouldBe(1); + } + + [Test] + public async Task Middleware_Uses_Framework_Collections_When_No_Local_Implementations_Exist() + { + var host = await new HostBuilder() + .ConfigureWebHost(webBuilder => + { + webBuilder + .UseTestServer() + .ConfigureServices(services => + { + services.AddScoped(); + services.AddSingleton(); + services.AddSingleton(); + }) + .Configure(app => + { + var adapter = app.ApplicationServices.ToContainer(); + var container = new Generated.DependencyInjectionContainer(adapter); + + app.UseFactoryGenerator(container); + + app.Run(async context => + { + var consumer = context.RequestServices.GetRequiredService(); + await context.Response.WriteAsync(consumer.Items.Count().ToString()); + }); + }); + }) + .StartAsync(); + + var client = host.GetTestClient(); + var response = await client.GetStringAsync("/"); + + response.ShouldBe("2"); } } diff --git a/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs b/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs index 5e9b480..f08435d 100644 --- a/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs +++ b/Tests/FactoryGenerator.Tests/ContainerRegistryTests.cs @@ -1,3 +1,4 @@ +using System.Runtime.CompilerServices; using FactoryGenerator; using Inherited; using Inheritor.Generated; @@ -7,17 +8,19 @@ namespace FactoryGenerator.Tests; public class ContainerRegistryTests { - [Fact] + [Test] public void ContainerEntryPointRegistersOnModuleLoad() { - // The Inheritor assembly's ModuleInitializer should have already registered - // its container factory in ContainerRegistry when the assembly was loaded. + EnsureContainerEntryPointModuleInitialized(); + ContainerRegistry.RegisteredAssemblies.ShouldContain("Inheritor"); } - [Fact] + [Test] public void ContainerEntryPointCreateBuildsWorkingContainer() { + EnsureContainerEntryPointModuleInitialized(); + // Create a base container var baseContainer = new DependencyInjectionContainer(default, default, new NonInjectedClass()); @@ -28,9 +31,11 @@ public void ContainerEntryPointCreateBuildsWorkingContainer() chained.ShouldBeAssignableTo(); } - [Fact] + [Test] public void BuildChainCreatesWorkingContainerPipeline() { + EnsureContainerEntryPointModuleInitialized(); + // Create a base container var baseContainer = new DependencyInjectionContainer(default, default, new NonInjectedClass()); @@ -42,9 +47,38 @@ public void BuildChainCreatesWorkingContainerPipeline() final.Resolve().ShouldNotBeNull(); } - [Fact] + [Test] + public void BuildChainWithoutAssemblyListSkipsCurrentContainerAssembly() + { + EnsureContainerEntryPointModuleInitialized(); + var baseContainer = new DependencyInjectionContainer(default, default, new NonInjectedClass()); + + var final = ContainerRegistry.BuildChain(baseContainer); + + ReferenceEquals(final, baseContainer).ShouldBeTrue(); + final.Inheritor.ShouldBeNull(); + } + + [Test] + public void BuildChainWithExplicitAssemblyListSkipsCurrentContainerAssembly() + { + EnsureContainerEntryPointModuleInitialized(); + var baseContainer = new DependencyInjectionContainer(default, default, new NonInjectedClass()); + + var final = ContainerRegistry.BuildChain(baseContainer, new[] { "Inheritor" }); + + ReferenceEquals(final, baseContainer).ShouldBeTrue(); + final.Inheritor.ShouldBeNull(); + } + + [Test] public void ContainerEntryPointAssemblyNameIsCorrect() { ContainerEntryPoint.AssemblyName.ShouldBe("Inheritor"); } + + private static void EnsureContainerEntryPointModuleInitialized() + { + RuntimeHelpers.RunModuleConstructor(typeof(ContainerEntryPoint).Module.ModuleHandle); + } } diff --git a/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj b/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj index 1bb0dd5..40d2225 100644 --- a/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj +++ b/Tests/FactoryGenerator.Tests/FactoryGenerator.Tests.csproj @@ -1,7 +1,8 @@ - net9.0 + net10.0 + preview enable enable false @@ -9,19 +10,12 @@ - + - - - - runtime; build; native; contentfiles; analyzers; buildtransitive - all - - - + diff --git a/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs b/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs new file mode 100644 index 0000000..a641c79 --- /dev/null +++ b/Tests/FactoryGenerator.Tests/GeneratorBehaviorTests.cs @@ -0,0 +1,521 @@ +using System; +using System.IO; +using System.Linq; +using FactoryGenerator; +using FactoryGenerator.Attributes; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using Shouldly; + +namespace FactoryGenerator.Tests; + +public class GeneratorBehaviorTests +{ + [Test] + public void GeneratorRejectsMultipleExternalValuesOfSameType() + { + const string source = """ +using FactoryGenerator.Attributes; + +namespace Sample +{ +public interface IService +{ +} + +public class ExternalValue +{ +} + +[Inject] +public class FirstConsumer : IService +{ + public FirstConsumer(ExternalValue first) + { + } +} + +[Inject, Self] +public class SecondConsumer +{ + public SecondConsumer(ExternalValue second) + { + } +} +} +"""; + + var compilation = CreateCompilation(source); + var (runResult, _) = RunGenerator(compilation); + runResult.Results.Length.ShouldBe(1); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldNotBeNull(); + generatorResult.Exception!.Message.ShouldContain("Multiple externally provided values of the same type"); + generatorResult.Exception.Message.ShouldContain("Sample.ExternalValue"); + } + + [Test] + public void GeneratorSupportsBooleanKeysThatAreNotIdentifiers() + { + const string source = """ +using FactoryGenerator.Attributes; + +namespace Sample +{ +public interface IService +{ +} + +[Inject, Boolean("feature-flag")] +public class EnabledService : IService +{ +} + +[Inject] +public class FallbackService : IService +{ +} +} +"""; + + var compilation = CreateCompilation(source); + var (runResult, outputCompilation) = RunGenerator(compilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldBeNull(); + outputCompilation.GetDiagnostics() + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .ToArray() + .ShouldBeEmpty(); + + var generatedSource = string.Join(Environment.NewLine, generatorResult.GeneratedSources.Select(sourceResult => sourceResult.SourceText.ToString())); + generatedSource.ShouldContain("\"feature-flag\""); + generatedSource.ShouldNotContain("bool feature-flag"); + generatedSource.ShouldContain("Resolve(DependencyInjectionContainer? container, bool boolean_feature_flag)"); + } + + [Test] + public void BooleanOnlyImplementationsThrowInsteadOfResolvingNull() + { + var assemblyName = "BooleanOnly" + Guid.NewGuid().ToString("N"); + var source = $$""" +using FactoryGenerator.Attributes; + +namespace {{assemblyName}} +{ +public interface IService +{ +} + +[Inject, Boolean("enabled")] +public class EnabledService : IService +{ +} +} +"""; + + var compilation = CreateCompilation(assemblyName, source); + var (runResult, outputCompilation) = RunGenerator(compilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldBeNull(); + outputCompilation.GetDiagnostics() + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .ToArray() + .ShouldBeEmpty(); + + var expectedMessage = $"Cannot resolve {assemblyName}.IService without a matching implementation"; + var declarations = generatorResult.GeneratedSources + .Single(sourceResult => sourceResult.HintName == "DependencyInjectionContainer.Declarations.g.cs") + .SourceText + .ToString(); + declarations.ShouldContain(expectedMessage); + declarations.ShouldNotContain("null!"); + + var staticExtensions = generatorResult.GeneratedSources + .Single(sourceResult => sourceResult.HintName == "DependencyInjectionContainer.StaticExtensions.g.cs") + .SourceText + .ToString(); + staticExtensions.ShouldContain(expectedMessage); + + var assembly = System.Reflection.Assembly.Load(EmitAssembly(outputCompilation)); + var containerType = assembly.GetType($"{assemblyName}.Generated.DependencyInjectionContainer"); + var serviceType = assembly.GetType($"{assemblyName}.IService"); + + containerType.ShouldNotBeNull(); + serviceType.ShouldNotBeNull(); + + var container = (IContainer)Activator.CreateInstance(containerType!, new object[] { false })!; + var exception = Should.Throw(() => container.Resolve(serviceType!)); + exception.Message.ShouldContain(expectedMessage); + } + + [Test] + public void GeneratorDetectsCyclesThroughInjectedMethods() + { + const string source = """ +using FactoryGenerator.Attributes; + +namespace Sample +{ +public interface IResult +{ +} + +public class Result : IResult +{ +} + +public interface IFactory +{ + [Inject] + IResult Create(); +} + +[Inject] +public class Factory : IFactory +{ + public Factory(IResult result) + { + } + + public IResult Create() => new Result(); +} +} +"""; + + var compilation = CreateCompilation(source); + var (runResult, _) = RunGenerator(compilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldNotBeNull(); + generatorResult.Exception!.Message.ShouldContain("Cyclic Dependency Detected"); + generatorResult.Exception.Message.ShouldContain("Sample.IResult"); + generatorResult.Exception.Message.ShouldContain("Sample.IFactory"); + } + + [Test] + public void GeneratorDetectsCyclesThroughInjectedProperties() + { + const string source = """ +using FactoryGenerator.Attributes; + +namespace Sample +{ +public interface IResult +{ +} + +public class Result : IResult +{ +} + +public interface IFactory +{ + [Inject] + IResult Value { get; } +} + +[Inject] +public class Factory : IFactory +{ + public Factory(IResult result) + { + } + + public IResult Value => new Result(); +} +} +"""; + + var compilation = CreateCompilation(source); + var (runResult, _) = RunGenerator(compilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldNotBeNull(); + generatorResult.Exception!.Message.ShouldContain("Cyclic Dependency Detected"); + generatorResult.Exception.Message.ShouldContain("Sample.IResult"); + generatorResult.Exception.Message.ShouldContain("Sample.IFactory"); + } + + [Test] + public void InjectedMethodsSurfaceExternalParametersAndHonorOptionalAndParamsArguments() + { + var assemblyName = "InjectedMethod" + Guid.NewGuid().ToString("N"); + var source = $$""" +using FactoryGenerator.Attributes; + +namespace {{assemblyName}} +{ +public sealed class ExternalInput +{ + public ExternalInput(string name) + { + Name = name; + } + + public string Name { get; } +} + +public interface IPart +{ +} + +[Inject] +public class PartOne : IPart +{ +} + +[Inject] +public class PartTwo : IPart +{ +} + +public interface IResult +{ +} + +public sealed class Result : IResult +{ + public Result(string summary, int partCount) + { + Summary = summary; + PartCount = partCount; + } + + public string Summary { get; } + public int PartCount { get; } +} + +public interface IFactory +{ + [Inject] + IResult Create(ExternalInput input, string label = "default", params IPart[] parts); +} + +[Inject] +public class Factory : IFactory +{ + public IResult Create(ExternalInput input, string label = "default", params IPart[] parts) + { + return new Result(input.Name + ":" + label, parts.Length); + } +} +} +"""; + + var compilation = CreateCompilation(assemblyName, source); + var (runResult, outputCompilation) = RunGenerator(compilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldBeNull(); + outputCompilation.GetDiagnostics() + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .ToArray() + .ShouldBeEmpty(); + + var assembly = System.Reflection.Assembly.Load(EmitAssembly(outputCompilation)); + var containerType = assembly.GetType($"{assemblyName}.Generated.DependencyInjectionContainer"); + var externalType = assembly.GetType($"{assemblyName}.ExternalInput"); + var serviceType = assembly.GetType($"{assemblyName}.IResult"); + + containerType.ShouldNotBeNull(); + externalType.ShouldNotBeNull(); + serviceType.ShouldNotBeNull(); + + var constructor = containerType!.GetConstructors() + .Single(ctor => + { + var parameters = ctor.GetParameters(); + return parameters.Length == 1 && parameters[0].ParameterType == externalType; + }); + + var external = Activator.CreateInstance(externalType!, "runtime"); + var container = (IContainer)constructor.Invoke(new[] { external! }); + var resolved = container.Resolve(serviceType!); + + resolved.GetType().GetProperty("Summary")!.GetValue(resolved).ShouldBe("runtime:default"); + resolved.GetType().GetProperty("PartCount")!.GetValue(resolved).ShouldBe(2); + } + + [Test] + public void InjectedConstructorsHonorOptionalAndParamsArguments() + { + var assemblyName = "InjectedConstructor" + Guid.NewGuid().ToString("N"); + var source = $$""" +using FactoryGenerator.Attributes; + +namespace {{assemblyName}} +{ +public interface IPart +{ +} + +[Inject] +public class PartOne : IPart +{ +} + +[Inject] +public class PartTwo : IPart +{ +} + +[Inject, Self] +public sealed class Consumer +{ + public Consumer(string label = "default", params IPart[] parts) + { + Summary = label; + PartCount = parts.Length; + } + + public string Summary { get; } + public int PartCount { get; } +} +} +"""; + + var compilation = CreateCompilation(assemblyName, source); + var (runResult, outputCompilation) = RunGenerator(compilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldBeNull(); + outputCompilation.GetDiagnostics() + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .ToArray() + .ShouldBeEmpty(); + + var assembly = System.Reflection.Assembly.Load(EmitAssembly(outputCompilation)); + var containerType = assembly.GetType($"{assemblyName}.Generated.DependencyInjectionContainer"); + var consumerType = assembly.GetType($"{assemblyName}.Consumer"); + + containerType.ShouldNotBeNull(); + consumerType.ShouldNotBeNull(); + + var container = (IContainer)Activator.CreateInstance(containerType!)!; + var resolved = container.Resolve(consumerType!); + + resolved.GetType().GetProperty("Summary")!.GetValue(resolved).ShouldBe("default"); + resolved.GetType().GetProperty("PartCount")!.GetValue(resolved).ShouldBe(2); + } + + [Test] + public void AssemblyPriorityCanOverrideProjectGraphPrecedence() + { + var baseAssemblyName = "PriorityBase" + Guid.NewGuid().ToString("N"); + var derivedAssemblyName = "PriorityDerived" + Guid.NewGuid().ToString("N"); + + var baseSource = $$""" +using FactoryGenerator.Attributes; + +[assembly: InjectionPriority(9)] + +namespace {{baseAssemblyName}} +{ +public interface IService +{ +} + +[Inject] +public class BaseService : IService +{ +} +} +"""; + + var derivedSource = $$""" +using FactoryGenerator.Attributes; +using {{baseAssemblyName}}; + +namespace {{derivedAssemblyName}} +{ +[Inject] +public class DerivedService : IService +{ +} +} +"""; + + var baseCompilation = CreateCompilation(baseAssemblyName, baseSource); +var (baseReference, _) = EmitReference(baseCompilation); + var derivedCompilation = CreateCompilation(derivedAssemblyName, derivedSource, baseReference); + + var (runResult, outputCompilation) = RunGenerator(derivedCompilation); + var generatorResult = runResult.Results[0]; + + generatorResult.Exception.ShouldBeNull(); + outputCompilation.GetDiagnostics() + .Where(diagnostic => diagnostic.Severity == DiagnosticSeverity.Error) + .ToArray() + .ShouldBeEmpty(); + + var serviceMemberName = baseAssemblyName + "_IService()"; + var prioritizedImplementationMemberName = baseAssemblyName + "_BaseService()"; + var nonPrioritizedImplementationMemberName = derivedAssemblyName + "_DerivedService()"; + var generatedSource = string.Join(Environment.NewLine, generatorResult.GeneratedSources.Select(sourceResult => sourceResult.SourceText.ToString())); + generatedSource.ShouldContain($"internal {baseAssemblyName}.IService {serviceMemberName} => {prioritizedImplementationMemberName};"); + generatedSource.ShouldNotContain($"internal {baseAssemblyName}.IService {serviceMemberName} => {nonPrioritizedImplementationMemberName};"); + } + + private static CSharpCompilation CreateCompilation(string assemblyName, string source, params MetadataReference[] additionalReferences) + { + var syntaxTree = CSharpSyntaxTree.ParseText(source, new CSharpParseOptions(LanguageVersion.Preview)); + var excludedAssemblies = new[] + { + "Benchmarks", + "FactoryGenerator", + "FactoryGenerator.Attributes", + "FactoryGenerator.Extensions.AspNetCore", + "FactoryGenerator.Extensions.AspNetCore.Tests", + "FactoryGenerator.Tests", + "Inherited", + "Inheritor", + "TestWebApp" + }; + var references = ((string?)AppContext.GetData("TRUSTED_PLATFORM_ASSEMBLIES"))! + .Split(Path.PathSeparator) + .Where(path => !excludedAssemblies.Contains(Path.GetFileNameWithoutExtension(path), StringComparer.Ordinal)) + .Select(path => (MetadataReference)MetadataReference.CreateFromFile(path)) + .ToList(); + + references.Add(MetadataReference.CreateFromFile(typeof(InjectAttribute).Assembly.Location)); + references.AddRange(additionalReferences); + + return CSharpCompilation.Create( + assemblyName: assemblyName, + syntaxTrees: new[] { syntaxTree }, + references: references, + options: new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary)); + } + + private static CSharpCompilation CreateCompilation(string source) + { + return CreateCompilation("GeneratorBehaviorTests", source); + } + + private static (GeneratorDriverRunResult RunResult, Compilation OutputCompilation) RunGenerator(CSharpCompilation compilation) + { + var parseOptions = (CSharpParseOptions)compilation.SyntaxTrees.First().Options; + GeneratorDriver driver = CSharpGeneratorDriver.Create( + [new global::FactoryGenerator.FactoryGenerator().AsSourceGenerator()], + parseOptions: parseOptions); + driver = driver.RunGeneratorsAndUpdateCompilation(compilation, out var outputCompilation, out _); + return (driver.GetRunResult(), outputCompilation); + } + + private static (MetadataReference Reference, byte[] Image) EmitReference(CSharpCompilation compilation) + { + var image = EmitAssembly(compilation); + return (MetadataReference.CreateFromImage(image), image); + } + + private static byte[] EmitAssembly(Compilation compilation) + { + using var stream = new MemoryStream(); + var result = compilation.Emit(stream); + result.Success.ShouldBeTrue(string.Join(Environment.NewLine, result.Diagnostics)); + return stream.ToArray(); + } +} diff --git a/Tests/FactoryGenerator.Tests/GlobalUsings.cs b/Tests/FactoryGenerator.Tests/GlobalUsings.cs index 8c927eb..7aa3922 100644 --- a/Tests/FactoryGenerator.Tests/GlobalUsings.cs +++ b/Tests/FactoryGenerator.Tests/GlobalUsings.cs @@ -1 +1 @@ -global using Xunit; \ No newline at end of file +global using TUnit.Core; \ No newline at end of file diff --git a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs index 5cceb2e..4332eb4 100644 --- a/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs +++ b/Tests/FactoryGenerator.Tests/InjectionDetectionTests.cs @@ -1,8 +1,8 @@ -using System.ComponentModel; using Inherited; using Inheritor; using Inheritor.Generated; using Shouldly; +using System.Threading.Tasks; using Type = Inherited.Type; namespace FactoryGenerator.Tests; @@ -11,13 +11,16 @@ public class InjectionDetectionTests() { private readonly IContainer m_container = new DependencyInjectionContainer(default, default, new NonInjectedClass()); - [Fact] + [After(Test)] + public void DisposeContainer() => m_container.Dispose(); + + [Test] public void InjectedTypesAreResolvable() { m_container.Resolve().ShouldBeOfType(); } - [Fact] + [Test] public void SingletonInjectionsResolveToTheSameInstanceEverytime() { var first = m_container.Resolve(); @@ -25,7 +28,7 @@ public void SingletonInjectionsResolveToTheSameInstanceEverytime() ReferenceEquals(first, second).ShouldBeTrue(); } - [Fact] + [Test] public void NonSingleInjectionsResolveToDifferentInstanceEverytime() { var first = m_container.Resolve(); @@ -33,7 +36,7 @@ public void NonSingleInjectionsResolveToDifferentInstanceEverytime() ReferenceEquals(first, second).ShouldBeFalse(); } - [Fact] + [Test] public void ResolveUsesArguments() { var dummy = new NonInjectedClass(); @@ -41,22 +44,64 @@ public void ResolveUsesArguments() myContainer.Resolve().NonInjectedClassArgument.ShouldBe(dummy); } - [Theory] - [InlineData(true, typeof(EnabledImplementation))] - [InlineData(false, typeof(FallbackImplementation))] + [Test] + public void StaticExtensionsPropagateDirectExternalArguments() + { + var dummy = new NonInjectedClass(); + Constructed.Resolve(dummy).NonInjectedClassArgument.ShouldBe(dummy); + } + + [Test] + public void StaticExtensionsPropagateTransitiveExternalArguments() + { + var dummy = new NonInjectedClass(); + ConstructedConsumer.Resolve(dummy).Value.NonInjectedClassArgument.ShouldBe(dummy); + } + + [Test] + public void StaticExtensionsPropagateExternalArgumentsIntoCollections() + { + var dummy = new NonInjectedClass(); + ConstructedArrayConsumer.Resolve(dummy).Items.ShouldHaveSingleItem().NonInjectedClassArgument.ShouldBe(dummy); + } + + [Test] + [Arguments(true, typeof(EnabledImplementation))] + [Arguments(false, typeof(FallbackImplementation))] public void PickupSingleInjectionWithBoolean(bool value, System.Type expected) { var myContainer = new DependencyInjectionContainer(value, default, default!); myContainer.Resolve().ShouldBeOfType(expected); } - [Fact] + [Test] + public void StaticExtensionsPropagateDirectBooleanArguments() + { + ISwitchableInterface.Resolve(true).ShouldBeOfType(); + ISwitchableInterface.Resolve(false).ShouldBeOfType(); + } + + [Test] + public void StaticExtensionsPropagateTransitiveBooleanArguments() + { + BooleanConsumer.Resolve(true).Value.ShouldBeOfType(); + BooleanConsumer.Resolve(false).Value.ShouldBeOfType(); + } + + [Test] + public void StaticExtensionsPropagateBooleanArgumentsIntoCollections() + { + SwitchableArrayConsumer.Resolve(false).Items.Count().ShouldBe(1); + SwitchableArrayConsumer.Resolve(true).Items.Count().ShouldBe(2); + } + + [Test] public void PickupSingleInjectionFromMethod() { m_container.Resolve().ShouldBeOfType(); } - [Fact] + [Test] public void DoNotPickupNonInjection() { try @@ -71,7 +116,7 @@ public void DoNotPickupNonInjection() true.ShouldBeFalse(); } - [Fact] + [Test] public void DontPickupIDisposable() { try @@ -86,7 +131,34 @@ public void DontPickupIDisposable() true.ShouldBeFalse(); } - [Fact] + [Test] + public void DontPickupIAsyncDisposable() + { + try + { + m_container.Resolve(); + } + catch (Exception) + { + return; + } + + true.ShouldBeFalse(); + } + + [Test] + public void InterfacesContainingIDisposableInTheNameRemainResolvable() + { + m_container.Resolve().ShouldBeOfType(); + } + + [Test] + public void NonSystemIDisposableInterfacesRemainResolvable() + { + m_container.Resolve().ShouldBeOfType(); + } + + [Test] public void DontPickupExcluded() { try @@ -101,26 +173,38 @@ public void DontPickupExcluded() true.ShouldBeFalse(); } - [Fact] + [Test] public void PickupTypesSpecifiedByAs() { m_container.Resolve().ShouldBeOfType(); } - [Fact] + [Test] public void PickupInheritedInterfaces() { m_container.Resolve().ShouldBeOfType(); } - [Fact] + [Test] public void InheritorsOverride() { m_container.Resolve().ShouldBeOfType(); } + [Test] + public void OverrideImplementationsPreventFalsePositiveCycleDetection() + { + m_container.Resolve().ShouldBeOfType(); + } + + [Test] + public void BestConstructorsPreventFalsePositiveCycleDetection() + { + m_container.Resolve().ShouldBeOfType(); + } - [Fact] + + [Test] public void DisposingContainerDisposesSingletons() { ISingletonDisposer singleton; @@ -132,7 +216,7 @@ public void DisposingContainerDisposesSingletons() ((DisposableSingleton) singleton).WasDisposed.ShouldBeTrue(); } - [Fact] + [Test] public void DisposingLifetimeContainerDoesNotDisposeSingletons() { ISingletonDisposer singleton; @@ -150,7 +234,7 @@ public void DisposingLifetimeContainerDoesNotDisposeSingletons() ((DisposableSingleton) singleton).WasDisposed.ShouldBeTrue(); } - [Fact] + [Test] public void DisposingLifetimeContainerDisposesScoped() { IScoped singleton; @@ -164,7 +248,7 @@ public void DisposingLifetimeContainerDisposesScoped() singleton.WasDisposed.ShouldBeTrue(); } - [Fact] + [Test] public void DisposingContainerDoesNotDisposeUntrackedInstances() { IDisposer singleton; @@ -176,50 +260,108 @@ public void DisposingContainerDoesNotDisposeUntrackedInstances() ((DisposableNonSingleton) singleton).WasDisposed.ShouldBeTrue(); } - [Fact] + [Test] public void DisposingContainerDoesNotDisposesUnreferencedSingletons() { using var myContainer = new DependencyInjectionContainer(false, default, default!); } - [Fact] + [Test] + public void DisposingContainerSynchronouslyWaitsForAsyncOnlyServices() + { + var myContainer = new DependencyInjectionContainer(false, default, default!); + var singleton = myContainer.Resolve(); + + myContainer.Dispose(); + + singleton.ShouldBeOfType(); + singleton.WasDisposed.ShouldBeTrue(); + } + + [Test] + public async Task AsyncDisposingContainerDisposesAsyncSingletons() + { + IAsyncSingletonDisposer singleton; + var myContainer = new DependencyInjectionContainer(false, default, default!); + singleton = myContainer.Resolve(); + + await myContainer.DisposeAsync(); + + singleton.ShouldBeOfType(); + singleton.WasDisposed.ShouldBeTrue(); + } + + [Test] + public async Task AsyncDisposingLifetimeContainerDisposesAsyncScoped() + { + var myContainer = new DependencyInjectionContainer(false, default, default!); + var lifetime = myContainer.BeginLifetimeScope(); + var scoped = lifetime.Resolve(); + + await lifetime.DisposeAsync(); + + scoped.ShouldBeOfType(); + scoped.WasDisposed.ShouldBeTrue(); + myContainer.Dispose(); + } + + [Test] public void ArrayExpressionsCollect() { m_container.Resolve().Arrays.Count().ShouldBe(3); } - [Fact] + [Test] + public void StaticExtensionsResolveCollectionsInNullContainerMode() + { + ArrayConsumer.Resolve(null).Arrays.Count().ShouldBe(3); + } + + [Test] public void RequestedArraysArePresent() { Program.Method().Count().ShouldBe(3); } - [Fact] + [Test] + public void EnumerablesAreResolvableWithoutUsageSites() + { + m_container.Resolve>().Count().ShouldBe(2); + } + + [Test] + public void DuplicateRequestedArrayUsagesDoNotDuplicateLookupKeys() + { + Program.Method().Count().ShouldBe(3); + Program.MethodAgain().Count().ShouldBe(3); + } + + [Test] public void BooleanFallbackIsOverriden() { m_container.Resolve().ShouldBeOfType(); } - [Fact] + [Test] public void TryResolveWithTypeArgumentsWorks() { m_container.TryResolve(out var type).ShouldBeTrue(); type.ShouldBeOfType(); } - [Fact] + [Test] public void TryResolveWithTypeParameterWorks() { m_container.TryResolve(typeof(IType), out var type).ShouldBeTrue(); type.ShouldBeOfType(); } - [Fact] + [Test] public void ClassesInsideOtherClassesCanBeInjected() { m_container.Resolve(); } - [Fact] + [Test] public void ContainerMayCreateItself() { var newContainer = new DependencyInjectionContainer(m_container); @@ -227,20 +369,41 @@ public void ContainerMayCreateItself() resolved.Count().ShouldBe(6); var nonInjected = m_container.Resolve(); } - [Fact] + [Test] public void HierarchicalContainersResolveArraysProperly() { var newContainer = new DependencyInjectionContainer(m_container); newContainer.Resolve().Arrays.Count().ShouldBe(6); } - [Fact] + [Test] public void HierarchicalContainersResolveUsesFallBackIfItCannotFindImplementation() { var newContainer = new DependencyInjectionContainer(new DummyContainer()); newContainer.Resolve().ShouldBe(DummyContainer.DummyText); } - [Fact] + [Test] + public void HierarchicalContainersResolveCollectionsFromBaseWhenNoLocalImplementationExists() + { + var newContainer = new DependencyInjectionContainer(new DummyContainer()); + newContainer.Resolve().Items.Count().ShouldBe(2); + } + + [Test] + public void DirectCollectionResolveUsesBaseWhenNoLocalImplementationExists() + { + var newContainer = new DependencyInjectionContainer(new DummyContainer()); + newContainer.Resolve>().Count().ShouldBe(2); + } + + [Test] + public void StaticExtensionsUseContainerFallbackForCollectionsWithoutLocalImplementations() + { + var newContainer = new DependencyInjectionContainer(new DummyContainer()); + FallbackCollectionConsumer.Resolve(newContainer).Items.Count().ShouldBe(2); + } + + [Test] public void ContainerPropgatesRelevantBooleansCreateItself() { var baseContainer = new DependencyInjectionContainer(true, false, new()); @@ -252,7 +415,7 @@ public void ContainerPropgatesRelevantBooleansCreateItself() newContainer.GetBoolean("A").ShouldBeFalse(); newContainer.GetBoolean("TestBool").ShouldBeTrue(); } - [Fact] + [Test] public void HierarchicalContainersPropgatesBooleansUnknownToIt() { var newContainer = new DependencyInjectionContainer(new DummyContainer()); @@ -260,15 +423,46 @@ public void HierarchicalContainersPropgatesBooleansUnknownToIt() newContainer.GetBoolean("C").ShouldBe(false); } + [Test] + public void DisposingChildContainerDoesNotDisposeBaseContainer() + { + var baseContainer = new DependencyInjectionContainer(false, default, default!); + var singleton = baseContainer.Resolve().ShouldBeOfType(); + var child = new DependencyInjectionContainer(baseContainer); + + child.Dispose(); + + singleton.WasDisposed.ShouldBeFalse(); + baseContainer.Inheritor.ShouldBeNull(); + + baseContainer.Dispose(); + singleton.WasDisposed.ShouldBeTrue(); + } + + [Test] + public void DisposingChildContainerUnregistersItFromParentCollections() + { + var parent = new DependencyInjectionContainer(false, false, new NonInjectedClass()); + var child = new DependencyInjectionContainer(parent); + + parent.Resolve>().Count().ShouldBe(6); + + child.Dispose(); + + parent.Resolve>().Count().ShouldBe(3); + parent.Inheritor.ShouldBeNull(); + parent.Dispose(); + } + // ── Nullable parameter tests ────────────────────────────────────────────── - [Fact] + [Test] public void NullableUnregisteredParameterDefaultsToNull() { m_container.Resolve().Optional.ShouldBeNull(); } - [Fact] + [Test] public void NullableRegisteredParameterIsResolved() { m_container.Resolve().Optional.ShouldBeOfType(); @@ -276,32 +470,82 @@ public void NullableRegisteredParameterIsResolved() // ── Collection constructor parameter tests ──────────────────────────────── - [Fact] + [Test] public void ArrayConstructorParameterIsResolved() { m_container.Resolve().Arrays.Length.ShouldBe(3); } - [Fact] + [Test] public void ListConstructorParameterIsResolved() { m_container.Resolve().Arrays.Count.ShouldBe(3); } - [Fact] + [Test] public void ImmutableArrayConstructorParameterIsResolved() { m_container.Resolve().Arrays.Length.ShouldBe(3); } - [Fact] + [Test] public void ReadOnlySpanConstructorParameterIsResolved() { m_container.Resolve().Count.ShouldBe(3); } + + // ── Cross-array reentrancy tests ────────────────────────────────────────── + // Ensures that reentrancy guards are per-array-type, not global. Resolving + // IEnumerable triggers construction of CrossA3 which needs + // IEnumerable. That second resolution must not be blocked. + + [Test] + public void CrossArrayReentrancyResolvesAllA() + { + var items = m_container.Resolve().Items.ToList(); + items.Count.ShouldBe(3); + } + + [Test] + public void CrossArrayReentrancyResolvesBInsideCrossA3() + { + var items = m_container.Resolve().Items.ToList(); + var crossA3 = items.OfType().ShouldHaveSingleItem(); + crossA3.Deps.Count().ShouldBe(2); + } + + // ── Inheritor + Base array tests ────────────────────────────────────────── + + [Test] + public void InheritorAndBaseContainerMergeArrays() + { + var parent = new DependencyInjectionContainer(false, false, new NonInjectedClass()); + var child = new DependencyInjectionContainer(parent); + // Inherited defines SplitBase1 + SplitBase2 (2 items per container). + // Inheritor defines SplitInheritor1..3 (3 more per container). + // Each standalone container has 5. After merging, the child sees its own 5 + // plus the parent's 5 = 10. + child.Resolve().Items.Count().ShouldBe(10); + } + + [Test] + public void BaseContainerSeesInheritorArraysAfterLinking() + { + var parent = new DependencyInjectionContainer(false, false, new NonInjectedClass()); + var child = new DependencyInjectionContainer(parent); + // After linking, the parent's Inheritor is the child. Resolving on the + // parent should now include its own 5 plus the child's 5 = 10. + parent.Resolve().Items.Count().ShouldBe(10); + } + private class DummyContainer : IContainer { public const string DummyText = "I am a bit of text"; + private static readonly IFallbackCollectionItem[] s_fallbackCollectionItems = + [ + new DummyFallbackCollectionItem(), + new DummyFallbackCollectionItem() + ]; public static NonInjectedClass m_dummy = new(); public IContainer? Base => null; @@ -335,12 +579,14 @@ public bool IsRegistered() public T Resolve() { if (typeof(T) == typeof(string)) return (T) (object) DummyText; + if (typeof(T) == typeof(IEnumerable)) return (T) (object) s_fallbackCollectionItems; return (T) (object) m_dummy; } public object Resolve(System.Type type) { if (type == typeof(string)) return DummyText; + if (type == typeof(IEnumerable)) return s_fallbackCollectionItems; return m_dummy; } @@ -348,6 +594,7 @@ public bool TryResolve(System.Type type, out object? resolved) { resolved = null; if (type == typeof(string)) resolved = DummyText; + if (type == typeof(IEnumerable)) resolved = s_fallbackCollectionItems; return resolved != null; } @@ -355,6 +602,7 @@ public bool TryResolve(out T? resolved) { resolved = default; if (typeof(T) == typeof(string)) resolved = (T) (object) DummyText; + if (typeof(T) == typeof(IEnumerable)) resolved = (T) (object) s_fallbackCollectionItems; return resolved != null; } public IEnumerable<(string Key, bool Value)> GetBooleans() @@ -363,4 +611,6 @@ public bool TryResolve(out T? resolved) } } + + private sealed class DummyFallbackCollectionItem : IFallbackCollectionItem; } \ No newline at end of file diff --git a/Tests/FactoryGenerator.Tests/ResolvedInstanceTrackerTests.cs b/Tests/FactoryGenerator.Tests/ResolvedInstanceTrackerTests.cs new file mode 100644 index 0000000..a65ccab --- /dev/null +++ b/Tests/FactoryGenerator.Tests/ResolvedInstanceTrackerTests.cs @@ -0,0 +1,102 @@ +using System.Collections.Concurrent; +using System.Linq; +using System.Threading; +using System.Threading.Tasks; +using Shouldly; + +namespace FactoryGenerator.Tests; + +public class ResolvedInstanceTrackerTests +{ + [Test] + public void ConcurrentTrackingDuringDisposeDisposesEveryInstance() + { + const int count = 128; + var tracker = new ResolvedInstanceTracker(); + var instances = new ConcurrentBag(); + var start = new ManualResetEventSlim(false); + + var tasks = Enumerable.Range(0, count) + .Select(_ => Task.Run(() => + { + var instance = new SyncDisposableProbe(); + instances.Add(instance); + start.Wait(); + tracker.Track(instance); + })) + .ToArray(); + + start.Set(); + tracker.Dispose(); + Task.WhenAll(tasks).GetAwaiter().GetResult(); + + instances.Count.ShouldBe(count); + instances.All(instance => instance.WasDisposed).ShouldBeTrue(); + } + + [Test] + public void SynchronousDisposeWaitsForAsyncOnlyInstances() + { + var tracker = new ResolvedInstanceTracker(); + var instance = new AsyncDisposableProbe(); + tracker.Track(instance); + + tracker.Dispose(); + + instance.WasDisposed.ShouldBeTrue(); + } + + [Test] + public async Task AsynchronousDisposeDisposesAsyncOnlyInstances() + { + var tracker = new ResolvedInstanceTracker(); + var instance = new AsyncDisposableProbe(); + tracker.Track(instance); + + await tracker.DisposeAsync(); + + instance.WasDisposed.ShouldBeTrue(); + } + + [Test] + public void TrackingDoesNotKeepObjectsAlive() + { + var tracker = new ResolvedInstanceTracker(); + var weakReference = CreateTrackedWeakReference(tracker); + + GC.Collect(); + GC.WaitForPendingFinalizers(); + GC.Collect(); + + weakReference.TryGetTarget(out _).ShouldBeFalse(); + tracker.Dispose(); + } + + private static WeakReference CreateTrackedWeakReference(ResolvedInstanceTracker tracker) + { + var instance = new SyncDisposableProbe(); + tracker.Track(instance); + return new WeakReference(instance); + } + + private sealed class SyncDisposableProbe : IDisposable + { + public bool WasDisposed { get; private set; } + + public void Dispose() + { + WasDisposed = true; + } + } + + private sealed class AsyncDisposableProbe : IAsyncDisposable + { + public bool WasDisposed { get; private set; } + + public ValueTask DisposeAsync() + { + WasDisposed = true; + return default; + } + } +} diff --git a/Tests/TestData/Inherited/Inherited.csproj b/Tests/TestData/Inherited/Inherited.csproj index 039b8de..6ffb18e 100644 --- a/Tests/TestData/Inherited/Inherited.csproj +++ b/Tests/TestData/Inherited/Inherited.csproj @@ -5,7 +5,7 @@ - net9.0 + net10.0 enable enable diff --git a/Tests/TestData/Inherited/Types.cs b/Tests/TestData/Inherited/Types.cs index 70b1483..4729fea 100644 --- a/Tests/TestData/Inherited/Types.cs +++ b/Tests/TestData/Inherited/Types.cs @@ -1,6 +1,6 @@ using FactoryGenerator.Attributes; -using System.Collections.Generic; using System.Collections.Immutable; +using System.Threading.Tasks; namespace Inherited; @@ -8,12 +8,20 @@ public interface IType; public interface IOverridable; +public interface IOverrideCycle; + [Inject] public class Type : IType; [Inject] public class Overriden : IOverridable; +[Inject] +public class OverrideCycleBase(IOverrideCycle self) : IOverrideCycle +{ + public IOverrideCycle Self { get; } = self; +} + public interface ISingleton; [Inject, Singleton] @@ -44,8 +52,49 @@ public class Constructed(NonInjectedClass nonInjectedClassArgument, ISingleton i public ISingleton InjectedArgument { get; } = injectedArgument; } +[Inject, Self] +public class ConstructedConsumer(Constructed value) +{ + public Constructed Value { get; } = value; +} + +[Inject, Self] +public class ConstructedArrayConsumer(IEnumerable items) +{ + public IEnumerable Items { get; } = items; +} + +[Inject, Self] +public class BooleanConsumer(ISwitchableInterface value) +{ + public ISwitchableInterface Value { get; } = value; +} + +[Inject, Self] +public class SwitchableArrayConsumer(IEnumerable items) +{ + public IEnumerable Items { get; } = items; +} + public interface IMethodResult; +public interface IMultiConstructorCycle; + +public class ExternalOnlyDependency; + +[Inject] +public class MultiConstructorCycle : IMultiConstructorCycle +{ + public MultiConstructorCycle() + { + } + + public MultiConstructorCycle(IMultiConstructorCycle self, ExternalOnlyDependency externalOnlyDependency) + : this() + { + } +} + public class MethodResult : IMethodResult; public interface IMethodSource @@ -113,8 +162,37 @@ public class RequestedArray2 : IRequestedArray; [Inject] public class RequestedArray3 : IRequestedArray; +public interface IUnrequestedEnumerable; + +[Inject] +public class UnrequestedEnumerable1 : IUnrequestedEnumerable; + +[Inject] +public class UnrequestedEnumerable2 : IUnrequestedEnumerable; + +public interface IFallbackCollectionItem; + +[Inject, Self] +public class FallbackCollectionConsumer(IEnumerable items) +{ + public IEnumerable Items { get; } = items; +} + public interface IDisposer; +public interface INotIDisposable; + +[Inject] +public class NotDisposableNameMatch : INotIDisposable; + +public class CustomDisposableTypes +{ + public interface IDisposable; + + [Inject] + public class CustomDisposable : IDisposable; +} + [Inject] public class DisposableNonSingleton : IDisposer, IDisposable { @@ -139,6 +217,23 @@ public void Dispose() } } +public interface IAsyncSingletonDisposer +{ + bool WasDisposed { get; } +} + +[Inject, Singleton] +public class AsyncDisposableSingleton : IAsyncSingletonDisposer, IAsyncDisposable +{ + public bool WasDisposed { get; private set; } + + public ValueTask DisposeAsync() + { + WasDisposed = true; + return default; + } +} + public interface IOverrideBoolean; [Inject, Boolean("A")] @@ -158,6 +253,11 @@ public interface IScoped bool WasDisposed { get; } } +public interface IAsyncScoped +{ + bool WasDisposed { get; } +} + public interface ISelfish; public interface ISelfReferentialFactory @@ -191,6 +291,18 @@ public void Dispose() } } +[Inject, Scoped] +public class AsyncScoped : IAsyncScoped, IAsyncDisposable +{ + public bool WasDisposed { get; private set; } + + public ValueTask DisposeAsync() + { + WasDisposed = true; + return default; + } +} + // ── Nullable parameter tests ───────────────────────────────────────────────── /// Interface with no [Inject] implementation — intentionally unregistered. @@ -237,4 +349,62 @@ public class ImmutableArrayConsumer(ImmutableArray arrays) public class ReadOnlySpanConsumer(ReadOnlySpan arrays) { public int Count { get; } = arrays.Length; +} + +// ── Cross-array reentrancy tests ───────────────────────────────────────────── +// Resolving IEnumerable should work even though CrossA3 depends on +// IEnumerable. The reentrancy flag is per-collection type, so +// resolving the B array must not be blocked by the A array's reentrancy guard. + +public interface ICrossArrayB; + +[Inject] +public class CrossB1 : ICrossArrayB; + +[Inject] +public class CrossB2 : ICrossArrayB; + +public interface ICrossArrayA; + +[Inject] +public class CrossA1 : ICrossArrayA; + +[Inject] +public class CrossA2 : ICrossArrayA; + +/// +/// Implementation of ICrossArrayA that depends on an array of ICrossArrayB. +/// When the container builds IEnumerable<ICrossArrayA> and encounters CrossA3 +/// it must resolve IEnumerable<ICrossArrayB>. This must succeed because the +/// reentrancy guard is local to each array type. +/// +[Inject] +public class CrossA3(IEnumerable deps) : ICrossArrayA +{ + public IEnumerable Deps { get; } = deps; +} + +[Inject, Self] +public class CrossArrayConsumer(IEnumerable items) +{ + public IEnumerable Items { get; } = items; +} + +// ── Inheritor + Base array tests ───────────────────────────────────────────── +// Interface whose implementations are split across the Inherited and Inheritor +// projects, so we can verify that arrays merge correctly across container +// hierarchies (both Base → child and Inheritor → child directions). + +public interface ISplitArray; + +[Inject] +public class SplitBase1 : ISplitArray; + +[Inject] +public class SplitBase2 : ISplitArray; + +[Inject, Self] +public class SplitArrayConsumer(IEnumerable items) +{ + public IEnumerable Items { get; } = items; } \ No newline at end of file diff --git a/Tests/TestData/Inheritor/Inheritor.csproj b/Tests/TestData/Inheritor/Inheritor.csproj index 1b4aa7e..7e0c344 100644 --- a/Tests/TestData/Inheritor/Inheritor.csproj +++ b/Tests/TestData/Inheritor/Inheritor.csproj @@ -6,14 +6,16 @@ + - net9.0 + net10.0 true enable enable true + preview \ No newline at end of file diff --git a/Tests/TestData/Inheritor/Types.cs b/Tests/TestData/Inheritor/Types.cs index cb2ad0b..1163552 100644 --- a/Tests/TestData/Inheritor/Types.cs +++ b/Tests/TestData/Inheritor/Types.cs @@ -10,6 +10,9 @@ public class Overrider : IOverridable; [Inject] public class OverridingBoolean : IOverrideBoolean; +[Inject] +public class OverrideCycleResolved : IOverrideCycle; + [Inject] public class ChainA(ChainB B, ChainC C, ChainD D) { @@ -50,4 +53,24 @@ public static IEnumerable Method() var array = container.Resolve>(); return array; } -} \ No newline at end of file + + public static IEnumerable MethodAgain() + { + var container = new DependencyInjectionContainer(false, false, null!); + var array = container.Resolve>(); + return array; + } +} +// ── Inheritor + Base array tests ───────────────────────────────────────────── +// Additional ISplitArray implementations in the Inheritor project. When a child +// container is created from a parent, the merged IEnumerable should +// contain items from both Inherited (Base) and Inheritor. + +[Inject] +public class SplitInheritor1 : ISplitArray; + +[Inject] +public class SplitInheritor2 : ISplitArray; + +[Inject] +public class SplitInheritor3 : ISplitArray; \ No newline at end of file diff --git a/Tests/TestWebApp/TestWebApp.csproj b/Tests/TestWebApp/TestWebApp.csproj index be9b04a..09bfeaa 100644 --- a/Tests/TestWebApp/TestWebApp.csproj +++ b/Tests/TestWebApp/TestWebApp.csproj @@ -1,7 +1,7 @@  - net8.0 + net10.0 enable enable @@ -10,6 +10,7 @@ + diff --git a/global.json b/global.json new file mode 100644 index 0000000..e163e86 --- /dev/null +++ b/global.json @@ -0,0 +1,5 @@ +{ + "test": { + "runner": "Microsoft.Testing.Platform" + } +} \ No newline at end of file