diff --git a/CHANGELOG.md b/CHANGELOG.md index 3f2a178..837c95f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,3 +1,6 @@ +## 0.12.0 +- Add suffix to member name for interfaces with the same name from different namespaces to resolve conflict (https://github.com/pakrym/jab/pull/187). + ## 0.11.0 - Factory-created services are now correctly disposed (https://github.com/pakrym/jab/pull/179). diff --git a/README.md b/README.md index 223ce94..d341f8c 100644 --- a/README.md +++ b/README.md @@ -18,7 +18,7 @@ Jab provides a [C# Source Generator](https://devblogs.microsoft.com/dotnet/intro Add Jab package reference: ```xml - + ``` @@ -283,7 +283,7 @@ A minimal example ends up looking like this: } ], "dependencies": { - "com.pakrym.jab": "0.11.0", + "com.pakrym.jab": "0.12.0", ... } } diff --git a/src/Jab.Attributes/Jab.Attributes.csproj b/src/Jab.Attributes/Jab.Attributes.csproj index e733897..7c8cb52 100644 --- a/src/Jab.Attributes/Jab.Attributes.csproj +++ b/src/Jab.Attributes/Jab.Attributes.csproj @@ -6,7 +6,7 @@ preview Jab $(DefineConstants);JAB_ATTRIBUTES_PACKAGE;GENERIC_ATTRIBUTES - 0.11.0 + 0.12.0 $(ReleaseVersion) README.md MIT diff --git a/src/Jab.FunctionalTests.Common/ContainerTests.cs b/src/Jab.FunctionalTests.Common/ContainerTests.cs index b81b634..1655f3b 100644 --- a/src/Jab.FunctionalTests.Common/ContainerTests.cs +++ b/src/Jab.FunctionalTests.Common/ContainerTests.cs @@ -1,11 +1,12 @@ #nullable enable +using Jab; using System; using System.Collections.Generic; +using System.Linq; +using System.Reflection; using System.Threading.Tasks; using Xunit; -using Jab; - namespace JabTests { public partial class ContainerTests @@ -1448,6 +1449,36 @@ public void SupportsNoneTransientValueType() internal partial class SupportsNoneTransientValueTypeContainer { } + + + [Fact] + public void GeneratesDistinctMembersForSameLocalNameAcrossNamespaces() + { + var c = new NameCollisionContainer(); + + var foo1 = c.GetService(); + var foo2 = c.GetService(); + + Assert.NotNull(foo1); + Assert.NotNull(foo2); + Assert.NotSame(foo1, foo2); + + // Verify distinct private cache fields were generated (e.g. _IFoo_0, _IFoo_1) + var fields = typeof(NameCollisionContainer) + .GetFields(BindingFlags.Instance | BindingFlags.NonPublic) + .Where(f => f.Name.StartsWith("_IService", StringComparison.Ordinal)) + .Select(f => f.Name) + .Distinct() + .ToArray(); + + Assert.Equal(2, fields.Length); + Assert.NotEqual(fields[0], fields[1]); + } + + [ServiceProvider] + [Singleton(typeof(IService), typeof(ServiceImplementation))] + [Singleton(typeof(NestedNS.IService), typeof(NestedNS.ServiceImplementation))] + internal partial class NameCollisionContainer { } } } diff --git a/src/Jab.FunctionalTests.Common/Mocks/NestedNS/IService.cs b/src/Jab.FunctionalTests.Common/Mocks/NestedNS/IService.cs new file mode 100644 index 0000000..da1d729 --- /dev/null +++ b/src/Jab.FunctionalTests.Common/Mocks/NestedNS/IService.cs @@ -0,0 +1,12 @@ +using Jab; + +namespace JabTests.NestedNS +{ + internal interface IService + { + } + internal interface IService + { + T InnerService { get; } + } +} \ No newline at end of file diff --git a/src/Jab.FunctionalTests.Common/Mocks/NestedNS/ServiceImplementation.cs b/src/Jab.FunctionalTests.Common/Mocks/NestedNS/ServiceImplementation.cs new file mode 100644 index 0000000..6831c6a --- /dev/null +++ b/src/Jab.FunctionalTests.Common/Mocks/NestedNS/ServiceImplementation.cs @@ -0,0 +1,18 @@ +using Jab; + +namespace JabTests.NestedNS +{ + internal class ServiceImplementation : IService + { + } + + internal class ServiceImplementation : IService + { + public ServiceImplementation(T innerService) + { + InnerService = innerService; + } + + public T InnerService { get; } + } +} \ No newline at end of file diff --git a/src/Jab.Unity/package.json b/src/Jab.Unity/package.json index 5c66ef9..0698cc9 100644 --- a/src/Jab.Unity/package.json +++ b/src/Jab.Unity/package.json @@ -1,6 +1,6 @@ { "name": "com.pakrym.jab", - "version": "0.11.0", + "version": "0.12.0", "description": "C# Source Generator based dependency injection container implementation.", "licensesUrl": "https://github.com/pakrym/jab/blob/main/LICENSE", "repository": { diff --git a/src/Jab/ContainerGenerator.cs b/src/Jab/ContainerGenerator.cs index 5f503f0..819db06 100644 --- a/src/Jab/ContainerGenerator.cs +++ b/src/Jab/ContainerGenerator.cs @@ -1,5 +1,7 @@ namespace Jab; +using System.Linq; + [Generator] #pragma warning disable RS1001 // We don't want this to be discovered as analyzer but it simplifies testing public partial class ContainerGenerator : DiagnosticAnalyzer @@ -8,7 +10,12 @@ public partial class ContainerGenerator : DiagnosticAnalyzer /// Code for a [GeneratedCode] attribute to put on the top-level generated members. private static readonly string _generatedCodeAttribute = $"[global::System.CodeDom.Compiler.GeneratedCodeAttribute(\"{typeof(ContainerGenerator).Assembly.GetName().Name}\", \"{typeof(ContainerGenerator).Assembly.GetName().Version}\")]"; - private void GenerateCallSiteWithCache(CodeWriter codeWriter, string rootReference, ServiceCallSite serviceCallSite, Action valueCallback) + private void GenerateCallSiteWithCache( + CodeWriter codeWriter, + string rootReference, + ServiceCallSite serviceCallSite, + Action valueCallback, + Dictionary typeBaseNameMap) { if (serviceCallSite is ErrorCallSite errorCallSite) { @@ -23,7 +30,7 @@ private void GenerateCallSiteWithCache(CodeWriter codeWriter, string rootReferen if (serviceCallSite.Lifetime != ServiceLifetime.Transient) { - var cacheLocation = GetCacheLocation(serviceCallSite.Identity); + var cacheLocation = GetCacheLocation(serviceCallSite.Identity, typeBaseNameMap); codeWriter.Line($"if ({cacheLocation} == null)"); codeWriter.Line($"lock (this)"); using (codeWriter.Scope($"if ({cacheLocation} == null)")) @@ -35,7 +42,8 @@ private void GenerateCallSiteWithCache(CodeWriter codeWriter, string rootReferen (w, v) => { w.Line($"{cacheLocation} = {v};"); - }); + }, + typeBaseNameMap); } if (serviceCallSite.ImplementationType.IsValueType) @@ -52,17 +60,17 @@ private void GenerateCallSiteWithCache(CodeWriter codeWriter, string rootReferen GenerateCallSite(codeWriter, rootReference, serviceCallSite, (w, v) => { w.Line($"{serviceCallSite.ImplementationType} service = {v};"); - }); + }, typeBaseNameMap); codeWriter.Line($"TryAddDisposable(service);"); valueCallback(codeWriter, w => w.Append($"service")); } else { - GenerateCallSite(codeWriter, rootReference, serviceCallSite, valueCallback); + GenerateCallSite(codeWriter, rootReference, serviceCallSite, valueCallback, typeBaseNameMap); } } - private void WriteResolutionCall(CodeWriter codeWriter, ServiceIdentity other, string reference) + private void WriteResolutionCall(CodeWriter codeWriter, ServiceIdentity other, string reference, Dictionary typeBaseNameMap) { if (other.IsMainImplementation) { @@ -70,7 +78,7 @@ private void WriteResolutionCall(CodeWriter codeWriter, ServiceIdentity other, s } else { - codeWriter.Append($"{reference}.{GetResolutionServiceName(other)}()"); + codeWriter.Append($"{reference}.{GetResolutionServiceName(other, typeBaseNameMap)}()"); } } @@ -117,24 +125,33 @@ private static void AppendMemberGenericParameters(CodeWriter codeWriter, ISymbol } } - private void AppendParameters(CodeWriter codeWriter, ServiceCallSite[] parameters, KeyValuePair[] optionalParameters) + private void AppendParameters( + CodeWriter codeWriter, + ServiceCallSite[] parameters, + KeyValuePair[] optionalParameters, + Dictionary typeBaseNameMap) { foreach (var parameter in parameters) { - WriteResolutionCall(codeWriter, parameter.Identity, "this"); + WriteResolutionCall(codeWriter, parameter.Identity, "this", typeBaseNameMap); codeWriter.AppendRaw(", "); } foreach (var pair in optionalParameters) { codeWriter.Append($"{pair.Key.Name}: "); - WriteResolutionCall(codeWriter, pair.Value.Identity, "this"); + WriteResolutionCall(codeWriter, pair.Value.Identity, "this", typeBaseNameMap); codeWriter.AppendRaw(", "); } codeWriter.RemoveTrailingComma(); } - private void GenerateCallSite(CodeWriter codeWriter, string rootReference, ServiceCallSite serviceCallSite, Action valueCallback) + private void GenerateCallSite( + CodeWriter codeWriter, + string rootReference, + ServiceCallSite serviceCallSite, + Action valueCallback, + Dictionary typeBaseNameMap) { switch (serviceCallSite) { @@ -142,7 +159,7 @@ private void GenerateCallSite(CodeWriter codeWriter, string rootReference, Servi valueCallback(codeWriter, w => { w.Append($"new {transientCallSite.ImplementationType}("); - AppendParameters(w, transientCallSite.Parameters, transientCallSite.OptionalParameters); + AppendParameters(w, transientCallSite.Parameters, transientCallSite.OptionalParameters, typeBaseNameMap); w.Append($")"); }); break; @@ -159,7 +176,7 @@ private void GenerateCallSite(CodeWriter codeWriter, string rootReference, Servi AppendMemberGenericParameters(w, methodCallSite.Member); w.AppendRaw("("); - AppendParameters(w, methodCallSite.Parameters, methodCallSite.OptionalParameters); + AppendParameters(w, methodCallSite.Parameters, methodCallSite.OptionalParameters, typeBaseNameMap); w.Append($")"); }); break; @@ -170,7 +187,7 @@ private void GenerateCallSite(CodeWriter codeWriter, string rootReference, Servi { foreach (var item in arrayServiceCallSite.Items) { - WriteResolutionCall(codeWriter, item.Identity, "this"); + WriteResolutionCall(codeWriter, item.Identity, "this", typeBaseNameMap); w.LineRaw(", "); } } @@ -192,6 +209,9 @@ private void Execute(GeneratorContext context) { var roots = new ServiceProviderBuilder(context).BuildRoots(); + // Pre-compute unique base names for all service types across all roots + var typeBaseNameMap = BuildTypeBaseNameMap(roots); + foreach (var root in roots) { var codeWriter = new CodeWriter(); @@ -215,7 +235,7 @@ private void Execute(GeneratorContext context) using (codeWriter.Scope()) { codeWriter.Line($"private Scope? _rootScope;"); - WriteCacheLocations(root, codeWriter, isScope: false); + WriteCacheLocations(root, codeWriter, isScope: false, typeBaseNameMap); foreach (var rootService in root.RootCallSites) { @@ -226,7 +246,7 @@ private void Execute(GeneratorContext context) } else { - codeWriter.Append($"private {rootServiceType} {GetResolutionServiceName(rootService.Identity)}()"); + codeWriter.Append($"private {rootServiceType} {GetResolutionServiceName(rootService.Identity, typeBaseNameMap)}()"); } if (rootService.Lifetime == ServiceLifetime.Scoped) @@ -241,16 +261,17 @@ private void Execute(GeneratorContext context) GenerateCallSiteWithCache(codeWriter, "this", rootService, - (w, v) => w.Line($"return {v};")); + (w, v) => w.Line($"return {v};"), + typeBaseNameMap); } } codeWriter.Line(); } - WriteNamedServiceProvider(codeWriter, root); - WriteServiceProvider(codeWriter, root); - WriteDispose(codeWriter, root, isScoped: false); + WriteNamedServiceProvider(codeWriter, root, typeBaseNameMap); + WriteServiceProvider(codeWriter, root, typeBaseNameMap); + WriteDispose(codeWriter, root, isScoped: false, typeBaseNameMap); WritePublicGetServiceMethods(codeWriter); codeWriter.Line($"public Scope CreateScope() => new Scope(this);"); @@ -289,7 +310,7 @@ private void Execute(GeneratorContext context) WriteInterfaces(codeWriter, root, true); using (codeWriter.Scope()) { - WriteCacheLocations(root, codeWriter, isScope: true); + WriteCacheLocations(root, codeWriter, isScope: true, typeBaseNameMap); codeWriter.Line($"private {root.Type} _root;"); codeWriter.Line(); @@ -307,12 +328,12 @@ private void Execute(GeneratorContext context) using (rootService.Identity.IsMainImplementation ? codeWriter.Scope($"{rootServiceType} IServiceProvider<{rootServiceType}>.GetService()") : - codeWriter.Scope($"private {rootServiceType} {GetResolutionServiceName(rootService.Identity)}()")) + codeWriter.Scope($"private {rootServiceType} {GetResolutionServiceName(rootService.Identity, typeBaseNameMap)}()")) { if (rootService.Lifetime == ServiceLifetime.Singleton) { codeWriter.Append($"return "); - WriteResolutionCall(codeWriter, rootService.Identity, "_root"); + WriteResolutionCall(codeWriter, rootService.Identity, "_root", typeBaseNameMap); codeWriter.Line($";"); } else @@ -320,21 +341,22 @@ private void Execute(GeneratorContext context) GenerateCallSiteWithCache(codeWriter, "_root", rootService, - (w, v) => w.Line($"return {v};")); + (w, v) => w.Line($"return {v};"), + typeBaseNameMap); } } codeWriter.Line(); } - WriteServiceProvider(codeWriter, root); - WriteNamedServiceProvider(codeWriter, root); + WriteServiceProvider(codeWriter, root, typeBaseNameMap); + WriteNamedServiceProvider(codeWriter, root, typeBaseNameMap); if (root.KnownTypes.IServiceScopeType != null) { codeWriter.Line($"{root.KnownTypes.IServiceProviderType} {root.KnownTypes.IServiceScopeType}.ServiceProvider => this;"); codeWriter.Line(); } - WriteDispose(codeWriter, root, isScoped: true); + WriteDispose(codeWriter, root, isScoped: true, typeBaseNameMap); } using (codeWriter.Scope($"private Scope GetRootScope()")) @@ -358,13 +380,90 @@ private void Execute(GeneratorContext context) } } + private Dictionary BuildTypeBaseNameMap(IEnumerable roots) + { + var allTypes = new HashSet(SymbolEqualityComparer.Default); + + foreach (var r in roots) + { + foreach (var cs in r.RootCallSites) + { + if (cs.Identity.Type is INamedTypeSymbol nts) + { + allTypes.Add(nts); + } + } + } + + // Group by raw base name + var groups = new Dictionary>(StringComparer.Ordinal); + foreach (var t in allTypes) + { + var baseName = BuildRawBaseName(t); + if (!groups.TryGetValue(baseName, out var list)) + { + list = new List(); + groups[baseName] = list; + } + list.Add(t); + } + + var result = new Dictionary(SymbolEqualityComparer.Default); + + foreach (var kvp in groups) + { + var baseName = kvp.Key; + var list = kvp.Value; + + if (list.Count == 1) + { + result[list[0]] = baseName; + } + else + { + var ordered = list + .OrderBy(t => t.ToDisplayString(SymbolDisplayFormat.FullyQualifiedFormat), StringComparer.Ordinal) + .ToList(); + + for (int i = 0; i < ordered.Count; i++) + { + result[ordered[i]] = baseName + "_" + i; + } + } + } + + return result; + } + + private static string BuildRawBaseName(INamedTypeSymbol symbol) + { + var sb = new StringBuilder(); + + void Traverse(ITypeSymbol s) + { + sb.Append(s.Name); + if (s is INamedTypeSymbol { IsGenericType: true } g) + { + sb.Append("_"); + foreach (var ta in g.TypeArguments) + { + Traverse(ta); + } + } + } + + Traverse(symbol); + return sb.ToString(); + } + private IEnumerable> GroupNamedServices(ServiceProvider root) { return root.RootCallSites .Where(static s => s.Identity.IsMainNamedImplementation) .GroupBy(static s => s.Identity.Type, SymbolEqualityComparer.Default); } - private void WriteNamedServiceProvider(CodeWriter codeWriter, ServiceProvider root) + + private void WriteNamedServiceProvider(CodeWriter codeWriter, ServiceProvider root, Dictionary typeBaseNameMap) { foreach (var serviceGroup in GroupNamedServices(root)) { @@ -376,7 +475,7 @@ private void WriteNamedServiceProvider(CodeWriter codeWriter, ServiceProvider ro foreach (var callSite in serviceGroup) { codeWriter.Append($"case \"{callSite.Identity.Name}\": return "); - WriteResolutionCall(codeWriter, callSite.Identity, "this"); + WriteResolutionCall(codeWriter, callSite.Identity, "this", typeBaseNameMap); codeWriter.Line($";"); } @@ -387,7 +486,7 @@ private void WriteNamedServiceProvider(CodeWriter codeWriter, ServiceProvider ro } } - private void WriteServiceProvider(CodeWriter codeWriter, ServiceProvider root) + private void WriteServiceProvider(CodeWriter codeWriter, ServiceProvider root, Dictionary typeBaseNameMap) { using (codeWriter.Scope($"{typeof(object)}? {typeof(IServiceProvider)}.GetService({typeof(Type)} type)")) { @@ -396,7 +495,7 @@ private void WriteServiceProvider(CodeWriter codeWriter, ServiceProvider root) if (rootRootCallSite.Identity.IsMainImplementation) { codeWriter.Append($"if (type == typeof({rootRootCallSite.Identity.Type})) return "); - WriteResolutionCall(codeWriter, rootRootCallSite.Identity, "this"); + WriteResolutionCall(codeWriter, rootRootCallSite.Identity, "this", typeBaseNameMap); codeWriter.Line($";"); } } @@ -406,11 +505,10 @@ private void WriteServiceProvider(CodeWriter codeWriter, ServiceProvider root) codeWriter.Line(); - WriteKeyedServiceProvider(codeWriter, root); + WriteKeyedServiceProvider(codeWriter, root, typeBaseNameMap); } - - private void WriteKeyedServiceProvider(CodeWriter codeWriter, ServiceProvider root) + private void WriteKeyedServiceProvider(CodeWriter codeWriter, ServiceProvider root, Dictionary typeBaseNameMap) { var iface = root.KnownTypes.IKeyedServiceProviderType; if (iface == null) @@ -430,7 +528,7 @@ private void WriteKeyedServiceProvider(CodeWriter codeWriter, ServiceProvider ro foreach (var callSite in serviceGroup) { codeWriter.Append($"case \"{callSite.Identity.Name}\": return "); - WriteResolutionCall(codeWriter, callSite.Identity, "this"); + WriteResolutionCall(codeWriter, callSite.Identity, "this", typeBaseNameMap); codeWriter.Line($";"); } } @@ -459,7 +557,7 @@ private void WritePublicGetServiceMethods(CodeWriter codeWriter) codeWriter.Line(); } - private void WriteDispose(CodeWriter codeWriter, ServiceProvider root, bool isScoped) + private void WriteDispose(CodeWriter codeWriter, ServiceProvider root, bool isScoped, Dictionary typeBaseNameMap) { codeWriter.Line($"private {typeof(List)}? _disposables;"); codeWriter.Line(); @@ -491,7 +589,7 @@ private void WriteDispose(CodeWriter codeWriter, ServiceProvider root, bool isSc (rootService.Lifetime == ServiceLifetime.Scoped && !isScoped) || rootService.Lifetime == ServiceLifetime.Transient) continue; - codeWriter.Line($"TryDispose({GetCacheLocation(rootService.Identity)});"); + codeWriter.Line($"TryDispose({GetCacheLocation(rootService.Identity, typeBaseNameMap)});"); } if (!isScoped) @@ -533,7 +631,7 @@ private void WriteDispose(CodeWriter codeWriter, ServiceProvider root, bool isSc (rootService.Lifetime == ServiceLifetime.Scoped && !isScoped) || rootService.Lifetime == ServiceLifetime.Transient) continue; - codeWriter.Line($"await TryDispose({GetCacheLocation(rootService.Identity)});"); + codeWriter.Line($"await TryDispose({GetCacheLocation(rootService.Identity, typeBaseNameMap)});"); } if (!isScoped) @@ -549,7 +647,6 @@ private void WriteDispose(CodeWriter codeWriter, ServiceProvider root, bool isSc } } - codeWriter.Line(); } @@ -608,7 +705,7 @@ private static void WriteInterfaces(CodeWriter codeWriter, ServiceProvider root, codeWriter.Line(); } - private void WriteCacheLocations(ServiceProvider root, CodeWriter codeWriter, bool isScope) + private void WriteCacheLocations(ServiceProvider root, CodeWriter codeWriter, bool isScope, Dictionary typeBaseNameMap) { foreach (var rootService in root.RootCallSites) { @@ -616,44 +713,41 @@ private void WriteCacheLocations(ServiceProvider root, CodeWriter codeWriter, bo (rootService.Lifetime == ServiceLifetime.Scoped && !isScope) || rootService.Lifetime == ServiceLifetime.Transient) continue; - codeWriter.Line($"private {rootService.ImplementationType}? {GetCacheLocation(rootService.Identity)};"); + codeWriter.Line($"private {rootService.ImplementationType}? {GetCacheLocation(rootService.Identity, typeBaseNameMap)};"); } codeWriter.Line(); } - private string GetResolutionServiceName(ServiceIdentity identity) + private string GetResolutionServiceName(ServiceIdentity identity, Dictionary typeBaseNameMap) { if (!identity.IsMainImplementation) { - return $"Get{GetServiceExpandedName(identity)}"; + return $"Get{GetServiceExpandedName(identity, typeBaseNameMap)}"; } throw new InvalidOperationException("Main implementation should be resolved via GetService call"); } - private string GetCacheLocation(ServiceIdentity identity) + private string GetCacheLocation(ServiceIdentity identity, Dictionary typeBaseNameMap) { - return $"_{GetServiceExpandedName(identity)}"; + return $"_{GetServiceExpandedName(identity, typeBaseNameMap)}"; } - private string GetServiceExpandedName(ServiceIdentity identity) + private string GetServiceExpandedName(ServiceIdentity identity, Dictionary typeBaseNameMap) { - StringBuilder builder = new(); + var typeSymbol = (INamedTypeSymbol)identity.Type; + string baseName; - void Traverse(ITypeSymbol symbol) + if (typeBaseNameMap.TryGetValue(typeSymbol, out var mapped)) { - builder.Append(symbol.Name); - if (symbol is INamedTypeSymbol { IsGenericType: true } genericType) - { - builder.Append("_"); - foreach (var typeArgument in genericType.TypeArguments) - { - Traverse(typeArgument); - } - } + baseName = mapped; + } + else + { + baseName = BuildRawBaseName(typeSymbol); } - Traverse(identity.Type); + var builder = new StringBuilder(baseName); if (identity.Name != null) { @@ -666,6 +760,7 @@ void Traverse(ITypeSymbol symbol) builder.Append("_"); builder.Append(identity.ReverseIndex); } + return builder.ToString(); } @@ -719,4 +814,4 @@ private static string ReadAttributesFile() using var reader = new StreamReader(manifestResourceStream); return reader.ReadToEnd(); } -} +} \ No newline at end of file diff --git a/src/Jab/Jab.Common.props b/src/Jab/Jab.Common.props index 04c0d7e..76ee260 100644 --- a/src/Jab/Jab.Common.props +++ b/src/Jab/Jab.Common.props @@ -4,7 +4,7 @@ netstandard2.0 latest enable - 0.11.0 + 0.12.0 $(ReleaseVersion) false Jab diff --git a/src/samples/ModuleSample/ModuleSample.csproj b/src/samples/ModuleSample/ModuleSample.csproj index f73a0ec..f3db0bf 100644 --- a/src/samples/ModuleSample/ModuleSample.csproj +++ b/src/samples/ModuleSample/ModuleSample.csproj @@ -10,4 +10,4 @@ - + \ No newline at end of file