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
+}