diff --git a/src/CommonLib/Processors/ACLProcessor.cs b/src/CommonLib/Processors/ACLProcessor.cs index 119a6679..9c8e69c6 100644 --- a/src/CommonLib/Processors/ACLProcessor.cs +++ b/src/CommonLib/Processors/ACLProcessor.cs @@ -12,15 +12,83 @@ using SharpHoundCommonLib.Enums; using SharpHoundCommonLib.OutputTypes; using System.Linq; +using System.Threading; namespace SharpHoundCommonLib.Processors { + /// + /// Owns state shared by processor instances and gives that state an explicit lifetime. + /// + public sealed class ACLProcessorContext : IDisposable { + private readonly ACLProcessor.GuidCache _aclGuidCache = new(); + private int _disposed; + + /// + /// Creates an that shares its GUID cache with other + /// ACL processors created by this context. + /// + public ACLProcessor CreateACLProcessor(ILdapUtils utils, ILogger log = null) { + if (Volatile.Read(ref _disposed) != 0) { + throw new ObjectDisposedException(nameof(ACLProcessorContext)); + } + + return new ACLProcessor(utils, _aclGuidCache, log); + } + + /// + /// Clears the shared processor state. Processors created by this context must not + /// be used after the context is disposed. + /// + public void Dispose() { + if (Interlocked.Exchange(ref _disposed, 1) != 0) { + return; + } + + _aclGuidCache.Dispose(); + } + } + public class ACLProcessor { private static readonly Dictionary BaseGuids; - private readonly ConcurrentDictionary _guidMap = new(); private readonly ILogger _log; private readonly ILdapUtils _utils; - private readonly ConcurrentHashSet _builtDomainCaches = new(StringComparer.OrdinalIgnoreCase); - private readonly object _lock = new(); + private readonly GuidCache _guidCache; + + internal sealed class GuidCache : IDisposable { + private readonly ConcurrentDictionary _guidMap = new(); + private readonly ConcurrentDictionary> _buildTasks = + new(StringComparer.OrdinalIgnoreCase); + private int _disposed; + + public Lazy GetOrAddBuildTask(string domain, Func> buildTaskFactory) { + ThrowIfDisposed(); + return _buildTasks.GetOrAdd(domain, _ => buildTaskFactory()); + } + + public void AddGuid(string guid, string name) { + ThrowIfDisposed(); + _guidMap.TryAdd(guid, name); + } + + public bool TryGetGuid(string guid, out string name) { + ThrowIfDisposed(); + return _guidMap.TryGetValue(guid, out name); + } + + public void Dispose() { + if (Interlocked.Exchange(ref _disposed, 1) != 0) { + return; + } + + _buildTasks.Clear(); + _guidMap.Clear(); + } + + private void ThrowIfDisposed() { + if (Volatile.Read(ref _disposed) != 0) { + throw new ObjectDisposedException(nameof(ACLProcessorContext)); + } + } + } static ACLProcessor() { //Create a dictionary with the base GUIDs of each object type @@ -42,9 +110,12 @@ static ACLProcessor() { }; } - public ACLProcessor(ILdapUtils utils, ILogger log = null) - { + public ACLProcessor(ILdapUtils utils, ILogger log = null) : this(utils, new GuidCache(), log) { + } + + internal ACLProcessor(ILdapUtils utils, GuidCache guidCache, ILogger log = null) { _utils = utils; + _guidCache = guidCache; _log = log ?? Logging.LogProvider.CreateLogger("ACLProc"); } @@ -73,14 +144,14 @@ public override string ToString() { /// LAPS /// private async Task BuildGuidCache(string domain) { - lock (_lock) { - if (_builtDomainCaches.Contains(domain)) { - return; - } + var buildTask = _guidCache.GetOrAddBuildTask(domain, + // The ExecutionAndPublication mode ensures that only one thread can execute the factory method at a time, and all other threads will wait for the result of that execution. This prevents multiple threads from building the cache simultaneously for the same domain. + () => new Lazy(() => BuildGuidCacheCore(domain), LazyThreadSafetyMode.ExecutionAndPublication)); - _builtDomainCaches.Add(domain); - } + await buildTask.Value; + } + private async Task BuildGuidCacheCore(string domain) { _log.LogInformation("Building GUID Cache for {Domain}", domain); await foreach (var result in _utils.PagedQuery(new LdapQueryParameters { DomainName = domain, @@ -108,7 +179,7 @@ private async Task BuildGuidCache(string domain) { if (name is LDAPProperties.LAPSPlaintextPassword or LDAPProperties.LAPSEncryptedPassword or LDAPProperties.LegacyLAPSPassword) { _log.LogInformation("Found GUID for ACL Right {Name}: {Guid} in domain {Domain}", name, guid, domain); - _guidMap.TryAdd(guid, name); + _guidCache.AddGuid(guid, name); } } else { _log.LogDebug("Error while building GUID cache for {Domain}: {Message}", domain, result.Error); @@ -676,7 +747,7 @@ public async IAsyncEnumerable ProcessACL(byte[] ntSecurityDescriptor, strin IsPermissionForOwnerRightsSid = isPermissionForOwnerRightsSid, IsInheritedPermissionForOwnerRightsSid = isInheritedPermissionForOwnerRightsSid, }; - else if (_guidMap.TryGetValue(aceType, out var lapsAttribute)) { + else if (_guidCache.TryGetGuid(aceType, out var lapsAttribute)) { // Compare the retrieved attribute name against LDAPProperties values if (lapsAttribute == LDAPProperties.LegacyLAPSPassword || lapsAttribute == LDAPProperties.LAPSPlaintextPassword || diff --git a/test/unit/ACLProcessorTest.cs b/test/unit/ACLProcessorTest.cs index a8e4d3b3..a01f752b 100644 --- a/test/unit/ACLProcessorTest.cs +++ b/test/unit/ACLProcessorTest.cs @@ -55,6 +55,67 @@ public void SanityCheck() { Assert.True(true); } + [Fact] + public async Task ProcessorContext_ACLProcessors_QueryOncePerDomain() { + var mockLdapUtils = new Mock(); + mockLdapUtils + .Setup(x => x.PagedQuery(It.IsAny(), It.IsAny())) + .Returns(Array.Empty>().ToAsyncEnumerable); + var domain = $"{Guid.NewGuid():N}.TEST"; + using var context = new ACLProcessorContext(); + var processors = Enumerable.Range(0, 50) + .Select(_ => context.CreateACLProcessor(mockLdapUtils.Object)) + .ToArray(); + + await Task.WhenAll(processors.Select(processor => + processor.ProcessACL(null, domain, Label.Computer, false).ToArrayAsync())); + + mockLdapUtils.Verify( + x => x.PagedQuery(It.Is(parameters => parameters.DomainName == domain), + It.IsAny()), + Times.Once); + } + + [Fact] + public async Task ProcessorContext_ACLProcessors_DoNotShareCacheAcrossContexts() { + var mockLdapUtils = new Mock(); + mockLdapUtils + .Setup(x => x.PagedQuery(It.IsAny(), It.IsAny())) + .Returns(Array.Empty>().ToAsyncEnumerable); + var domain = $"{Guid.NewGuid():N}.TEST"; + using var firstContext = new ACLProcessorContext(); + using var secondContext = new ACLProcessorContext(); + + await Task.WhenAll( + firstContext.CreateACLProcessor(mockLdapUtils.Object) + .ProcessACL(null, domain, Label.Computer, false).ToArrayAsync(), + secondContext.CreateACLProcessor(mockLdapUtils.Object) + .ProcessACL(null, domain, Label.Computer, false).ToArrayAsync()); + + mockLdapUtils.Verify( + x => x.PagedQuery(It.Is(parameters => parameters.DomainName == domain), + It.IsAny()), + Times.Exactly(2)); + } + + [Fact] + public void ProcessorContext_CreateACLProcessor_AfterDispose_Throws() { + var context = new ACLProcessorContext(); + context.Dispose(); + + Assert.Throws(() => context.CreateACLProcessor(new MockLdapUtils())); + } + + [Fact] + public async Task ProcessorContext_ACLProcessor_AfterDispose_Throws() { + var context = new ACLProcessorContext(); + var processor = context.CreateACLProcessor(new MockLdapUtils()); + context.Dispose(); + + await Assert.ThrowsAsync(() => + processor.ProcessACL(null, "TEST.LOCAL", Label.Computer, false).ToArrayAsync()); + } + [Fact] public void ACLProcessor_IsACLProtected_NullNTSD_ReturnsFalse() { var processor = new ACLProcessor(new MockLdapUtils()); @@ -2289,4 +2350,4 @@ public async Task ACLProcessor_ProcessACL_GenericWrite_Computer_WritePublicInfor Assert.Equal(actual.RightName, expectedRightName); } } -} \ No newline at end of file +}