diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 000000000..136ecbaff --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,42 @@ +# Repository guide for coding agents + +Read `coding_standards.md` and `CONTRIBUTING.md` before changing code. Treat this file as a practical map of the repository, not as a substitute for inspecting the affected component. + +## Repository map + +- `src/CommonLib/`: `SharpHoundCommonLib`, the higher-level library for LDAP resolution, caching, processors, metrics, and collection support. +- `src/SharpHoundRPC/`: `SharpHoundRPC`, the lower-level Windows RPC, native interop, handle, and registry library. CommonLib references this project. +- `test/unit/`: xUnit tests for CommonLib, including mocks and facades. +- `RPCTest/`: xUnit tests for SharpHoundRPC. +- `docfx/`: documentation project and generated coverage output. +- `.github/workflows/build-and-test.yml`: authoritative CI build and test sequence. + +The two shipping projects target `net472` via `Directory.Build.props`; both test projects target `net8.0`. The root README's older prerequisite text does not override the project files. Windows is the CI platform and the code uses Windows and Active Directory APIs. + +## Working on a change + +1. Inspect the affected project, nearby implementation and tests, and any relevant public contract before editing. +2. Keep changes scoped. Follow the existing style in each file; do not reformat unrelated code or change target frameworks, dependencies, package metadata, or generated files without a task reason. +3. Put high-level behavior in CommonLib and native/RPC details in SharpHoundRPC. Preserve existing result, error, cancellation, and handle ownership behavior unless the task calls for changing it. +4. Add focused tests for behavior changes using local mocks or facades. Do not require a live domain, credentials, or external hosts for routine tests. +5. Run the relevant test project, then the CI sequence when shared behavior or build configuration changes. Report commands run, failures, and any environment limitation accurately. +6. Update relevant README or API documentation when a consumer-facing contract changes. + +## Commands + +Run from the repository root on Windows with the .NET 8 SDK: + +```powershell +dotnet restore +dotnet build --no-restore +dotnet test --no-build +``` + +For a focused check, use `dotnet test test/unit/CommonLibTest.csproj` or `dotnet test RPCTest/RPCTest.csproj`. `dotnet test` produces coverage files under `docfx/coverage/`. + +## Workspace care + +- Inspect `git status` before and after changes. Preserve user edits and untracked files. +- Do not commit generated coverage, `bin/`, or `obj/` output. +- Do not put secrets, credentials, or sensitive collected directory data into code, tests, logs, or documentation. +- If the requested change needs a real AD environment or Windows-only behavior that cannot be exercised locally, use the available unit tests and state the remaining validation gap. \ No newline at end of file diff --git a/CODING_STANDARDS.md b/CODING_STANDARDS.md new file mode 100644 index 000000000..11f301752 --- /dev/null +++ b/CODING_STANDARDS.md @@ -0,0 +1,41 @@ +# Coding standards + +These standards apply to new work in SharpHoundCommon. Keep changes focused and follow the style of the file being edited; the repository does not have a single enforced formatter. + +## Project boundaries and compatibility + +- `src/CommonLib` contains higher-level LDAP, cache, processor, and collection behavior. `src/SharpHoundRPC` contains lower-level Windows RPC, native interop, handle, and registry code. Keep the dependency direction from CommonLib to RPC. +- Both shipping libraries target **.NET Framework 4.7.2** through `Directory.Build.props`. Do not use an API or language feature in library code unless it builds for that target and its configured compiler. The `net8.0` test projects can use newer language features; their syntax is not a compatibility guide for library code. +- Preserve public signatures, serialized output shapes, result and error meanings, and package behavior unless a change deliberately updates that contract. Add a regression test for a behavior change. +- Keep Windows and Active Directory specifics behind the existing interfaces and wrappers so behavior can be tested without a live domain. + +## C# style + +- Use four spaces for indentation. Match the surrounding file's namespace and brace layout; both end-of-line and next-line braces exist in this repository. Avoid formatting unrelated code. +- Use `PascalCase` for types, public members, and constants; `camelCase` for parameters and locals; and `_camelCase` for private fields. Use names that reflect the AD, LDAP, RPC, or registry concept involved. +- Prefer small methods with explicit inputs and outcomes. Reuse existing interfaces, result types, and helpers instead of introducing a parallel abstraction for the same operation. +- Use `async`/`await` for asynchronous I/O. Propagate cancellation where an API accepts a `CancellationToken`; do not hide cancellation as an ordinary failure. +- Use structured `ILogger` messages with named placeholders. Do not log credentials, tokens, private keys, or raw sensitive directory data. +- In RPC and interop code, make ownership clear. Dispose native handles, buffers, and other disposable resources on success and failure paths; keep conversions and lifetime boundaries close together. +- Add XML documentation when a public API's purpose, parameters, error behavior, or ownership is not clear from its name. Update package READMEs for consumer-facing changes. + +## Tests + +- Add or update focused xUnit tests in `test/unit` for CommonLib changes and `RPCTest` for RPC changes. Put tests near the existing tests for the affected component. +- Test observable behavior and important failure paths, including null or missing LDAP values, RPC status failures, cancellation, and resource cleanup when relevant. Use the existing mocks and facades for directory, network, and native boundaries. +- Keep routine tests deterministic and independent of a live AD domain or remote host. Avoid timing-sensitive assertions and shared mutable state when practical. +- Use a descriptive test name consistent with neighboring tests; `[Theory]` is useful for related input cases. Do not add tests that only repeat implementation details. + +## Validation and review + +CI runs on Windows with the .NET 8 SDK. From the repository root, its core sequence is: + +```powershell +dotnet restore +dotnet build --no-restore +dotnet test --no-build +``` + +Run the relevant test project during development, then the full sequence for changes that affect shared code or project configuration. `dotnet test` also generates coverage under `docfx/coverage/` as described in `CONTRIBUTING.md`. + +Before review, check for unintended public API changes, compatibility with `net472`, resource leaks, sensitive logging, and unrelated formatting changes. Explain behavior changes and test evidence in the pull request. \ No newline at end of file diff --git a/src/CommonLib/ConcurrentHashSet.cs b/src/CommonLib/ConcurrentHashSet.cs index 670175cea..d29e15517 100644 --- a/src/CommonLib/ConcurrentHashSet.cs +++ b/src/CommonLib/ConcurrentHashSet.cs @@ -53,8 +53,12 @@ public IEnumerable Values() { return _backingDictionary.Keys; } + public void Clear() { + _backingDictionary?.Clear(); + } + public void Dispose() { _backingDictionary = null; GC.SuppressFinalize(this); } -} \ No newline at end of file +} diff --git a/src/CommonLib/ConnectionPoolManager.cs b/src/CommonLib/ConnectionPoolManager.cs index 828e39791..5adecbfe7 100644 --- a/src/CommonLib/ConnectionPoolManager.cs +++ b/src/CommonLib/ConnectionPoolManager.cs @@ -83,11 +83,8 @@ public void ReleaseConnection(LdapConnectionWrapper connectionWrapper, bool conn } var resolved = ResolveIdentifier(identifier); - if (!_pools.TryGetValue(resolved, out var pool)) { - pool = new LdapConnectionPool(identifier, resolved, _ldapConfig, scanner: _portScanner); - _pools.TryAdd(resolved, pool); - } - + var pool = _pools.GetOrAdd(resolved, _ => new LdapConnectionPool(identifier, resolved, _ldapConfig, scanner: _portScanner)); + return (true, pool); } @@ -139,7 +136,7 @@ private string ResolveIdentifier(string identifier) { if (Cache.GetDomainSidMapping(domainName, out var domainSid)) return (true, domainSid); try { - var entry = new DirectoryEntry($"LDAP://{domainName}").ToDirectoryObject(); + var entry = Helpers.CreateDirectoryEntry($"LDAP://{domainName}", _ldapConfig); if (entry.TryGetSecurityIdentifier(out var sid)) { Cache.AddDomainSidMapping(domainName, sid); return (true, sid); @@ -149,17 +146,12 @@ private string ResolveIdentifier(string identifier) { //we expect this to fail sometimes } - if (LdapUtils.GetDomain(domainName, _ldapConfig, out var domainObject)) - try { - // TODO: MC - Confirm GetDirectoryEntry is not a Blocking External Call - if (domainObject.GetDirectoryEntry().ToDirectoryObject().TryGetSecurityIdentifier(out domainSid)) { - Cache.AddDomainSidMapping(domainName, domainSid); - return (true, domainSid); - } - } - catch { - //we expect this to fail sometimes (not sure why, but better safe than sorry) - } + if (LdapUtils.GetDomain(domainName, _ldapConfig, out var domainObject) && + !string.IsNullOrWhiteSpace(domainObject.DomainSid)) { + domainSid = domainObject.DomainSid; + Cache.AddDomainSidMapping(domainName, domainSid); + return (true, domainSid); + } foreach (var name in _translateNames) try { diff --git a/src/CommonLib/Helpers.cs b/src/CommonLib/Helpers.cs index 19a51a37c..b90b1bdb4 100644 --- a/src/CommonLib/Helpers.cs +++ b/src/CommonLib/Helpers.cs @@ -1,5 +1,6 @@ using System; using System.Collections.Generic; +using System.DirectoryServices; using System.Globalization; using System.Linq; using System.Security.Principal; @@ -15,8 +16,11 @@ public static class Helpers { private static readonly HashSet Computers = new() { "805306369" }; private static readonly HashSet Users = new() { "805306368", "805306370" }; - private static readonly Regex DCReplaceRegex = new("DC=", RegexOptions.IgnoreCase | RegexOptions.Compiled); private static readonly Regex SPNRegex = new(@".*\/.*", RegexOptions.Compiled); + + // Splits a DN on commas that are not escaped (i.e. not preceded by a backslash). + // Plain Split(',') would break on OU/CN values that contain \, (e.g. "OU=Sales\, West"). + private static readonly Regex UnescapedCommaRegex = new(@"(?" → "LDAP://dc01.corp.com/" + // "LDAP://domain.com" → "LDAP://dc01.corp.com/DC=domain,DC=com" + // "LDAP://domain/RootDSE" → "LDAP://dc01.corp.com/RootDSE" + // Note: DisableCertVerification cannot be honoured here — there is no ADSI API for it. + var serverTarget = config.GetServerTarget(); + if (serverTarget != null) { + const string ldapPrefix = "LDAP://"; + var serverPrefix = $"{ldapPrefix}{serverTarget}/"; + + // Guard: if the path already begins with "LDAP:///" the server has + // already been injected (e.g. the caller constructed the path from a previous + // result). Injecting again would corrupt the path, so leave it unchanged. + if (!path.StartsWith(serverPrefix, StringComparison.OrdinalIgnoreCase)) { + var afterPrefix = path.Substring(ldapPrefix.Length); + + // Detect domain-shortcut targets: the component before the first '/' contains no '=' + // so it is a plain domain name (e.g. "domain.com", "domain") rather than an + // already-valid DN ("DC=domain,DC=com") or an ADSI special moniker (""). + var slashIndex = afterPrefix.IndexOf('/'); + var firstComponent = slashIndex >= 0 ? afterPrefix.Substring(0, slashIndex) : afterPrefix; + var suffix = slashIndex >= 0 ? afterPrefix.Substring(slashIndex) : string.Empty; + + if (!firstComponent.Contains('=')) { + // RootDSE is a special ADSI moniker – when the caller passes a path like + // "LDAP://domain.com/RootDSE" (used by GetNamingContextPath) the domain + // portion is only there for server selection. We must NOT convert it to + // a DN component; instead just point at the server's RootDSE directly. + if (suffix.Equals("/RootDSE", StringComparison.OrdinalIgnoreCase)) { + path = $"{ldapPrefix}{serverTarget}/RootDSE"; + } else { + // "domain.com" → "DC=domain,DC=com"; single-label "domain" → "DC=domain" + var dn = string.Join(",", firstComponent.Split('.').Select(part => $"DC={part}")); + path = $"{ldapPrefix}{serverTarget}/{dn}{suffix}"; + } + } else { + path = $"{ldapPrefix}{serverTarget}/{afterPrefix}"; + } + } + } + + if (config.Username != null) { + return new DirectoryEntry(path, config.Username, config.Password, authType) + .ToDirectoryObject(); + } + + return new DirectoryEntry(path) { AuthenticationType = authType }.ToDirectoryObject(); + } /// /// Splits a GPLink property into its representative parts @@ -129,21 +198,23 @@ public static string ConvertGuidToHexGuid(string guid) { /// Distinguished Name to extract domain from /// String representing the domain name of this object public static string DistinguishedNameToDomain(string distinguishedName) { - int idx; - if (distinguishedName.ToUpper().Contains("DELETED OBJECTS")) { - idx = distinguishedName.IndexOf("DC=", 3, StringComparison.Ordinal); - } - else { - idx = distinguishedName.IndexOf("DC=", - StringComparison.CurrentCultureIgnoreCase); + // Split on commas and collect only the trailing DC= RDNs (which always form the + // DNS domain suffix in AD DNs). Walking backward and stopping at the first non-DC= + // component correctly skips leading DC= RDNs on deleted-object tombstones and any + // over-split pieces from escaped commas in CN/OU values — DC= values are DNS labels + // and never contain commas themselves. + var rdns = UnescapedCommaRegex.Split(distinguishedName); + var dcValues = new List(); + for (var i = rdns.Length - 1; i >= 0; i--) { + var rdn = rdns[i].Trim(); + if (!rdn.StartsWith("DC=", StringComparison.OrdinalIgnoreCase)) break; + dcValues.Add(rdn.Substring(3)); } - if (idx < 0) - return null; - - var temp = distinguishedName.Substring(idx); - temp = DCReplaceRegex.Replace(temp, "").Replace(",", ".").ToUpper(); - return temp; + if (dcValues.Count == 0) return null; + dcValues.Reverse(); + // DNS identity must not depend on the process culture (for example, Turkish casing of i). + return string.Join(".", dcValues).ToUpperInvariant(); } /// diff --git a/src/CommonLib/ILdapUtils.cs b/src/CommonLib/ILdapUtils.cs index 753e53bd1..5b59e66dd 100644 --- a/src/CommonLib/ILdapUtils.cs +++ b/src/CommonLib/ILdapUtils.cs @@ -1,9 +1,10 @@ -using System; +using System; using System.Collections.Generic; using System.Security.Principal; using System.Threading; using System.Threading.Tasks; using SharpHoundCommonLib.Enums; +using SharpHoundCommonLib.Models; using SharpHoundCommonLib.OutputTypes; namespace SharpHoundCommonLib { @@ -76,18 +77,18 @@ IAsyncEnumerable> RangedRetrieval(string distinguishedName, /// A tuple containing success state as well as the resolved domain sid if successful Task<(bool Success, string DomainSid)> GetDomainSidFromDomainName(string domainName); /// - /// Attempts to retrieve the Domain object for the specified domain + /// Attempts to resolve plain domain metadata using configured LDAP settings. /// - /// The domain name to retrieve the Domain object for - /// The domain object + /// The requested domain; null selects the configured target or discovery hint. + /// The resolved metadata, or null on failure. /// True if the domain was found, false if not - bool GetDomain(string domainName, out System.DirectoryServices.ActiveDirectory.Domain domain); + bool GetDomain(string domainName, out LdapDomainInfo domain); /// - /// Attempts to retrieve the Domain object for the user's current domain + /// Attempts to resolve plain domain metadata using the configured target or discovery hint. /// - /// The domain object + /// The resolved metadata, or null on failure. /// True if the domain was found, false if not - bool GetDomain(out System.DirectoryServices.ActiveDirectory.Domain domain); + bool GetDomain(out LdapDomainInfo domain); Task<(bool Success, string ForestName)> GetForest(string domain); /// diff --git a/src/CommonLib/LdapConfig.cs b/src/CommonLib/LdapConfig.cs index 3f3e84e40..3a3a40615 100644 --- a/src/CommonLib/LdapConfig.cs +++ b/src/CommonLib/LdapConfig.cs @@ -7,12 +7,25 @@ public class LdapConfig { public string Username { get; set; } = null; public string Password { get; set; } = null; + /// + /// The DNS or NetBIOS domain associated with the user's credentials, used as an endpoint + /// hint by controlled domain resolution when neither a server nor a domain argument is supplied. + /// Takes precedence over USERDNSDOMAIN. Does not change the authentication context or + /// set explicit credentials; leave Username unset to use ambient credentials under /netonly. + /// + public string UserDomain { get; set; } = null; public string Server { get; set; } = null; public int Port { get; set; } = 0; public int SSLPort { get; set; } = 0; public bool ForceSSL { get; set; } = false; public bool DisableSigning { get; set; } = false; public bool DisableCertVerification { get; set; } = false; + /// + /// Permits legacy framework domain resolution after controlled LDAP resolution fails. + /// Authentication rejection stops resolution without invoking this fallback. + /// This fallback may ignore configured LDAP settings. Disabled by default. + /// + public bool AllowUncontrolledDomainFallback { get; set; } = false; public AuthType AuthType { get; set; } = AuthType.Kerberos; public int MaxConcurrentQueries { get; set; } = 15; @@ -35,12 +48,27 @@ public int GetGCPort(bool ssl) return ssl ? 3269 : 3268; } + /// + /// Returns the server-target string used in ADSI paths and bindings: + /// "server" when the port is the protocol default, or "server:port" when a + /// non-default port is configured. Returns null when is not set. + /// + public string GetServerTarget() + { + if (string.IsNullOrWhiteSpace(Server)) return null; + var port = GetPort(ForceSSL); + var isDefaultPort = port == (ForceSSL ? 636 : 389); + return isDefaultPort ? Server : $"{Server}:{port}"; + } + public override string ToString() { var sb = new StringBuilder(); sb.AppendLine($"Server: {Server}"); + sb.AppendLine($"UserDomain: {UserDomain}"); sb.AppendLine($"LdapPort: {GetPort(false)}"); sb.AppendLine($"LdapSSLPort: {GetPort(true)}"); sb.AppendLine($"ForceSSL: {ForceSSL}"); + sb.AppendLine($"AllowUncontrolledDomainFallback: {AllowUncontrolledDomainFallback}"); sb.AppendLine($"AuthType: {AuthType.ToString()}"); sb.AppendLine($"MaxConcurrentQueries: {MaxConcurrentQueries}"); if (!string.IsNullOrWhiteSpace(Username)) { @@ -52,5 +80,40 @@ public override string ToString() { } return sb.ToString(); } + + public string GetConfigWarnings() { + var builder = new StringBuilder(); + var hasWarning = false; + if (!string.IsNullOrWhiteSpace(Server)) { + hasWarning = true; + builder.AppendLine("-------------LDAP CONFIG WARNINGS-------------"); + builder.AppendLine($"-Explicit Server has been set to {Server}, this can degrade cross domain lookups"); + } + + if (ForceSSL && DisableCertVerification) { + if (!hasWarning) { + builder.AppendLine("-------------LDAP CONFIG WARNINGS-------------"); + } + + hasWarning = true; + builder.AppendLine("-Not all calls are able to respect DisableCertVerification, lookups may fail"); + } + + if (DisableSigning) { + if (!hasWarning) { + builder.AppendLine("-------------LDAP CONFIG WARNINGS-------------"); + } + + hasWarning = true; + builder.AppendLine("-Signing is disabled, regular LDAP traffic will be in plaintext"); + } + + if (hasWarning) { + builder.AppendLine("----------------------------------------------"); + return builder.ToString(); + } + + return string.Empty; + } } -} \ No newline at end of file +} diff --git a/src/CommonLib/LdapConnectionFactory.cs b/src/CommonLib/LdapConnectionFactory.cs new file mode 100644 index 000000000..ad46de9c5 --- /dev/null +++ b/src/CommonLib/LdapConnectionFactory.cs @@ -0,0 +1,40 @@ +using System; +using System.DirectoryServices.Protocols; +using System.Net; + +namespace SharpHoundCommonLib { + internal static class LdapConnectionFactory { + // Creates an unbound connection. The caller owns binding, retries, and disposal. + internal static LdapConnection Create(LdapConfig config, string target, bool ssl, + bool globalCatalog = false, bool pinServer = false) { + var port = globalCatalog ? config.GetGCPort(ssl) : config.GetPort(ssl); + var identifier = new LdapDirectoryIdentifier(target, port, pinServer, false); + var connection = new LdapConnection(identifier); + try { + connection.Timeout = TimeSpan.FromMinutes(5); + connection.SessionOptions.ProtocolVersion = 3; + // Referral chasing does not work with paged searches. + connection.SessionOptions.ReferralChasing = ReferralChasingOptions.None; + if (pinServer) connection.SessionOptions.AutoReconnect = false; + if (ssl) connection.SessionOptions.SecureSocketLayer = true; + + var signing = !config.DisableSigning && !ssl; + connection.SessionOptions.Signing = signing; + connection.SessionOptions.Sealing = signing; + + if (config.DisableCertVerification) + connection.SessionOptions.VerifyServerCertificate = (_, _) => true; + + if (config.Username != null) + connection.Credential = new NetworkCredential(config.Username, config.Password); + + connection.AuthType = config.AuthType; + return connection; + } + catch { + connection.Dispose(); + throw; + } + } + } +} diff --git a/src/CommonLib/LdapConnectionPool.cs b/src/CommonLib/LdapConnectionPool.cs index 5a094e865..4cff8d911 100644 --- a/src/CommonLib/LdapConnectionPool.cs +++ b/src/CommonLib/LdapConnectionPool.cs @@ -1,10 +1,8 @@ using System; using System.Collections.Concurrent; using System.Collections.Generic; -using System.DirectoryServices.ActiveDirectory; using System.DirectoryServices.Protocols; using System.Linq; -using System.Net; using System.Runtime.CompilerServices; using System.Threading; using System.Threading.Tasks; @@ -23,10 +21,10 @@ namespace SharpHoundCommonLib { internal class LdapConnectionPool : IDisposable { private readonly ConcurrentBag _connections; private readonly ConcurrentBag _globalCatalogConnection; - private readonly SemaphoreSlim _semaphore = null; private readonly string _identifier; private readonly string _poolIdentifier; private readonly LdapConfig _ldapConfig; + private readonly LdapDomainResolver _domainResolver; private readonly ILogger _log; private readonly IPortScanner _portScanner; private readonly NativeMethods _nativeMethods; @@ -44,24 +42,18 @@ internal class LdapConnectionPool : IDisposable { private readonly IMetricRouter _metric; // Tracks domains we know we've determined we shouldn't try to connect to - private static readonly ConcurrentHashSet _excludedDomains = new(); + private static readonly ConcurrentHashSet ExcludedDomains = new(); public LdapConnectionPool(string identifier, string poolIdentifier, LdapConfig config, - IPortScanner scanner = null, NativeMethods nativeMethods = null, ILogger log = null, IMetricRouter metric = null) { - _connections = new ConcurrentBag(); - _globalCatalogConnection = new ConcurrentBag(); - //TODO: Re-enable this once we track down the semaphore deadlock - // if (config.MaxConcurrentQueries > 0) { - // _semaphore = new SemaphoreSlim(config.MaxConcurrentQueries, config.MaxConcurrentQueries); - // } else { - // //If MaxConcurrentQueries is 0, we'll just disable the semaphore entirely - // _semaphore = null; - // } - + IPortScanner scanner = null, NativeMethods nativeMethods = null, ILogger log = null, IMetricRouter metric = null, + LdapDomainResolver domainResolver = null) { + _connections = []; + _globalCatalogConnection = []; _identifier = identifier; _poolIdentifier = poolIdentifier; _ldapConfig = config; _log = log ?? Logging.LogProvider.CreateLogger("LdapConnectionPool"); + _domainResolver = domainResolver ?? new LdapDomainResolver(config, _log); _metric = metric ?? Metrics.Factory.CreateMetricRouter(); _portScanner = scanner ?? new PortScanner(); _nativeMethods = nativeMethods ?? new NativeMethods(); @@ -103,30 +95,28 @@ public async IAsyncEnumerable> Query(LdapQueryParam var queryRetryCount = 0; var busyRetryCount = 0; - LdapResult tempResult = null; + LdapResult errorResult = null; var querySuccess = false; SearchResponse response = null; - while (!cancellationToken.IsCancellationRequested) { - //Grab our semaphore here to take one of our query slots - if (_semaphore != null) { - _log.LogTrace("Query entering semaphore with {Count} remaining for query {Info}", - _semaphore.CurrentCount, queryParameters.GetQueryInfo()); - await _semaphore.WaitAsync(cancellationToken); - _log.LogTrace("Query entered semaphore with {Count} remaining for query {Info}", - _semaphore.CurrentCount, queryParameters.GetQueryInfo()); - } + /* + * Retry loop structure: + * queryRetryCount — tracks connection-level failures (null response, ServerDown). On ServerDown, + * an inner retry loop attempts to establish a new connection before the outer loop continues. + * busyRetryCount — tracks server-busy/timeout conditions. The same connection is reused after + * an exponential backoff delay, since the server is expected to become available again. + */ + while (!cancellationToken.IsCancellationRequested) { try { _log.LogTrace("Sending ldap request - {Info}", queryParameters.GetQueryInfo()); response = await SendRequestWithTimeout(connectionWrapper.Connection, searchRequest, _queryAdaptiveTimeout); if (response != null) { querySuccess = true; - } - else if (queryRetryCount == MaxRetries) { + } else if (queryRetryCount == MaxRetries) { _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); - tempResult = + errorResult = LdapResult.Fail($"Failed to get a response after {MaxRetries} attempts", queryParameters); } @@ -153,8 +143,10 @@ public async IAsyncEnumerable> Query(LdapQueryParam queryRetryCount++; _log.LogDebug("Query - Attempting to recover from ServerDown for query {Info} (Attempt {Count})", queryParameters.GetQueryInfo(), queryRetryCount); + //Call ReleaseConnection with faulted = true ReleaseConnection(connectionWrapper, true); + //Try MaxRetries times to get a new connection for (var retryCount = 0; retryCount < MaxRetries; retryCount++) { var backoffDelay = GetNextBackoff(retryCount); await Task.Delay(backoffDelay, cancellationToken); @@ -168,13 +160,14 @@ public async IAsyncEnumerable> Query(LdapQueryParam break; } - //If we hit our max retries for making a new connection, set tempResult so we can yield it after this logic + //If we hit our max retries for making a new connection, set errorResult so we can yield it after this logic if (retryCount == MaxRetries - 1) { _log.LogError("Query - Failed to get a new connection after ServerDown.\n{Info}", queryParameters.GetQueryInfo()); - tempResult = + errorResult = LdapResult.Fail( - "Query - Failed to get a new connection after ServerDown.", queryParameters); + "Query - Failed to get a new connection after ServerDown.", queryParameters, + le.ErrorCode); } } } @@ -209,9 +202,9 @@ public async IAsyncEnumerable> Query(LdapQueryParam */ _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); - tempResult = LdapResult.Fail( + errorResult = LdapResult.Fail( $"Query - Caught unrecoverable ldap exception: {le.Message} (ServerMessage: {le.ServerErrorMessage}) (ErrorCode: {le.ErrorCode})", - queryParameters); + queryParameters, le.ErrorCode); } catch (Exception e) { /* @@ -219,31 +212,21 @@ public async IAsyncEnumerable> Query(LdapQueryParam */ _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); - tempResult = + errorResult = LdapResult.Fail($"Query - Caught unrecoverable exception: {e.Message}", queryParameters); } - finally { - // Always release our semaphore to prevent deadlocks - if (_semaphore != null) { - _log.LogTrace("Query releasing semaphore with {Count} remaining for query {Info}", - _semaphore.CurrentCount, queryParameters.GetQueryInfo()); - _semaphore.Release(); - _log.LogTrace("Query released semaphore with {Count} remaining for query {Info}", - _semaphore.CurrentCount, queryParameters.GetQueryInfo()); - } - } - //If we have a tempResult set it means we hit an error we couldn't recover from, so yield that result and then break out of the function - if (tempResult != null) { - if (tempResult.ErrorCode == (int)LdapErrorCodes.ServerDown) { + //If we have a errorResult set it means we hit an error we couldn't recover from, so yield that result and then break out of the function + if (errorResult != null) { + if (errorResult.ErrorCode == (int)LdapErrorCodes.ServerDown) { ReleaseConnection(connectionWrapper, true); } else { ReleaseConnection(connectionWrapper); } - yield return tempResult; + yield return errorResult; yield break; } @@ -283,17 +266,20 @@ public async IAsyncEnumerable> PagedQuery(LdapQuery PageResultResponseControl pageResponse = null; var busyRetryCount = 0; var queryRetryCount = 0; - LdapResult tempResult = null; - + + LdapResult errorResult = null; + + /* + * Retry loop structure: + * queryRetryCount — tracks connection-level failures (null response, ServerDown). On ServerDown, + * an inner retry loop attempts to establish a new connection to the same server before the + * outer loop continues. Same-server reconnection is required because the LDAP server holds + * paging state (the cookie) per-connection; switching servers would invalidate that state. + * busyRetryCount — tracks server-busy/timeout conditions. The same connection is reused after + * an exponential backoff delay, since the server is expected to become available again. + * This counter resets to 0 after each successful page to give every page a fresh budget. + */ while (!cancellationToken.IsCancellationRequested) { - if (_semaphore != null) { - _log.LogTrace("PagedQuery entering semaphore with {Count} remaining for query {Info}", - _semaphore.CurrentCount, queryParameters.GetQueryInfo()); - await _semaphore.WaitAsync(cancellationToken); - _log.LogTrace("PagedQuery entered semaphore with {Count} remaining for query {Info}", - _semaphore.CurrentCount, queryParameters.GetQueryInfo()); - } - SearchResponse response = null; try { _log.LogTrace("Sending paged ldap request - {Info}", queryParameters.GetQueryInfo()); @@ -301,12 +287,10 @@ public async IAsyncEnumerable> PagedQuery(LdapQuery if (response != null) { pageResponse = (PageResultResponseControl)response.Controls .Where(x => x is PageResultResponseControl).DefaultIfEmpty(null).FirstOrDefault(); - queryRetryCount = 0; - } - else if (queryRetryCount == MaxRetries) { + } else if (queryRetryCount == MaxRetries) { _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); - tempResult = LdapResult.Fail( + errorResult = LdapResult.Fail( $"PagedQuery - Failed to get a response after {MaxRetries} attempts", queryParameters); } @@ -316,20 +300,21 @@ public async IAsyncEnumerable> PagedQuery(LdapQuery queryRetryCount++; } } - catch (LdapException le) when (le.ErrorCode == (int)LdapErrorCodes.ServerDown) { + catch (LdapException le) when (le.ErrorCode == (int)LdapErrorCodes.ServerDown && queryRetryCount < MaxRetries) { /* * A ServerDown exception indicates that our connection is no longer valid for one of many reasons. - * We'll want to release our connection back to the pool, but dispose it. We need a new connection, - * and because this is not a paged query, we can get this connection from anywhere. + * We'll want to release our connection back to the pool, but dispose it. We need a new connection. * - * We use queryRetryCount here to prevent an infinite retry loop from occurring + * Unlike non-paged queries, paged queries MUST reconnect to the same server because the server + * maintains paging state (the cookie) per-connection. Connecting to a different server would + * invalidate the cookie and the query would have to restart from scratch. * - * Release our connection in a faulted state since the connection is defunct. - * Paged queries require a connection to be made to the same server which we started the paged query on + * We use queryRetryCount here to prevent an infinite retry loop from occurring. */ _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); if (serverName == null) { + // Because we MUST connect back to the original server, if we don't have the server to connect too, we just have to exit out here _log.LogError( "PagedQuery - Received server down exception without a known servername. Unable to generate new connection\n{Info}", queryParameters.GetQueryInfo()); @@ -337,6 +322,7 @@ public async IAsyncEnumerable> PagedQuery(LdapQuery yield break; } + queryRetryCount++; _log.LogDebug( "PagedQuery - Attempting to recover from ServerDown for query {Info} (Attempt {Count})", queryParameters.GetQueryInfo(), queryRetryCount); @@ -357,8 +343,9 @@ public async IAsyncEnumerable> PagedQuery(LdapQuery if (retryCount == MaxRetries - 1) { _log.LogError("PagedQuery - Failed to get a new connection after ServerDown.\n{Info}", queryParameters.GetQueryInfo()); - tempResult = - LdapResult.Fail("Failed to get a new connection after serverdown", + errorResult = + LdapResult.Fail( + "PagedQuery - Failed to get a new connection after ServerDown.", queryParameters, le.ErrorCode); } } @@ -391,36 +378,27 @@ public async IAsyncEnumerable> PagedQuery(LdapQuery catch (LdapException le) { _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); - tempResult = LdapResult.Fail( + errorResult = LdapResult.Fail( $"PagedQuery - Caught unrecoverable ldap exception: {le.Message} (ServerMessage: {le.ServerErrorMessage}) (ErrorCode: {le.ErrorCode})", queryParameters, le.ErrorCode); } catch (Exception e) { _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); - tempResult = + errorResult = LdapResult.Fail($"PagedQuery - Caught unrecoverable exception: {e.Message}", queryParameters); } - finally { - if (_semaphore != null) { - _log.LogTrace("PagedQuery releasing semaphore with {Count} remaining for query {Info}", - _semaphore.CurrentCount, queryParameters.GetQueryInfo()); - _semaphore.Release(); - _log.LogTrace("PagedQuery released semaphore with {Count} remaining for query {Info}", - _semaphore.CurrentCount, queryParameters.GetQueryInfo()); - } - } - if (tempResult != null) { - if (tempResult.ErrorCode == (int)LdapErrorCodes.ServerDown) { + if (errorResult != null) { + if (errorResult.ErrorCode == (int)LdapErrorCodes.ServerDown) { ReleaseConnection(connectionWrapper, true); } else { ReleaseConnection(connectionWrapper); } - yield return tempResult; + yield return errorResult; yield break; } @@ -434,6 +412,10 @@ public async IAsyncEnumerable> PagedQuery(LdapQuery continue; } + // Reset busy and query retry count after a successfully delivered page so each page starts with a fresh budget + busyRetryCount = 0; + queryRetryCount = 0; + foreach (SearchResultEntry entry in response.Entries) { if (cancellationToken.IsCancellationRequested) { ReleaseConnection(connectionWrapper); @@ -521,51 +503,56 @@ public async IAsyncEnumerable> RangedRetrieval(string distinguish var queryRetryCount = 0; var busyRetryCount = 0; - LdapResult tempResult = null; - + LdapResult errorResult = null; + + /* + * Retry loop structure: + * queryRetryCount — tracks connection-level failures (null response, ServerDown). On ServerDown, + * an inner retry loop attempts to establish a new connection before the outer loop continues. + * This counter resets to 0 after each successful range chunk to give subsequent chunks a fresh budget. + * busyRetryCount — tracks server-busy/timeout conditions. The same connection is reused after + * an exponential backoff delay, since the server is expected to become available again. + * This counter resets to 0 after each successful range chunk to give subsequent chunks a fresh budget. + */ while (!cancellationToken.IsCancellationRequested) { SearchResponse response = null; - if (_semaphore != null) { - _log.LogTrace("RangedRetrieval entering semaphore with {Count} remaining for query {Info}", - _semaphore.CurrentCount, queryParameters.GetQueryInfo()); - await _semaphore.WaitAsync(cancellationToken); - _log.LogTrace("RangedRetrieval entered semaphore with {Count} remaining for query {Info}", - _semaphore.CurrentCount, queryParameters.GetQueryInfo()); - } try { response = await SendRequestWithTimeout(connectionWrapper.Connection, searchRequest, _rangedRetrievalAdaptiveTimeout); - } - catch (LdapException le) when (le.ErrorCode == (int)ResultCode.Busy && busyRetryCount < MaxRetries) { - _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, - new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); - busyRetryCount++; - _log.LogDebug("RangedRetrieval - Executing busy backoff for query {Info} (Attempt {Count})", - queryParameters.GetQueryInfo(), busyRetryCount); - var backoffDelay = GetNextBackoff(busyRetryCount); - await Task.Delay(backoffDelay, cancellationToken); - } - catch (TimeoutException) when (busyRetryCount < MaxRetries) { - /* - * Treat a timeout as a busy error - */ - _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, - new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); - busyRetryCount++; - _log.LogDebug("RangedRetrieval - Timeout: Executing busy backoff for query {Info} (Attempt {Count})", - queryParameters.GetQueryInfo(), busyRetryCount); - var backoffDelay = GetNextBackoff(busyRetryCount); - await Task.Delay(backoffDelay, cancellationToken); + if (response != null) { + // Reset retry counters on a successful response so the next range chunk starts with a fresh budget + queryRetryCount = 0; + busyRetryCount = 0; + } + else if (queryRetryCount == MaxRetries) { + errorResult = LdapResult.Fail( + $"RangedRetrieval - Failed to get a response after {MaxRetries} attempts", + queryParameters); + } + else { + _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, + new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); + queryRetryCount++; + } } catch (LdapException le) when (le.ErrorCode == (int)LdapErrorCodes.ServerDown && queryRetryCount < MaxRetries) { + /* + * A ServerDown exception indicates that our connection is no longer valid for one of many reasons. + * We'll want to release our connection back to the pool, but dispose it. We need a new connection, + * and because ranged retrieval does not use paging state, we can connect to any server in the domain. + * + * We use queryRetryCount here to prevent an infinite retry loop from occurring. + */ _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); queryRetryCount++; _log.LogDebug( "RangedRetrieval - Attempting to recover from ServerDown for query {Info} (Attempt {Count})", queryParameters.GetQueryInfo(), queryRetryCount); + //Call ReleaseConnection with faulted = true ReleaseConnection(connectionWrapper, true); + for (var retryCount = 0; retryCount < MaxRetries; retryCount++) { var backoffDelay = GetNextBackoff(retryCount); await Task.Delay(backoffDelay, cancellationToken); @@ -584,51 +571,79 @@ public async IAsyncEnumerable> RangedRetrieval(string distinguish _log.LogError( "RangedRetrieval - Failed to get a new connection after ServerDown for path {Path}", distinguishedName); - tempResult = + errorResult = LdapResult.Fail( "RangedRetrieval - Failed to get a new connection after ServerDown.", queryParameters, le.ErrorCode); } } } + catch (LdapException le) when (le.ErrorCode == (int)ResultCode.Busy && busyRetryCount < MaxRetries) { + /* + * If we get a busy error, we want to do an exponential backoff, but maintain the current connection. + * The expectation is that given enough time, the server should stop being busy and service our query appropriately. + */ + _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, + new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); + busyRetryCount++; + _log.LogDebug("RangedRetrieval - Executing busy backoff for query {Info} (Attempt {Count})", + queryParameters.GetQueryInfo(), busyRetryCount); + var backoffDelay = GetNextBackoff(busyRetryCount); + await Task.Delay(backoffDelay, cancellationToken); + } + catch (TimeoutException) when (busyRetryCount < MaxRetries) { + /* + * Treat a timeout as a busy error + */ + _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, + new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); + busyRetryCount++; + _log.LogDebug("RangedRetrieval - Timeout: Executing busy backoff for query {Info} (Attempt {Count})", + queryParameters.GetQueryInfo(), busyRetryCount); + var backoffDelay = GetNextBackoff(busyRetryCount); + await Task.Delay(backoffDelay, cancellationToken); + } catch (LdapException le) { + /* + * This is our fallback catch. If our retry counts have been exhausted this will trigger and break us out of our loop + */ _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); - tempResult = LdapResult.Fail( - $"Caught unrecoverable ldap exception: {le.Message} (ServerMessage: {le.ServerErrorMessage}) (ErrorCode: {le.ErrorCode})", + errorResult = LdapResult.Fail( + $"RangedRetrieval - Caught unrecoverable ldap exception: {le.Message} (ServerMessage: {le.ServerErrorMessage}) (ErrorCode: {le.ErrorCode})", queryParameters, le.ErrorCode); } catch (Exception e) { + /* + * Generic exception handling for unforeseen circumstances + */ _metric.Observe(LdapMetricDefinitions.FailedRequests, 1, new LabelValues([nameof(LdapConnectionPool), _poolIdentifier])); - tempResult = - LdapResult.Fail($"Caught unrecoverable exception: {e.Message}", queryParameters); - } - finally { - if (_semaphore != null) { - _log.LogTrace("RangedRetrieval releasing semaphore with {Count} remaining for query {Info}", - _semaphore.CurrentCount, queryParameters.GetQueryInfo()); - _semaphore.Release(); - _log.LogTrace("RangedRetrieval released semaphore with {Count} remaining for query {Info}", - _semaphore.CurrentCount, queryParameters.GetQueryInfo()); - } + errorResult = + LdapResult.Fail($"RangedRetrieval - Caught unrecoverable exception: {e.Message}", queryParameters); } - //If we have a tempResult set it means we hit an error we couldn't recover from, so yield that result and then break out of the function + //If we have a errorResult set it means we hit an error we couldn't recover from, so yield that result and then break out of the function //We handle connection release in the relevant exception blocks - if (tempResult != null) { - if (tempResult.ErrorCode == (int)LdapErrorCodes.ServerDown) { + if (errorResult != null) { + if (errorResult.ErrorCode == (int)LdapErrorCodes.ServerDown) { ReleaseConnection(connectionWrapper, true); } else { ReleaseConnection(connectionWrapper); } - yield return tempResult; + yield return errorResult; yield break; } + + if (response == null) { + // response is null when queryRetryCount has been incremented but MaxRetries not yet reached; + // the outer while loop will retry automatically + continue; + } - if (response?.Entries.Count == 1) { + if (response.Entries.Count == 1) { var entry = response.Entries[0]; //We dont know the name of our attribute, but there should only be one, so we're safe to just use a loop here foreach (string attr in entry.Attributes.AttributeNames) { @@ -678,25 +693,20 @@ private static TimeSpan GetNextBackoff(int retryCount) { basePath = queryParameters.SearchBase; } else if (!connectionWrapper.GetSearchBase(queryParameters.NamingContext, out basePath)) { - string tempPath; - if (CallDsGetDcName(queryParameters.DomainName, out var info) && info != null) { - tempPath = Helpers.DomainNameToDistinguishedName(info.Value.DomainName); - connectionWrapper.SaveContext(queryParameters.NamingContext, basePath); + // Native discovery supplies a default domain DN. Configuration and schema + // contexts must come from the server's advertised naming contexts. + if (queryParameters.NamingContext == NamingContext.Default && + CallDsGetDcName(queryParameters.DomainName, out var info) && info != null) { + basePath = Helpers.DomainNameToDistinguishedName(info.Value.DomainName); } - else if (LdapUtils.GetDomain(queryParameters.DomainName, _ldapConfig, out var domainObject)) { - tempPath = Helpers.DomainNameToDistinguishedName(domainObject.Name); + else if (_domainResolver.TryResolveWithFallback(queryParameters.DomainName, out var domainObject, out _)) { + basePath = domainObject.GetNamingContext(queryParameters.NamingContext); } else { return (false, null); } - basePath = queryParameters.NamingContext switch { - NamingContext.Configuration => $"CN=Configuration,{tempPath}", - NamingContext.Schema => $"CN=Schema,CN=Configuration,{tempPath}", - NamingContext.Default => tempPath, - _ => throw new ArgumentOutOfRangeException() - }; - + if (string.IsNullOrWhiteSpace(basePath)) return (false, null); connectionWrapper.SaveContext(queryParameters.NamingContext, basePath); } @@ -741,7 +751,7 @@ private bool CallDsGetDcName(string domainName, out NetAPIStructs.DomainControll public async Task<(bool Success, LdapConnectionWrapper ConnectionWrapper, string Message)> GetConnectionAsync() { - if (_excludedDomains.Contains(_identifier)) { + if (ExcludedDomains.Contains(_identifier)) { return (false, null, $"Identifier {_identifier} excluded for connection attempt"); } @@ -775,14 +785,13 @@ private bool CallDsGetDcName(string domainName, out NetAPIStructs.DomainControll public async Task<(bool Success, LdapConnectionWrapper ConnectionWrapper, string Message)> GetGlobalCatalogConnectionAsync() { - if (_excludedDomains.Contains(_identifier)) { + if (ExcludedDomains.Contains(_identifier)) { return (false, null, $"Identifier {_identifier} excluded for connection attempt"); } if (!_globalCatalogConnection.TryTake(out var connectionWrapper)) { var (success, connection, message) = await CreateNewConnection(true); if (!success) { - //If we didn't get a connection, immediately release the semaphore so we don't have hanging ones return (false, null, message); } @@ -802,13 +811,17 @@ public void ReleaseConnection(LdapConnectionWrapper connectionWrapper, bool conn } } else { - connectionWrapper.Connection.Dispose(); + connectionWrapper.Connection?.Dispose(); } } public void Dispose() { while (_connections.TryTake(out var wrapper)) { - wrapper.Connection.Dispose(); + wrapper.Connection?.Dispose(); + } + + while (_globalCatalogConnection.TryTake(out var wrapper)) { + wrapper.Connection?.Dispose(); } } @@ -856,12 +869,12 @@ await CreateLdapConnection(tempDomainName, globalCatalog) is (true, var connecti } } - if (!LdapUtils.GetDomain(_identifier, _ldapConfig, out var domainObject) || domainObject?.Name == null) { + if (!_domainResolver.TryResolveWithFallback(_identifier, out var domainObject, out _) || domainObject?.Name == null) { //If we don't get a result here, we effectively have no other ways to resolve this domain, so we'll just have to exit out _log.LogDebug( "Could not get domain object from GetDomain, unable to create ldap connection for domain {Domain}", _identifier); - _excludedDomains.Add(_identifier); + ExcludedDomains.Add(_identifier); return (false, null, "Unable to get domain object for further strategies"); } @@ -875,32 +888,24 @@ await CreateLdapConnection(tempDomainName, globalCatalog) is (true, var connecti return (true, connectionWrapper4, ""); } - var primaryDomainController = domainObject.PdcRoleOwner.Name; - var portConnectionResult = - await CreateLDAPConnectionWithPortCheck(primaryDomainController, globalCatalog); - if (portConnectionResult.success) { - _log.LogDebug( - "Successfully created ldap connection for domain: {Domain} using strategy 5 with to pdc {Server}", - _identifier, primaryDomainController); - return (true, portConnectionResult.connection, ""); - } - - // Blocking External Call - Possible on domainObject.DomainControllers as it calls DsGetDcNameWrapper - foreach (DomainController dc in domainObject.DomainControllers) { - portConnectionResult = - await CreateLDAPConnectionWithPortCheck(dc.Name, globalCatalog); - if (portConnectionResult.success) { + // Try the PDC first, then the remaining controller metadata. Missing hostnames + // are optional metadata and must not become connection targets. + var controllerNames = new[] { domainObject.PdcRoleOwnerName }.Concat(domainObject.DomainControllerNames); + foreach (var hostname in controllerNames) { + if (string.IsNullOrWhiteSpace(hostname)) continue; + var result = await CreateLDAPConnectionWithPortCheck(hostname, globalCatalog); + if (result.success) { _log.LogDebug( - "Successfully created ldap connection for domain: {Domain} using strategy 6 with to pdc {Server}", - _identifier, primaryDomainController); - return (true, portConnectionResult.connection, ""); + "Successfully created ldap connection for domain: {Domain} using controller metadata to server {Server}", + _identifier, hostname); + return (true, result.connection, ""); } } } catch (Exception e) { _log.LogInformation(e, "We will not be able to connect to domain {Domain} by any strategy, leaving it.", _identifier); - _excludedDomains.Add(_identifier); + ExcludedDomains.Add(_identifier); } return (false, null, "All attempted connections failed"); @@ -954,36 +959,7 @@ await CreateLdapConnection(tempDomainName, globalCatalog) is (true, var connecti private LdapConnection CreateBaseConnection(string directoryIdentifier, bool ssl, bool globalCatalog) { _log.LogDebug("Creating connection for identifier {Identifier}", directoryIdentifier); - var port = globalCatalog ? _ldapConfig.GetGCPort(ssl) : _ldapConfig.GetPort(ssl); - var identifier = new LdapDirectoryIdentifier(directoryIdentifier, port, false, false); - var connection = new LdapConnection(identifier) { Timeout = new TimeSpan(0, 0, 5, 0) }; - - //These options are important! - connection.SessionOptions.ProtocolVersion = 3; - //Referral chasing does not work with paged searches - connection.SessionOptions.ReferralChasing = ReferralChasingOptions.None; - if (ssl) connection.SessionOptions.SecureSocketLayer = true; - - if (_ldapConfig.DisableSigning || ssl) { - connection.SessionOptions.Signing = false; - connection.SessionOptions.Sealing = false; - } - else { - connection.SessionOptions.Signing = true; - connection.SessionOptions.Sealing = true; - } - - if (_ldapConfig.DisableCertVerification) - connection.SessionOptions.VerifyServerCertificate = (_, _) => true; - - if (_ldapConfig.Username != null) { - var cred = new NetworkCredential(_ldapConfig.Username, _ldapConfig.Password); - connection.Credential = cred; - } - - connection.AuthType = _ldapConfig.AuthType; - - return connection; + return LdapConnectionFactory.Create(_ldapConfig, directoryIdentifier, ssl, globalCatalog); } /// @@ -1079,7 +1055,7 @@ await _portScanner.CheckPort(target, _ldapConfig.GetGCPort(false)))) else { if (await _portScanner.CheckPort(target, _ldapConfig.GetPort(true)) || (!_ldapConfig.ForceSSL && await _portScanner.CheckPort(target, _ldapConfig.GetPort(false)))) - return await CreateLdapConnection(target, true); + return await CreateLdapConnection(target, false); } return (false, null); diff --git a/src/CommonLib/LdapDomainResolver.Dependencies.cs b/src/CommonLib/LdapDomainResolver.Dependencies.cs new file mode 100644 index 000000000..3cdd114d6 --- /dev/null +++ b/src/CommonLib/LdapDomainResolver.Dependencies.cs @@ -0,0 +1,47 @@ +using System; +using System.Collections.Generic; +using System.DirectoryServices.Protocols; +using Microsoft.Extensions.Logging; +using SharpHoundCommonLib.Enums; + +namespace SharpHoundCommonLib { + // Internal dependency contracts and injection for tests without AD access. + // Concrete LDAP and framework adapters stay beside their resolution logic. + internal sealed partial class LdapDomainResolver { + // Creates an unbound connection owned and disposed by the resolver. + internal delegate IConnection ConnectionFactory(string target, bool ssl, bool pinServer); + + // Reads the USERDNSDOMAIN endpoint hint when no explicit target or UserDomain hint is available. + internal delegate string EnvironmentDomainReader(); + + internal LdapDomainResolver(LdapConfig config, + ConnectionFactory createConnection, EnvironmentDomainReader getEnvironmentDomain, + ILogger log = null, Func getLegacyDomain = null) { + _config = config; + _createConnection = createConnection; + _getEnvironmentDomain = getEnvironmentDomain; + _log = log ?? Logging.LogProvider.CreateLogger("LdapDomainResolver"); + _getLegacyDomain = getLegacyDomain ?? OpenLegacyDomain; + } + + // Limited to direct resolver operations and connection ownership. + internal interface IConnection : IDisposable { + void Bind(); + IReadOnlyList Search(SearchRequest request); + // A null cookie means the response omitted the paging control; empty means complete. + IReadOnlyList SearchPage(SearchRequest request, out byte[] cookie); + } + + // Owned by the resolver; only plain values leave the private framework adapter. + internal interface ILegacyDomain : IDisposable { + string Name { get; } + string DefaultNamingContext { get; } + string ForestName { get; } + string DomainSid { get; } + string PdcRoleOwnerName { get; } + string ReadNamingContext(string attribute); + IReadOnlyList ReadControllerNames(); + IReadOnlyDictionary ReadTrustTypes(); + } + } +} diff --git a/src/CommonLib/LdapDomainResolver.Legacy.cs b/src/CommonLib/LdapDomainResolver.Legacy.cs new file mode 100644 index 000000000..f9bc72cbf --- /dev/null +++ b/src/CommonLib/LdapDomainResolver.Legacy.cs @@ -0,0 +1,204 @@ +using System; +using System.Collections.Generic; +using System.DirectoryServices; +using System.DirectoryServices.ActiveDirectory; +using System.Security.Principal; +using Microsoft.Extensions.Logging; +using SharpHoundCommonLib.Models; +using TrustType = SharpHoundCommonLib.Enums.TrustType; + +namespace SharpHoundCommonLib { + internal sealed partial class LdapDomainResolver { + // Reports which resolution path supplied the identity. + internal bool TryResolveWithFallback(string domainName, out LdapDomainInfo domain, out bool usedLegacy) { + return TryResolveWithFallback(domainName, out domain, out usedLegacy, out _); + } + + internal bool TryResolveWithFallback(string domainName, out LdapDomainInfo domain, out bool usedLegacy, + out MetadataState metadata) { + usedLegacy = false; + var success = TryResolveMetadata(domainName, null, out metadata, out var authenticationRejected); + domain = metadata?.Domain; + if (success) return true; + if (authenticationRejected || !_config.AllowUncontrolledDomainFallback) return false; + + usedLegacy = true; + _log.LogWarning("Using uncontrolled framework domain fallback for domain {Domain}; configured LDAP settings may be ignored", + domainName); + success = TryResolveLegacy(domainName, null, out metadata); + domain = metadata?.Domain; + return success; + } + + private bool TryResolveLegacy(string domainName, MetadataState previous, out MetadataState metadata) { + metadata = null; + try { + using (var legacy = _getLegacyDomain(domainName)) { + if (legacy == null) return false; + var name = Normalize(legacy.Name); + var namingContext = Normalize(legacy.DefaultNamingContext); + if (name == null || DomainFromNamingContext(namingContext) == null) return false; + + if (previous != null && + (!string.Equals(previous.Domain.Name, name, StringComparison.OrdinalIgnoreCase) || + !string.Equals(previous.Domain.DefaultNamingContext, namingContext, + StringComparison.OrdinalIgnoreCase))) return false; + + var result = previous?.Copy() ?? new MetadataState { + Domain = new LdapDomainInfo { Name = name, DefaultNamingContext = namingContext }, + UsedLegacy = true, + ForestRead = false, + ConfigurationRead = false, + SchemaRead = false, + TopologyRead = true + }; + ReadLegacyAdditionalMetadata(legacy, result); + metadata = result; + } + return true; + } + catch (Exception e) { + metadata = null; + _log.LogDebug(e, "Uncontrolled domain fallback failed for domain {Domain}", domainName); + return false; + } + } + + private void ReadLegacyAdditionalMetadata(ILegacyDomain legacy, MetadataState metadata) { + var domain = metadata.Domain; + // Each read is independent: missing optional data must preserve the core identity. + ReadLegacyMetadata(domain.Name, "forest name", ref metadata.ForestRead, () => + domain.ForestName = Normalize(legacy.ForestName)); + ReadLegacyMetadata(domain.Name, "configuration naming context", ref metadata.ConfigurationRead, () => + domain.ConfigurationNamingContext = Normalize(legacy.ReadNamingContext("configurationNamingContext"))); + ReadLegacyMetadata(domain.Name, "schema naming context", ref metadata.SchemaRead, () => + domain.SchemaNamingContext = Normalize(legacy.ReadNamingContext("schemaNamingContext"))); + ReadLegacyMetadata(domain.Name, "domain SID", ref metadata.SidRead, () => + domain.DomainSid = Normalize(legacy.DomainSid)); + ReadLegacyMetadata(domain.Name, "PDC hostname", ref metadata.PdcRead, () => + domain.PdcRoleOwnerName = Normalize(legacy.PdcRoleOwnerName)); + ReadLegacyMetadata(domain.Name, "controller hostnames", ref metadata.ControllersRead, () => { + var seen = new HashSet(StringComparer.OrdinalIgnoreCase); + var names = new List(); + foreach (var controller in legacy.ReadControllerNames()) { + var hostname = Normalize(controller); + if (hostname != null && seen.Add(hostname)) names.Add(hostname); + } + domain.DomainControllerNames.AddRange(names); + }); + ReadLegacyMetadata(domain.Name, "trust classifications", ref metadata.TrustsRead, () => { + var trusts = new Dictionary(StringComparer.OrdinalIgnoreCase); + foreach (var trust in legacy.ReadTrustTypes()) { + var target = Normalize(trust.Key); + if (target != null) trusts[target] = trust.Value; + } + domain.TrustTypes.Clear(); + foreach (var trust in trusts) domain.TrustTypes.Add(trust.Key, trust.Value); + }); + } + + private void ReadLegacyMetadata(string domainName, string metadata, ref bool completed, Action read) { + if (completed) return; + try { + read(); + completed = true; + } + catch (Exception e) { + _log.LogDebug(e, "Uncontrolled domain fallback could not read additional metadata {Metadata} for domain {Domain}", + metadata, domainName); + } + } + + private ILegacyDomain OpenLegacyDomain(string domainName) { + DirectoryContext context; + if (_config.Username != null) { + context = domainName != null + ? new DirectoryContext(DirectoryContextType.Domain, domainName, _config.Username, _config.Password) + : new DirectoryContext(DirectoryContextType.Domain, _config.Username, _config.Password); + } + else { + context = domainName != null + ? new DirectoryContext(DirectoryContextType.Domain, domainName) + : new DirectoryContext(DirectoryContextType.Domain); + } + var domain = Domain.GetDomain(context); + return domain == null ? null : new LegacyDomain(domain, _config); + } + + private sealed class LegacyDomain : ILegacyDomain { + private readonly Domain _domain; + private readonly LdapConfig _config; + + internal LegacyDomain(Domain domain, LdapConfig config) { + _domain = domain; + _config = config; + } + + public string Name => _domain.Name; + + public string DefaultNamingContext { + get { + using (var entry = _domain.GetDirectoryEntry()) { + return entry.Properties["distinguishedName"].Value as string; + } + } + } + + public string ForestName { + get { + using (var forest = _domain.Forest) { + return forest.Name; + } + } + } + + public string DomainSid { + get { + using (var entry = _domain.GetDirectoryEntry()) { + var bytes = entry.Properties["objectSid"].Value as byte[]; + return bytes == null ? null : new SecurityIdentifier(bytes, 0).Value; + } + } + } + + public string PdcRoleOwnerName { + get { + using (var controller = _domain.PdcRoleOwner) { + return controller.Name; + } + } + } + + public string ReadNamingContext(string attribute) { + using (var root = new DirectoryEntry("LDAP://" + _domain.Name + "/RootDSE", + _config.Username, _config.Username == null ? null : _config.Password)) { + return root.Properties[attribute].Value as string; + } + } + + public IReadOnlyList ReadControllerNames() { + var controllers = _domain.DomainControllers; + var names = new List(); + try { + foreach (DomainController controller in controllers) names.Add(controller.Name); + } + finally { + foreach (DomainController controller in controllers) controller.Dispose(); + } + return names; + } + + public IReadOnlyDictionary ReadTrustTypes() { + var trusts = new Dictionary(StringComparer.OrdinalIgnoreCase); + foreach (TrustRelationshipInformation trust in _domain.GetAllTrustRelationships()) { + // Match enum names explicitly; the framework and output enum values differ. + trusts[trust.TargetName] = Enum.TryParse(trust.TrustType.ToString(), out TrustType type) + ? type : TrustType.Unknown; + } + return trusts; + } + + public void Dispose() => _domain.Dispose(); + } + } +} diff --git a/src/CommonLib/LdapDomainResolver.Metadata.cs b/src/CommonLib/LdapDomainResolver.Metadata.cs new file mode 100644 index 000000000..a6d1ecbfd --- /dev/null +++ b/src/CommonLib/LdapDomainResolver.Metadata.cs @@ -0,0 +1,57 @@ +using System; +using System.Collections.Generic; +using SharpHoundCommonLib.Models; + +namespace SharpHoundCommonLib { + internal sealed partial class LdapDomainResolver { + // Completion is independent of values: null, empty, and Unknown can all be successful reads. + internal sealed class MetadataState { + internal LdapDomainInfo Domain; + internal string Endpoint; + internal bool UsedLegacy; + internal bool ForestRead = true; + internal bool ConfigurationRead = true; + internal bool SchemaRead = true; + internal bool SidRead; + internal bool PdcRead; + internal bool ControllersRead; + internal bool TopologyRead; + internal bool TrustsRead; + internal Dictionary Topology = new(StringComparer.OrdinalIgnoreCase); + internal List Trusts = new(); + + internal bool Complete => ForestRead && ConfigurationRead && SchemaRead && + SidRead && PdcRead && ControllersRead && TopologyRead && TrustsRead; + + internal MetadataState Copy() { + // Topology and trust records contain only plain values and are replaced as complete sets. + // The public snapshot needs its own collections so refresh never mutates an earlier result. + var copy = (MetadataState)MemberwiseClone(); + copy.Domain = new LdapDomainInfo { + Name = Domain.Name, + DefaultNamingContext = Domain.DefaultNamingContext, + ForestName = Domain.ForestName, + ConfigurationNamingContext = Domain.ConfigurationNamingContext, + SchemaNamingContext = Domain.SchemaNamingContext, + DomainSid = Domain.DomainSid, + PdcRoleOwnerName = Domain.PdcRoleOwnerName + }; + copy.Domain.DomainControllerNames.AddRange(Domain.DomainControllerNames); + foreach (var trust in Domain.TrustTypes) copy.Domain.TrustTypes.Add(trust.Key, trust.Value); + return copy; + } + } + + internal sealed class TopologyEntry { + internal string DistinguishedName; + internal string Parent; + internal bool Valid; + } + + internal sealed class TrustRecord { + internal string Target; + internal long? Type; + internal long? Attributes; + } + } +} diff --git a/src/CommonLib/LdapDomainResolver.cs b/src/CommonLib/LdapDomainResolver.cs new file mode 100644 index 000000000..5d46e8696 --- /dev/null +++ b/src/CommonLib/LdapDomainResolver.cs @@ -0,0 +1,429 @@ +using System; +using System.Collections.Generic; +using System.DirectoryServices.Protocols; +using System.Linq; +using Microsoft.Extensions.Logging; +using SharpHoundCommonLib.Enums; +using SharpHoundCommonLib.LDAPQueries; +using SharpHoundCommonLib.Models; + +namespace SharpHoundCommonLib { + // Resolves domain identity directly from LDAP. Do not use LdapUtils or the pools here: + // pool initialization itself needs domain resolution and would recurse back into this code. + internal sealed partial class LdapDomainResolver { + private readonly LdapConfig _config; + private readonly ILogger _log; + private readonly ConnectionFactory _createConnection; + private readonly EnvironmentDomainReader _getEnvironmentDomain; + private readonly Func _getLegacyDomain; + + // The shared factory preserves LDAP settings and leaves credentials unset when no + // username is configured, allowing the bind to use ambient outbound credentials. + internal LdapDomainResolver(LdapConfig config, ILogger log = null) + : this(config, (target, ssl, pinServer) => + new Connection(LdapConnectionFactory.Create(config, target, ssl, pinServer: pinServer)), + () => Environment.GetEnvironmentVariable("USERDNSDOMAIN"), log) { } + + /// + /// Resolves a domain name and default naming context; additional naming contexts may be null. + /// Returns false with a null result when the target cannot establish the requested identity. + /// + internal bool TryResolve(string domainName, out LdapDomainInfo domain) { + var success = TryResolveMetadata(domainName, null, out var metadata, out _); + domain = metadata?.Domain; + return success; + } + + // Keep the successful resolution path and identity when refreshing its failed metadata reads. + internal bool TryRefreshMetadata(MetadataState previous, out MetadataState metadata) => + previous.UsedLegacy + ? TryResolveLegacy(previous.Domain.Name, previous, out metadata) + : TryResolveMetadata(null, previous, out metadata, out _); + + private bool TryResolveMetadata(string domainName, MetadataState previous, out MetadataState metadata, + out bool authenticationRejected) { + metadata = null; + authenticationRejected = false; + var suppliedDomain = Normalize(domainName); + var server = Normalize(_config.Server); + // A configured server selects the endpoint, but does not override validation of + // an explicitly supplied domain. USERDNSDOMAIN is only a last-resort endpoint hint. + var target = previous?.Endpoint ?? server ?? suppliedDomain; + if (target == null) { + // UserDomain describes the credential domain, which can differ from the local + // logon environment under /netonly. It guides discovery without changing credentials + // or constraining the domain advertised by a configured server. + target = Normalize(_config.UserDomain); + } + if (target == null) { + target = Normalize(_getEnvironmentDomain()); + } + if (target == null) return false; + + // Pinning disables referrals and automatic reconnection in the shared factory. + // Both protocol attempts and every search must retain this configured host. + var pinServer = server != null; + try { + using (var connection = _createConnection(target, ssl: true, pinServer: pinServer)) { + connection.Bind(); + // A mismatch or missing core data is definitive; do not retry over plaintext. + return TryReadIdentity(connection, suppliedDomain, target, previous, out metadata); + } + } + catch (Exception e) when (e is LdapException || e is DirectoryOperationException || + e is InvalidOperationException || e is ArgumentException) { + metadata = null; + _log.LogDebug(e, "Controlled domain resolution failed for endpoint {Endpoint} using SSL {SSL}", + target, true); + // Authentication rejection is definitive; transport and legacy retries would reuse credentials. + authenticationRejected = IsAuthenticationRejection(e); + if (authenticationRejected) return false; + } + + if (_config.ForceSSL) return false; + + // The SSL operation failed and plaintext is permitted. Keep the same endpoint. + try { + using (var connection = _createConnection(target, ssl: false, pinServer: pinServer)) { + connection.Bind(); + return TryReadIdentity(connection, suppliedDomain, target, previous, out metadata); + } + } + catch (Exception e) when (e is LdapException || e is DirectoryOperationException || + e is InvalidOperationException || e is ArgumentException) { + metadata = null; + _log.LogDebug(e, "Controlled domain resolution failed for endpoint {Endpoint} using SSL {SSL}", + target, false); + authenticationRejected = IsAuthenticationRejection(e); + } + + return false; + } + + private static bool IsAuthenticationRejection(Exception exception) => + exception is LdapException ldapException && + ldapException.ErrorCode is (int)LdapErrorCodes.InvalidCredentials + or (int)ResultCode.InappropriateAuthentication; + + private bool TryReadIdentity(IConnection connection, string suppliedDomain, string target, + MetadataState previous, out MetadataState metadata) { + metadata = null; + // RootDSE is the server's naming-context advertisement. The empty DN and base + // scope address that entry without needing to know a domain search base first. + var rootDseRequest = new SearchRequest("", "(objectClass=*)", SearchScope.Base, + "defaultNamingContext", "rootDomainNamingContext", "configurationNamingContext", + "schemaNamingContext"); + var entries = connection.Search(rootDseRequest); + if (entries.Count != 1) return false; + + var root = entries[0]; + var defaultNamingContext = ReadString(root, "defaultNamingContext"); + // Derive the identity from the returned DN, rather than assuming the endpoint + // or environment hint names the domain actually served by this connection. + var domainName = DomainFromNamingContext(defaultNamingContext); + if (domainName == null) return false; + + var configurationNamingContext = ReadString(root, "configurationNamingContext"); + if (previous != null) { + if (!string.Equals(previous.Domain.Name, domainName, StringComparison.OrdinalIgnoreCase) || + !string.Equals(previous.Domain.DefaultNamingContext, defaultNamingContext, + StringComparison.OrdinalIgnoreCase)) return false; + metadata = previous.Copy(); + ReadAdditionalMetadata(connection, metadata); + return true; + } + + if (!MatchesSuppliedDomain(connection, suppliedDomain, domainName, defaultNamingContext, + configurationNamingContext)) { + return false; + } + + // Only the default naming context and its domain name are required for success. + // Missing forest, configuration, or schema metadata must preserve that success. + var domain = new LdapDomainInfo { + Name = domainName, + DefaultNamingContext = defaultNamingContext, + ForestName = DomainFromNamingContext(ReadString(root, "rootDomainNamingContext")), + ConfigurationNamingContext = configurationNamingContext, + SchemaNamingContext = ReadString(root, "schemaNamingContext") + }; + metadata = new MetadataState { Domain = domain, Endpoint = target }; + ReadAdditionalMetadata(connection, metadata); + return true; + } + + private void ReadAdditionalMetadata(IConnection connection, MetadataState metadata) { + var domain = metadata.Domain; + ReadDomainMetadata(connection, metadata); + ReadOptionalMetadata(domain.Name, "controller hostnames", ref metadata.ControllersRead, () => { + // Publish only a complete search; a later-page failure leaves the list empty. + domain.DomainControllerNames.AddRange(ReadControllerNames(connection, domain.DefaultNamingContext)); + }); + ReadTrustMetadata(connection, metadata); + } + + private void ReadDomainMetadata(IConnection connection, MetadataState metadata) { + if (metadata.SidRead && metadata.PdcRead) return; + var domain = metadata.Domain; + IDirectoryObject domainRoot = null; + var rootRead = false; + ReadOptionalMetadata(domain.Name, "domain root", ref rootRead, () => { + var entries = connection.Search(new SearchRequest(domain.DefaultNamingContext, + "(objectClass=*)", SearchScope.Base, "objectSid", "fSMORoleOwner")); + if (entries.Count == 1) domainRoot = entries[0]; + }); + if (!rootRead) return; + if (domainRoot == null) { + // A completed search without a usable root is unavailable data, not a failed read. + metadata.SidRead = metadata.PdcRead = true; + return; + } + + // Retry the shared root dependency, but preserve each successful metadata read. + ReadOptionalMetadata(domain.Name, "domain SID", ref metadata.SidRead, () => { + if (domainRoot.TryGetSecurityIdentifier(out var sid)) domain.DomainSid = Normalize(sid); + }); + ReadOptionalMetadata(domain.Name, "PDC hostname", ref metadata.PdcRead, () => + domain.PdcRoleOwnerName = ReadPdcHostname(connection, domainRoot)); + } + + private void ReadTrustMetadata(IConnection connection, MetadataState metadata) { + var domain = metadata.Domain; + if (domain.ConfigurationNamingContext == null) metadata.TopologyRead = true; + ReadOptionalMetadata(domain.Name, "domain trust topology", ref metadata.TopologyRead, () => { + var request = new SearchRequest("CN=Partitions," + domain.ConfigurationNamingContext, + "(&(objectClass=crossRef)(systemFlags:1.2.840.113556.1.4.803:=2))", + SearchScope.OneLevel, "nCName", "trustParent", "distinguishedName"); + // Materialize every page before publishing topology. An incomplete search + // cannot establish that a missing trustParent denotes a tree root. + var entries = ReadPages(connection, request); + var resolved = new Dictionary(StringComparer.OrdinalIgnoreCase); + foreach (var entry in entries) { + var name = DomainFromNamingContext(ReadString(entry, "nCName")); + if (name == null) continue; + var hasDn = entry.TryGetDistinguishedName(out var dn); + var parent = ReadString(entry, "trustParent"); + var parentCount = entry.PropertyCount("trustParent"); + resolved.Add(name, new TopologyEntry { + DistinguishedName = dn, + Parent = parent, + Valid = hasDn && (parentCount == 0 || parentCount == 1 && parent != null) + }); + } + metadata.Topology = resolved; + }); + + ReadOptionalMetadata(domain.Name, "trust records", ref metadata.TrustsRead, () => { + var request = new SearchRequest(domain.DefaultNamingContext, CommonFilters.TrustedDomains, + SearchScope.Subtree, "trustPartner", "trustType", "trustAttributes"); + var trusts = new List(); + foreach (var entry in ReadPages(connection, request)) { + var target = ReadString(entry, "trustPartner"); + if (target == null) continue; + trusts.Add(new TrustRecord { + Target = target, + Type = entry.TryGetLongProperty("trustType", out var type) ? type : (long?)null, + Attributes = entry.TryGetLongProperty("trustAttributes", out var attributes) ? attributes : (long?)null + }); + } + metadata.Trusts = trusts; + }); + + // Topology recovery must reclassify even when the trust records were already read successfully. + domain.TrustTypes.Clear(); + foreach (var trust in metadata.Trusts) { + domain.TrustTypes[trust.Target] = ClassifyTrust(trust, domain, metadata.Topology); + } + } + + private static TrustType ClassifyTrust(TrustRecord trust, LdapDomainInfo domain, + IReadOnlyDictionary topology) { + // AD trustType 3 denotes an MIT Kerberos realm and takes precedence over attributes. + if (trust.Type == 3) return TrustType.Kerberos; + if ((trust.Type != 1 && trust.Type != 2) || !trust.Attributes.HasValue) return TrustType.Unknown; + var attributes = (TrustAttributes)trust.Attributes.Value; + if (!attributes.HasFlag(TrustAttributes.WithinForest)) { + return attributes.HasFlag(TrustAttributes.ForestTransitive) ? TrustType.Forest : TrustType.External; + } + return ClassifyWithinForestTrust(domain, trust.Target, topology); + } + + private static TrustType ClassifyWithinForestTrust(LdapDomainInfo domain, string target, + IReadOnlyDictionary topology) { + if (!topology.TryGetValue(domain.Name, out var source) || !topology.TryGetValue(target, out var destination)) { + return TrustType.Unknown; + } + if (!source.Valid || !destination.Valid) return TrustType.Unknown; + var sourceParent = source.Parent; + var destinationParent = destination.Parent; + if (string.Equals(sourceParent, destination.DistinguishedName, StringComparison.OrdinalIgnoreCase) || + string.Equals(destinationParent, source.DistinguishedName, StringComparison.OrdinalIgnoreCase)) { + return TrustType.ParentChild; + } + if (sourceParent == null && destinationParent == null) { + if (domain.ForestName == null) return TrustType.Unknown; + if (string.Equals(domain.Name, domain.ForestName, StringComparison.OrdinalIgnoreCase) || + string.Equals(target, domain.ForestName, StringComparison.OrdinalIgnoreCase)) return TrustType.TreeRoot; + } + return TrustType.CrossLink; + } + + private static string ReadPdcHostname(IConnection connection, IDirectoryObject domainRoot) { + var owner = ReadString(domainRoot, "fSMORoleOwner"); + const string ntdsPrefix = "CN=NTDS Settings,"; + if (owner == null || !owner.StartsWith(ntdsPrefix, StringComparison.OrdinalIgnoreCase)) return null; + // Remove only the fixed NTDS Settings RDN, preserving escaped commas in + // the parent server DN. The hostname is data, never a connection target. + var serverDn = owner.Substring(ntdsPrefix.Length); + if (string.IsNullOrWhiteSpace(serverDn)) return null; + var entries = connection.Search(new SearchRequest(serverDn, "(objectClass=server)", + SearchScope.Base, "dNSHostName")); + return entries.Count == 1 ? ReadHostname(entries[0]) : null; + } + + private static List ReadControllerNames(IConnection connection, string defaultNamingContext) { + // RODCs carry PARTIAL_SECRETS_ACCOUNT rather than SERVER_TRUST_ACCOUNT. + var controllerFilter = new LdapFilter() + .AddFilter(CommonFilters.DomainControllers, false) + .AddFilter("(userAccountControl:1.2.840.113556.1.4.803:=67108864)", false) + .GetFilter(); + var request = new SearchRequest(defaultNamingContext, controllerFilter, + SearchScope.Subtree, "dNSHostName"); + var seen = new HashSet(StringComparer.OrdinalIgnoreCase); + foreach (var entry in ReadPages(connection, request)) { + var name = ReadHostname(entry); + if (name != null) seen.Add(name); + } + return seen.ToList(); + } + + private static List ReadPages(IConnection connection, SearchRequest request) { + var pageControl = new PageResultRequestControl(500); + request.Controls.Add(pageControl); + var results = new List(); + do { + var entries = connection.SearchPage(request, out var cookie); + // Without the response control we cannot know that all pages were read. + if (cookie == null) throw new InvalidOperationException("Missing LDAP paging response control"); + results.AddRange(entries); + pageControl.Cookie = cookie; + } while (pageControl.Cookie.Length != 0); + + return results; + } + + private void ReadOptionalMetadata(string domainName, string metadata, ref bool completed, Action read) { + if (completed) return; + try { + read(); + completed = true; + } + catch (Exception e) when (e is LdapException or DirectoryOperationException or InvalidOperationException or ArgumentException or FormatException) { + _log.LogDebug(e, "Controlled domain resolution could not read additional metadata {Metadata} for domain {Domain}", + metadata, domainName); + } + } + + private static string ReadHostname(IDirectoryObject entry) { + var name = ReadString(entry, "dNSHostName"); + return name != null && Uri.CheckHostName(name) == UriHostNameType.Dns ? name : null; + } + + private static bool MatchesSuppliedDomain(IConnection connection, string suppliedDomain, + string domainName, string defaultNamingContext, string configurationNamingContext) { + // Automatic endpoint selection accepts the advertised identity. Validation applies + // only to the domain the caller explicitly requested. + if (suppliedDomain == null) return true; + + // DNS identities may be single-label names. Accept them before requiring alias metadata. + if (string.Equals(suppliedDomain, domainName, StringComparison.OrdinalIgnoreCase)) return true; + + // Dotted input must match the DNS identity; other single-label input may be an alias. + if (suppliedDomain.IndexOf('.') >= 0) { + // A single terminal dot denotes the DNS root. Remove it only for identity + // comparison, preserving the caller's endpoint and any other empty labels. + var dnsDomain = suppliedDomain; + if (dnsDomain.EndsWith(".", StringComparison.Ordinal)) { + dnsDomain = dnsDomain.Substring(0, dnsDomain.Length - 1); + } + return string.Equals(dnsDomain, domainName, StringComparison.OrdinalIgnoreCase); + } + + // RootDSE does not advertise the NetBIOS alias. Read the cross-reference for this + // specific naming context, using the same connection even with a pinned server. + // Configuration metadata becomes required when it is needed to validate an alias. + if (configurationNamingContext == null) return false; + + var partitionsDn = "CN=Partitions," + configurationNamingContext; + var filter = "(&(objectClass=crossRef)(nCName=" + EscapeFilterValue(defaultNamingContext) + "))"; + var crossRefRequest = new SearchRequest(partitionsDn, filter, SearchScope.OneLevel, + "nCName", "nETBIOSName"); + var crossRefs = connection.Search(crossRefRequest); + if (crossRefs.Count != 1) return false; + + var crossRef = crossRefs[0]; + var advertisedNamingContext = ReadString(crossRef, "nCName"); + if (!string.Equals(advertisedNamingContext, defaultNamingContext, StringComparison.OrdinalIgnoreCase)) { + return false; + } + + var advertisedAlias = ReadString(crossRef, "nETBIOSName"); + return string.Equals(advertisedAlias, suppliedDomain, StringComparison.OrdinalIgnoreCase); + } + + private static string Normalize(string value) => string.IsNullOrWhiteSpace(value) ? null : value.Trim(); + + private static string ReadString(IDirectoryObject entry, string attributeName) { + // The wrapper handles LDAP value conversion. Keep the resolver's stricter + // single-value requirement so ambiguous naming contexts or aliases cannot match. + if (entry.PropertyCount(attributeName) != 1) return null; + return entry.TryGetProperty(attributeName, out var value) ? Normalize(value) : null; + } + + private static string DomainFromNamingContext(string namingContext) { + if (namingContext == null) return null; + var name = Helpers.DistinguishedNameToDomain(namingContext); + // The shared DN helper can yield an empty DNS label for a malformed DC component. + if (name == null || name.Split('.').Any(string.IsNullOrWhiteSpace)) return null; + return name; + } + + private static string EscapeFilterValue(string value) { + // A DN used as a filter value needs filter escaping, even though it came from LDAP. + // Escape backslashes first so later replacements do not escape their own sequences. + return value.Replace("\\", "\\5c") + .Replace("*", "\\2a") + .Replace("(", "\\28") + .Replace(")", "\\29") + .Replace("\0", "\\00"); + } + + // Thin adapter over direct LDAP operations; it performs no discovery or pool access. + private sealed class Connection : IConnection { + private readonly LdapConnection _connection; + + internal Connection(LdapConnection connection) => _connection = connection; + + public void Bind() => _connection.Bind(); + + public IReadOnlyList Search(SearchRequest request) { + var response = (SearchResponse)_connection.SendRequest(request); + return WrapEntries(response); + } + + public IReadOnlyList SearchPage(SearchRequest request, out byte[] cookie) { + var response = (SearchResponse)_connection.SendRequest(request); + cookie = response.Controls.OfType().FirstOrDefault()?.Cookie; + return WrapEntries(response); + } + + private static IReadOnlyList WrapEntries(SearchResponse response) { + // Reuse the common directory attribute accessors. + return response.Entries.Cast() + .Select(entry => (IDirectoryObject)new SearchResultEntryWrapper(entry)).ToArray(); + } + + public void Dispose() => _connection.Dispose(); + } + } +} diff --git a/src/CommonLib/LdapUtils.cs b/src/CommonLib/LdapUtils.cs index 9db6d7ec0..f1fdaa3b9 100644 --- a/src/CommonLib/LdapUtils.cs +++ b/src/CommonLib/LdapUtils.cs @@ -3,7 +3,6 @@ using System.Collections.Generic; using System.DirectoryServices; using System.DirectoryServices.AccountManagement; -using System.DirectoryServices.ActiveDirectory; using System.Linq; using System.Net; using System.Net.Sockets; @@ -23,18 +22,20 @@ using SharpHoundCommonLib.Static; using SharpHoundRPC.NetAPINative; using SharpHoundRPC.PortScanner; -using Domain = System.DirectoryServices.ActiveDirectory.Domain; using Group = SharpHoundCommonLib.OutputTypes.Group; using SearchScope = System.DirectoryServices.Protocols.SearchScope; namespace SharpHoundCommonLib { public class LdapUtils : ILdapUtils { - //This cache is indexed by domain sid - private static ConcurrentDictionary _domainCache = new(); + // Successful results from either enabled resolution path are cached by requested domain name. + private ConcurrentDictionary _domainCache = new(StringComparer.OrdinalIgnoreCase); + private readonly Func _createDomainResolver; + private readonly Func _utcNow = () => DateTime.UtcNow; + private static readonly TimeSpan MetadataRetryInterval = TimeSpan.FromSeconds(30); private static ConcurrentHashSet _domainControllers = new(StringComparer.OrdinalIgnoreCase); private static ConcurrentHashSet _unresolvablePrincipals = new(StringComparer.OrdinalIgnoreCase); - private static readonly ConcurrentDictionary DomainToForestCache = + private readonly ConcurrentDictionary DomainToForestCache = new(StringComparer.OrdinalIgnoreCase); private static readonly ConcurrentDictionary @@ -95,6 +96,20 @@ public LdapUtils(NativeMethods nativeMethods = null, PortScanner scanner = null, _connectionPool = new ConnectionPoolManager(_ldapConfig, scanner: _portScanner); } + internal LdapUtils(Func createDomainResolver, + Func utcNow = null) : this() { + _createDomainResolver = createDomainResolver; + _utcNow = utcNow ?? (() => DateTime.UtcNow); + } + + private sealed class DomainCacheEntry { + internal readonly LdapDomainResolver Resolver; + internal LdapDomainResolver.MetadataState Metadata; + internal DateTime RetryAfter; + + internal DomainCacheEntry(LdapDomainResolver resolver) => Resolver = resolver; + } + public IAsyncEnumerable> RangedRetrieval(string distinguishedName, string attributeName, CancellationToken cancellationToken = new()) { return _connectionPool.RangedRetrieval(distinguishedName, attributeName, cancellationToken); @@ -174,7 +189,7 @@ public IAsyncEnumerable> PagedQuery(LdapQueryParame } try { - var entry = CreateDirectoryEntry($"LDAP://"); + var entry = Helpers.CreateDirectoryEntry($"LDAP://", _ldapConfig); if (entry.GetLabel(out type)) { Cache.AddType(sid, type); return (true, type); @@ -185,7 +200,7 @@ public IAsyncEnumerable> PagedQuery(LdapQueryParame } try { - using (var ctx = new PrincipalContext(ContextType.Domain)) { + using (var ctx = CreatePrincipalContext(_ldapConfig, tempDomain)) { // Blocking External Call var principal = Principal.FindByIdentity(ctx, IdentityType.Sid, sid); if (principal != null) { @@ -223,7 +238,7 @@ public IAsyncEnumerable> PagedQuery(LdapQueryParame } try { - var entry = CreateDirectoryEntry($"LDAP://"); + var entry = Helpers.CreateDirectoryEntry($"LDAP://", _ldapConfig); if (entry.GetLabel(out type)) { Cache.AddType(guid, type); return (true, type); @@ -234,7 +249,7 @@ public IAsyncEnumerable> PagedQuery(LdapQueryParame } try { - using (var ctx = new PrincipalContext(ContextType.Domain)) { + using (var ctx = CreatePrincipalContext(_ldapConfig, domain)) { // Blocking External Call var principal = Principal.FindByIdentity(ctx, IdentityType.Guid, guid); if (principal != null) { @@ -301,15 +316,9 @@ public IAsyncEnumerable> PagedQuery(LdapQueryParame return (true, cachedForest); } - if (GetDomain(domain, out var domainObject)) { - try { - var forestName = domainObject.Forest.Name.ToUpper(); - DomainToForestCache.TryAdd(domain, forestName); - return (true, forestName); - } - catch { - //pass - } + if (GetDomain(domain, out var domainObject) && !string.IsNullOrWhiteSpace(domainObject.ForestName)) { + var forestName = domainObject.ForestName.ToUpper(); + return (true, forestName); } var (success, forest) = await GetForestFromLdap(domain); @@ -358,7 +367,7 @@ public IAsyncEnumerable> PagedQuery(LdapQueryParame } try { - var entry = CreateDirectoryEntry($"LDAP://"); + var entry = Helpers.CreateDirectoryEntry($"LDAP://", _ldapConfig); if (entry.TryGetDistinguishedName(out var dn)) { Cache.AddDomainSidMapping(domainSid, Helpers.DistinguishedNameToDomain(dn)); return (true, Helpers.DistinguishedNameToDomain(dn)); @@ -374,7 +383,7 @@ public IAsyncEnumerable> PagedQuery(LdapQueryParame } try { - using (var ctx = new PrincipalContext(ContextType.Domain)) { + using (var ctx = CreatePrincipalContext(_ldapConfig)) { // Blocking External Call var principal = Principal.FindByIdentity(ctx, IdentityType.Sid, sid); if (principal != null) { @@ -440,7 +449,7 @@ public IAsyncEnumerable> PagedQuery(LdapQueryParame if (Cache.GetDomainSidMapping(domainName, out var domainSid)) return (true, domainSid); try { - var entry = CreateDirectoryEntry($"LDAP://{domainName}"); + var entry = Helpers.CreateDirectoryEntry($"LDAP://{domainName}", _ldapConfig); //Force load objectsid into the object cache if (entry.TryGetSecurityIdentifier(out var sid)) { Cache.AddDomainSidMapping(domainName, sid); @@ -452,17 +461,11 @@ public IAsyncEnumerable> PagedQuery(LdapQueryParame //we expect this to fail sometimes } - if (GetDomain(domainName, out var domainObject)) - try { - var entry = domainObject.GetDirectoryEntry().ToDirectoryObject(); - if (entry.TryGetSecurityIdentifier(out domainSid)) { - Cache.AddDomainSidMapping(domainName, domainSid); - return (true, domainSid); - } - } - catch { - //we expect this to fail sometimes (not sure why, but better safe than sorry) - } + if (GetDomain(domainName, out var domainObject) && !string.IsNullOrWhiteSpace(domainObject.DomainSid)) { + domainSid = domainObject.DomainSid; + Cache.AddDomainSidMapping(domainName, domainSid); + return (true, domainSid); + } foreach (var name in _translateNames) try { @@ -492,98 +495,45 @@ public IAsyncEnumerable> PagedQuery(LdapQueryParame } /// - /// Attempts to get the Domain object representing the target domain. If null is specified for the domain name, gets - /// the user's current domain + /// Resolves plain domain metadata using configured LDAP settings. A null or blank name + /// selects the configured target or discovery hint. Successful results from either enabled path are cached. + /// Failed metadata reads are retried on a later call after 30 seconds; earlier results remain unchanged. /// - /// - /// - /// - public bool GetDomain(string domainName, out Domain domain) { + public bool GetDomain(string domainName, out LdapDomainInfo domain) { + domainName = string.IsNullOrWhiteSpace(domainName) ? null : domainName.Trim(); var cacheKey = domainName ?? _nullCacheKey; - if (_domainCache.TryGetValue(cacheKey, out domain)) return true; - - try { - DirectoryContext context; - if (_ldapConfig.Username != null) - context = domainName != null - ? new DirectoryContext(DirectoryContextType.Domain, domainName, _ldapConfig.Username, - _ldapConfig.Password) - : new DirectoryContext(DirectoryContextType.Domain, _ldapConfig.Username, - _ldapConfig.Password); - else - context = domainName != null - ? new DirectoryContext(DirectoryContextType.Domain, domainName) - : new DirectoryContext(DirectoryContextType.Domain); + var cache = _domainCache; + var entry = cache.GetOrAdd(cacheKey, _ => new DomainCacheEntry( + _createDomainResolver?.Invoke(_ldapConfig) ?? new LdapDomainResolver(_ldapConfig, _log))); + // Serialize both initial resolution and refresh for this requested domain. + lock (entry) { + if (entry.Metadata == null) { + if (!entry.Resolver.TryResolveWithFallback(domainName, out domain, out _, + out var metadata)) return false; + entry.Metadata = metadata; + entry.RetryAfter = _utcNow().Add(MetadataRetryInterval); + } + else if (!entry.Metadata.Complete && _utcNow() >= entry.RetryAfter) { + if (entry.Resolver.TryRefreshMetadata(entry.Metadata, out var metadata)) { + entry.Metadata = metadata; + } + // Connection failures also back off while preserving the cached identity. + entry.RetryAfter = _utcNow().Add(MetadataRetryInterval); + } - // Blocking External Call - domain = Domain.GetDomain(context); - if (domain == null) return false; - _domainCache.TryAdd(cacheKey, domain); + domain = entry.Metadata.Domain; return true; } - catch (Exception e) { - _log.LogDebug(e, "GetDomain call failed for domain name {Name}", domainName); - domain = null; - return false; - } } - public static bool GetDomain(string domainName, LdapConfig ldapConfig, out Domain domain) { - if (_domainCache.TryGetValue(domainName, out domain)) return true; - - try { - DirectoryContext context; - if (ldapConfig.Username != null) - context = domainName != null - ? new DirectoryContext(DirectoryContextType.Domain, domainName, ldapConfig.Username, - ldapConfig.Password) - : new DirectoryContext(DirectoryContextType.Domain, ldapConfig.Username, - ldapConfig.Password); - else - context = domainName != null - ? new DirectoryContext(DirectoryContextType.Domain, domainName) - : new DirectoryContext(DirectoryContextType.Domain); - - // Blocking External Call - domain = Domain.GetDomain(context); - if (domain == null) return false; - _domainCache.TryAdd(domainName, domain); - return true; - } - catch (Exception e) { - Logging.Logger.LogDebug("Static GetDomain call failed for domain {DomainName}: {Error}", domainName, - e.Message); - domain = null; - return false; - } + /// Resolves domain metadata without caching, using the supplied LDAP settings. + public static bool GetDomain(string domainName, LdapConfig ldapConfig, out LdapDomainInfo domain) { + return new LdapDomainResolver(ldapConfig).TryResolveWithFallback(domainName, out domain, out _); } - /// - /// Attempts to get the Domain object representing the target domain. If null is specified for the domain name, gets - /// the user's current domain - /// - /// - /// - /// - public bool GetDomain(out Domain domain) { - if (_domainCache.TryGetValue(_nullCacheKey, out domain)) return true; - - try { - var context = _ldapConfig.Username != null - ? new DirectoryContext(DirectoryContextType.Domain, _ldapConfig.Username, - _ldapConfig.Password) - : new DirectoryContext(DirectoryContextType.Domain); - - // Blocking External Call - domain = Domain.GetDomain(context); - _domainCache.TryAdd(_nullCacheKey, domain); - return true; - } - catch (Exception e) { - _log.LogDebug(e, "GetDomain call failed for blank domain"); - domain = null; - return false; - } + /// Resolves domain metadata using the configured target or discovery hint. + public bool GetDomain(out LdapDomainInfo domain) { + return GetDomain(null, out domain); } public async Task<(bool Success, TypedPrincipal Principal)> ResolveAccountName(string name, string domain) { @@ -951,7 +901,7 @@ public async Task IsDomainController(string computerObjectId, string domai } try { - using (var ctx = new PrincipalContext(ContextType.Domain)) { + using (var ctx = CreatePrincipalContext(_ldapConfig, domain)) { // Blocking External Call var lookupPrincipal = Principal.FindByIdentity(ctx, IdentityType.DistinguishedName, distinguishedName); @@ -1071,6 +1021,8 @@ await GetDomainSidFromDomainName(forestName) is (true, var forestDomainSid)) { public void SetLdapConfig(LdapConfig config) { _ldapConfig = config; + _domainCache = new ConcurrentDictionary(StringComparer.OrdinalIgnoreCase); + DomainToForestCache.Clear(); _log.LogInformation("New LDAP Config Set:\n {ConfigString}", config.ToString()); _connectionPool.Dispose(); _connectionPool = new ConnectionPoolManager(_ldapConfig, scanner: _portScanner); @@ -1096,7 +1048,7 @@ public void SetLdapConfig(LdapConfig config) { }; try { - var entry = CreateDirectoryEntry($"LDAP://{domain}/RootDSE"); + var entry = Helpers.CreateDirectoryEntry($"LDAP://{domain}/RootDSE", _ldapConfig); if (entry.TryGetProperty(property, out var searchBase)) { return (true, searchBase); } @@ -1106,29 +1058,8 @@ public void SetLdapConfig(LdapConfig config) { } if (GetDomain(domain, out var domainObj)) { - try { - var entry = domainObj.GetDirectoryEntry().ToDirectoryObject(); - if (entry.TryGetProperty(property, out var searchBase)) { - return (true, searchBase); - } - } - catch { - //pass - } - - var name = domainObj.Name; - if (!string.IsNullOrWhiteSpace(name)) { - var tempPath = Helpers.DomainNameToDistinguishedName(name); - - var searchBase = context switch { - NamingContext.Configuration => $"CN=Configuration,{tempPath}", - NamingContext.Schema => $"CN=Schema,CN=Configuration,{tempPath}", - NamingContext.Default => tempPath, - _ => throw new ArgumentOutOfRangeException() - }; - - return (true, searchBase); - } + var searchBase = domainObj.GetNamingContext(context); + if (!string.IsNullOrWhiteSpace(searchBase)) return (true, searchBase); } return (false, default); @@ -1136,7 +1067,8 @@ public void SetLdapConfig(LdapConfig config) { public void ResetUtils() { _unresolvablePrincipals = new ConcurrentHashSet(StringComparer.OrdinalIgnoreCase); - _domainCache = new ConcurrentDictionary(); + _domainCache = new ConcurrentDictionary(StringComparer.OrdinalIgnoreCase); + DomainToForestCache.Clear(); _domainControllers = new ConcurrentHashSet(StringComparer.OrdinalIgnoreCase); _connectionPool?.Dispose(); _connectionPool = new ConnectionPoolManager(_ldapConfig, scanner: _portScanner); @@ -1145,12 +1077,75 @@ public void ResetUtils() { LdapMetrics.ResetInFlight(); } - private IDirectoryObject CreateDirectoryEntry(string path) { - if (_ldapConfig.Username != null) { - return new DirectoryEntry(path, _ldapConfig.Username, _ldapConfig.Password).ToDirectoryObject(); + /// + /// Computes the contextName and that should be passed + /// to a for a given . + /// + /// + /// Separated from so that the parameter-building logic + /// can be unit-tested without constructing a real (which + /// would require a live directory connection). + /// + /// + /// + /// When is set, the server hostname is returned as + /// contextName so that binds to that specific DC + /// rather than relying on domain-level DNS discovery. Non-standard ports are expressed as + /// host:port. Otherwise is returned as-is (null = let + /// the runtime discover the current domain). + /// + /// + /// + /// Signing and sealing are disabled when SSL is active, mirroring the mutual-exclusion rule + /// applied by . + /// + /// + internal static (string ContextName, ContextOptions Options) BuildPrincipalContextParameters( + LdapConfig config, string domainName = null) { + var options = ContextOptions.Negotiate; + + if (config.ForceSSL) { + options |= ContextOptions.SecureSocketLayer; + } + + // Signing and sealing are mutually exclusive with SSL (the transport provides integrity). + if (!config.DisableSigning && !config.ForceSSL) { + options |= ContextOptions.Signing | ContextOptions.Sealing; + } + + // GetServerTarget() returns null when Server is not set, so the ?? falls through to + // domainName — which itself may be null, meaning "let the runtime discover the domain". + var contextName = config.GetServerTarget() ?? domainName; + + return (contextName, options); + } + + /// + /// Creates a that targets the same DC as the connection pool + /// and applies the same SSL / signing / credential settings from . + /// + /// + /// is always used. When is + /// set, the server hostname is passed as the name argument so the runtime binds to + /// that specific DC rather than performing domain-level DNS discovery. + /// + /// + /// Note: cannot be applied here — + /// exposes no API for it. + /// + /// This is intentionally static so that Moq's Castle.DynamicProxy does not + /// encounter in an instance-method signature when building + /// test proxies against on non-Windows runtimes. + /// + private static PrincipalContext CreatePrincipalContext(LdapConfig config, string domainName = null) { + var (contextName, options) = BuildPrincipalContextParameters(config, domainName); + + if (config.Username != null) { + return new PrincipalContext(ContextType.Domain, contextName, null, options, + config.Username, config.Password); } - return new DirectoryEntry(path).ToDirectoryObject(); + return new PrincipalContext(ContextType.Domain, contextName, null, options); } public void Dispose() { @@ -1418,4 +1413,4 @@ private static string ComputeDisplayName(IDirectoryObject directoryObject, strin return displayName.ToUpper(); } } -} \ No newline at end of file +} diff --git a/src/CommonLib/Models/LdapDomainInfo.cs b/src/CommonLib/Models/LdapDomainInfo.cs new file mode 100644 index 000000000..43f54f59f --- /dev/null +++ b/src/CommonLib/Models/LdapDomainInfo.cs @@ -0,0 +1,46 @@ +using System; +using System.Collections.Generic; +using SharpHoundCommonLib.Enums; + +namespace SharpHoundCommonLib.Models; + +/// +/// Plain domain metadata with no framework directory objects or resources to dispose. +/// Successful resolution requires and . +/// Unavailable additional strings are null and unavailable collections are empty. +/// +public class LdapDomainInfo { + internal string GetNamingContext(NamingContext context) => context switch { + NamingContext.Default => DefaultNamingContext, + NamingContext.Configuration => ConfigurationNamingContext, + NamingContext.Schema => SchemaNamingContext, + _ => throw new ArgumentOutOfRangeException(nameof(context), context, null) + }; + + /// The resolved DNS domain name. + public string Name { get; set; } + + /// The DNS forest name, or null when unavailable. + public string ForestName { get; set; } + + /// The domain SID, or null when unavailable. + public string DomainSid { get; set; } + + /// The distinguished name of the domain's default naming context. + public string DefaultNamingContext { get; set; } + + /// The configuration naming context, or null when unavailable. + public string ConfigurationNamingContext { get; set; } + + /// The schema naming context, or null when unavailable. + public string SchemaNamingContext { get; set; } + + /// The PDC role owner's hostname, or null when unavailable. + public string PdcRoleOwnerName { get; set; } + + /// Available domain controller hostnames; empty when unavailable. + public List DomainControllerNames { get; } = new(); + + /// Trust classifications indexed by case-insensitive target domain name; empty when unavailable. + public Dictionary TrustTypes { get; } = new(StringComparer.OrdinalIgnoreCase); +} diff --git a/src/CommonLib/Ntlm/HttpClientFactory.cs b/src/CommonLib/Ntlm/HttpClientFactory.cs deleted file mode 100644 index 74b8e6f5b..000000000 --- a/src/CommonLib/Ntlm/HttpClientFactory.cs +++ /dev/null @@ -1,35 +0,0 @@ -using System; -using System.Net; -using System.Net.Http; - -namespace SharpHoundCommonLib.Ntlm; - -public interface IHttpClientFactory { - HttpClient CreateUnauthenticatedClient(); - HttpClient CreateAuthenticatedHttpClient(Uri Url, string authPackage = "Kerberos"); -} - -public class HttpClientFactory : IHttpClientFactory { - public HttpClient CreateUnauthenticatedClient() { - var handler = new HttpClientHandler { - ServerCertificateCustomValidationCallback = (httpRequestMessage, cert, cetChain, policyErrors) => true, - UseDefaultCredentials = false - }; - - return new HttpClient(handler); - } - - public HttpClient CreateAuthenticatedHttpClient(Uri Url, string authPackage = "Kerberos") { - var handler = new HttpClientHandler { - Credentials = new CredentialCache() { - { Url, authPackage, CredentialCache.DefaultNetworkCredentials } - }, - - PreAuthenticate = true, - ServerCertificateCustomValidationCallback = - (httpRequestMessage, cert, cetChain, policyErrors) => { return true; }, - }; - - return new HttpClient(handler); - } -} \ No newline at end of file diff --git a/src/CommonLib/Ntlm/HttpNtlmAuthenticationService.cs b/src/CommonLib/Ntlm/HttpNtlmAuthenticationService.cs index cd0a7b935..ac36bcae8 100644 --- a/src/CommonLib/Ntlm/HttpNtlmAuthenticationService.cs +++ b/src/CommonLib/Ntlm/HttpNtlmAuthenticationService.cs @@ -15,14 +15,14 @@ namespace SharpHoundCommonLib.Ntlm; /// public class HttpNtlmAuthenticationService { private readonly ILogger _logger; - private readonly IHttpClientFactory _httpClientFactory; + private readonly INtlmHttpClientFactory _ntlmHttpClientFactory; private readonly AdaptiveTimeout _getSupportedNTLMAuthSchemesAdaptiveTimeout; private readonly AdaptiveTimeout _ntlmAuthAdaptiveTimeout; private readonly AdaptiveTimeout _authWithChannelBindingAdaptiveTimeout; - public HttpNtlmAuthenticationService(IHttpClientFactory httpClientFactory, ILogger logger = null) { + public HttpNtlmAuthenticationService(INtlmHttpClientFactory ntlmHttpClientFactory, ILogger logger = null) { _logger = logger ?? Logging.LogProvider.CreateLogger(nameof(HttpNtlmAuthenticationService)); - _httpClientFactory = httpClientFactory; + _ntlmHttpClientFactory = ntlmHttpClientFactory ?? throw new ArgumentNullException(nameof(ntlmHttpClientFactory)); _getSupportedNTLMAuthSchemesAdaptiveTimeout = new AdaptiveTimeout(maxTimeout: TimeSpan.FromMinutes(2), Logging.LogProvider.CreateLogger(nameof(GetSupportedNtlmAuthSchemesAsync))); _ntlmAuthAdaptiveTimeout = new AdaptiveTimeout(maxTimeout: TimeSpan.FromMinutes(2), Logging.LogProvider.CreateLogger(nameof(NtlmAuthenticationHandler.PerformNtlmAuthenticationAsync))); _authWithChannelBindingAdaptiveTimeout = new AdaptiveTimeout(maxTimeout: TimeSpan.FromMinutes(2), Logging.LogProvider.CreateLogger(nameof(AuthWithBadChannelBindingsAsync))); @@ -57,7 +57,7 @@ public async Task EnsureRequiresAuth(Uri url, bool? useBadChannelBindings) { } private async Task GetSupportedNtlmAuthSchemesAsync(Uri url) { - var httpClient = _httpClientFactory.CreateUnauthenticatedClient(); + var httpClient = _ntlmHttpClientFactory.CreateUnauthenticatedClient(); using var getRequest = new HttpRequestMessage(HttpMethod.Get, url); var result = await _getSupportedNTLMAuthSchemesAdaptiveTimeout.ExecuteWithTimeout(async (timeoutToken) => { @@ -105,7 +105,7 @@ internal string[] ExtractAuthSchemes(HttpResponseMessage response) { } private async Task AuthWithBadChannelBindingsAsync(Uri url, string authScheme, NtlmAuthenticationHandler ntlmAuth = null) { - var httpClient = _httpClientFactory.CreateUnauthenticatedClient(); + var httpClient = _ntlmHttpClientFactory.CreateUnauthenticatedClient(); var transport = new HttpTransport(httpClient, url, authScheme, _logger); var ntlmAuthHandler = ntlmAuth ?? new NtlmAuthenticationHandler($"HTTP/{url.Host}"); @@ -143,18 +143,7 @@ private async Task AuthWithBadChannelBindingsAsync(Uri url, string authScheme, N } private async Task AuthWithChannelBindingAsync(Uri url, string authScheme) { - var handler = new HttpClientHandler { - ServerCertificateCustomValidationCallback = (httpRequestMessage, cert, cetChain, policyErrors) => true, - }; - - var credentialCache = new CredentialCache { - { url, authScheme, CredentialCache.DefaultNetworkCredentials } - }; - - handler.Credentials = credentialCache; - handler.PreAuthenticate = true; - - using var client = new HttpClient(handler); + using var client = _ntlmHttpClientFactory.CreateAuthenticatedHttpClient(url, authScheme); var result = await _authWithChannelBindingAdaptiveTimeout.ExecuteWithTimeout(async (timeoutToken) => { try { @@ -217,4 +206,4 @@ public AuthNotRequiredException() { public AuthNotRequiredException(string message) : base(message) { } -} \ No newline at end of file +} diff --git a/src/CommonLib/Ntlm/NtlmHttpClientFactory.cs b/src/CommonLib/Ntlm/NtlmHttpClientFactory.cs new file mode 100644 index 000000000..cb5e917d6 --- /dev/null +++ b/src/CommonLib/Ntlm/NtlmHttpClientFactory.cs @@ -0,0 +1,62 @@ +using System; +using System.Net; +using System.Net.Http; +using System.Security.Authentication; + +namespace SharpHoundCommonLib.Ntlm; + +public interface INtlmHttpClientFactory { + HttpClient CreateUnauthenticatedClient(); + HttpClient CreateAuthenticatedHttpClient(Uri Url, string authPackage = "Kerberos"); +} + +public class NtlmHttpClientFactory : INtlmHttpClientFactory { + private readonly SslProtocols _sslProtocols; + + /// + /// Creates an HttpClientFactory whose handlers will negotiate TLS using OS/framework defaults. + /// + public NtlmHttpClientFactory() : this(SslProtocols.None) { } + + /// + /// Creates an HttpClientFactory whose handlers will restrict TLS negotiation to the specified protocols. + /// Use this overload when a specific set of legacy protocols must be supported for a target service, + /// rather than setting process-wide. + /// + /// + /// The SSL/TLS protocols to allow. Pass to defer to OS/framework defaults. + /// + public NtlmHttpClientFactory(SslProtocols sslProtocols) { + _sslProtocols = sslProtocols; + } + + public HttpClient CreateUnauthenticatedClient() { + var handler = new HttpClientHandler { + ServerCertificateCustomValidationCallback = + (httpRequestMessage, cert, cetChain, policyErrors) => true, + UseDefaultCredentials = false + }; + + if (_sslProtocols != SslProtocols.None) + handler.SslProtocols = _sslProtocols; + + return new HttpClient(handler); + } + + public HttpClient CreateAuthenticatedHttpClient(Uri Url, string authPackage = "Kerberos") { + var handler = new HttpClientHandler { + Credentials = new CredentialCache() { + { Url, authPackage, CredentialCache.DefaultNetworkCredentials } + }, + + PreAuthenticate = true, + ServerCertificateCustomValidationCallback = + (httpRequestMessage, cert, cetChain, policyErrors) => true, + }; + + if (_sslProtocols != SslProtocols.None) + handler.SslProtocols = _sslProtocols; + + return new HttpClient(handler); + } +} \ No newline at end of file diff --git a/src/CommonLib/Processors/CAEnrollmentProcessor.cs b/src/CommonLib/Processors/CAEnrollmentProcessor.cs index 753e33215..a45c12afb 100644 --- a/src/CommonLib/Processors/CAEnrollmentProcessor.cs +++ b/src/CommonLib/Processors/CAEnrollmentProcessor.cs @@ -7,6 +7,7 @@ using System.Net; using System.Net.Http; using System.Net.Sockets; +using System.Security.Authentication; using System.Threading.Tasks; namespace SharpHoundCommonLib.Processors { @@ -18,13 +19,11 @@ public class CAEnrollmentProcessor { private readonly string _caName; private readonly ILogger _logger; - public CAEnrollmentProcessor(string caDnsHostname, string caName, ILogger log = null) { - ServicePointManager.SecurityProtocol |= - SecurityProtocolType.Ssl3 - | SecurityProtocolType.Tls12 - | SecurityProtocolType.Tls11 - | SecurityProtocolType.Tls; + // TLS1.3 is not available in .Net Framework 4.7.2, but the enum can still be assigned. + private const SslProtocols CaEnrollmentSslProtocols = + SslProtocols.Ssl3 | SslProtocols.Tls | SslProtocols.Tls11 | SslProtocols.Tls12 | (SslProtocols)12288; + public CAEnrollmentProcessor(string caDnsHostname, string caName, ILogger log = null) { _caDnsHostname = caDnsHostname; _caName = caName; _logger = log ?? Logging.LogProvider.CreateLogger("CAEnrollmentProcessor"); @@ -48,7 +47,7 @@ await Task.WhenAll( } catch (Exception ex) { _logger.LogError(ex, "An error occurred while scanning enrollment endpoints"); } - + endpoints = TagEndpoints(endpoints).ToList(); return endpoints; @@ -59,7 +58,7 @@ private IEnumerable> TagEndpoints(IEnumerable>> private async Task> GetNtlmEndpoint(Uri url, bool? useBadChannelBinding, CAEnrollmentEndpointType type, CAEnrollmentEndpointScanResult scanResult) { var authService = new HttpNtlmAuthenticationService( - new HttpClientFactory() + new NtlmHttpClientFactory(CaEnrollmentSslProtocols) ); var output = new CAEnrollmentEndpoint(url, type, scanResult); @@ -228,4 +227,4 @@ private async Task> GetNtlmEndpoint(Uri url, boo } } } -} \ No newline at end of file +} diff --git a/src/CommonLib/Processors/DomainTrustProcessor.cs b/src/CommonLib/Processors/DomainTrustProcessor.cs index 8047ae9fb..c41446054 100644 --- a/src/CommonLib/Processors/DomainTrustProcessor.cs +++ b/src/CommonLib/Processors/DomainTrustProcessor.cs @@ -1,6 +1,5 @@ -using System.Collections.Generic; +using System.Collections.Generic; using System.DirectoryServices.Protocols; -using System.Linq; using System.Security.Principal; using Microsoft.Extensions.Logging; using SharpHoundCommonLib.Enums; @@ -29,16 +28,9 @@ public async IAsyncEnumerable EnumerateDomainTrusts(string domain) { _log.LogDebug("Running trust enumeration for {Domain}", domain); - // Attempt to get trust type - var trustInfoList = new List<(string TargetName, System.DirectoryServices.ActiveDirectory.TrustType TrustType)>(); - try { - if (_utils.GetDomain(domain, out var domainObject)) { - trustInfoList.AddRange(from System.DirectoryServices.ActiveDirectory.TrustRelationshipInformation trust in domainObject.GetAllTrustRelationships() - select (trust.TargetName, trust.TrustType)); - } - } - catch { - _log.LogWarning("Trust type enumeration using non-LDAP for {Domain} failed", domain); + var trustTypes = new Dictionary(System.StringComparer.OrdinalIgnoreCase); + if (_utils.GetDomain(domain, out var domainInfo)) { + trustTypes = domainInfo.TrustTypes; } await foreach (var result in _utils.Query(new LdapQueryParameters { @@ -91,7 +83,7 @@ public async IAsyncEnumerable EnumerateDomainTrusts(string domain) trust.IsTransitive = !attributes.HasFlag(TrustAttributes.NonTransitive); if (entry.TryGetProperty(LDAPProperties.CanonicalName, out var cn)) { - trust.TargetDomainName = cn.ToUpper(); + trust.TargetDomainName = cn.ToUpperInvariant(); } trust.SidFilteringEnabled = @@ -104,9 +96,16 @@ public async IAsyncEnumerable EnumerateDomainTrusts(string domain) (attributes.HasFlag(TrustAttributes.WithinForest) || attributes.HasFlag(TrustAttributes.CrossOrganizationEnableTGTDelegation)); - var match = trustInfoList.FirstOrDefault(t => - t.TargetName.ToUpper().Equals(trust.TargetDomainName)); - trust.TrustType = !string.IsNullOrEmpty(match.TargetName) ? (TrustType) match.TrustType : TrustAttributesToType(attributes); + if (trust.TargetDomainName != null && trustTypes.TryGetValue(trust.TargetDomainName, out var classifiedType) && + classifiedType != TrustType.Unknown) { + trust.TrustType = classifiedType; + } + else if (entry.TryGetLongProperty(LDAPProperties.TrustType, out var ldapTrustType) && ldapTrustType == 3) { + trust.TrustType = TrustType.Kerberos; + } + else { + trust.TrustType = TrustAttributesToType(attributes); + } yield return trust; } @@ -114,19 +113,9 @@ public async IAsyncEnumerable EnumerateDomainTrusts(string domain) public static TrustType TrustAttributesToType(TrustAttributes attributes) { - TrustType trustType; - - if (attributes.HasFlag(TrustAttributes.WithinForest)) - trustType = TrustType.ParentChild; - else if (attributes.HasFlag(TrustAttributes.ForestTransitive)) - trustType = TrustType.Forest; - else if (!attributes.HasFlag(TrustAttributes.WithinForest) && - !attributes.HasFlag(TrustAttributes.ForestTransitive)) - trustType = TrustType.External; - else - trustType = TrustType.Unknown; - - return trustType; + if (attributes.HasFlag(TrustAttributes.WithinForest)) return TrustType.Unknown; + if (attributes.HasFlag(TrustAttributes.ForestTransitive)) return TrustType.Forest; + return TrustType.External; } } } diff --git a/src/CommonLib/Processors/GPOLocalGroupProcessor.cs b/src/CommonLib/Processors/GPOLocalGroupProcessor.cs index 28a6996f6..7fba2bc3f 100644 --- a/src/CommonLib/Processors/GPOLocalGroupProcessor.cs +++ b/src/CommonLib/Processors/GPOLocalGroupProcessor.cs @@ -1,4 +1,4 @@ -using System; +using System; using System.Collections.Concurrent; using System.Collections.Generic; using System.DirectoryServices.Protocols; @@ -66,7 +66,7 @@ public async Task ReadGPOLocalGroups(string gpLink, string string domain; //If our dn is null, use our default domain if (string.IsNullOrEmpty(distinguishedName)) { - if (!_utils.GetDomain(out var domainResult)) { + if (!_utils.GetDomain(out var domainResult) || string.IsNullOrWhiteSpace(domainResult?.Name)) { return ret; } diff --git a/src/CommonLib/README.md b/src/CommonLib/README.md index 831dddbfc..144bf1305 100644 --- a/src/CommonLib/README.md +++ b/src/CommonLib/README.md @@ -35,6 +35,40 @@ You may optionally provide an `ILogger` and a pre-created `Cache` instance to `C - Registry collection orchestration via `RegistryProcessor` - User rights, SPN, and certificate-related processing helpers +## Domain resolution metadata + +`SharpHoundCommonLib.Models.LdapDomainInfo` holds plain domain metadata: domain and forest names, the domain SID, naming contexts, the PDC hostname, `DomainControllerNames`, and `TrustTypes`. Additional strings may be null; collections start empty. `TrustTypes` compares target domain names case-insensitively. + +Controlled LDAP discovery includes both writable and read-only domain controllers in `DomainControllerNames`, providing candidates for the connection pool's controller fallback strategy. + +All `GetDomain` overloads now return `LdapDomainInfo` through the synchronous `bool`/`out` pattern instead of framework `Domain` objects. Success requires a resolved name and default naming context. Callers should use the returned naming contexts and check optional metadata before using it. Successful results from controlled LDAP and opted-in framework fallback are equally valid and cached per `LdapUtils` instance, case-insensitively; `SetLdapConfig` and `ResetUtils` clear that cache. Static calls are not cached. + +Failed metadata reads are retried by subsequent instance `GetDomain` calls after a 30-second backoff. Successful reads, including empty results or an `Unknown` trust classification, remain cached. Topology recovery recomputes classifications from the cached trust records. Refresh uses the resolution path that supplied the cached identity and verifies that identity; controlled refresh stays on the original endpoint and does not invoke framework fallback. Framework results also retry failed forest and naming-context reads. Failures preserve the previous metadata. Refresh publishes a new `LdapDomainInfo` snapshot, so callers must call `GetDomain` again to receive recovered metadata. Previously returned snapshots remain unchanged by refresh. + +`LdapConfig.AllowUncontrolledDomainFallback` defaults to `false` and appears in configuration logging. `GetDomain` attempts controlled LDAP first and permits legacy framework resolution only after core identity resolution fails and this flag is enabled. Authentication rejection (`InvalidCredentials` or `InappropriateAuthentication`) on either controlled transport stops resolution without legacy fallback, regardless of `ForceSSL`. Fallback use is logged and may ignore LDAP settings. Successful controlled results with unavailable additional metadata never trigger legacy enrichment. + +`LdapConfig.UserDomain` declares the DNS or NetBIOS domain associated with the user's credentials. The internal controlled resolver selects its endpoint in this order: `Server`, the supplied domain argument, `UserDomain`, then `USERDNSDOMAIN`. Null, empty, or whitespace hints are ignored. The resolved identity comes from LDAP; the hint does not restrict collection to the credential domain. + +For reliable `/netonly` use, supply an explicit `Server` or domain argument to `GetDomain` and leave `Username` unset so LDAP binding can use ambient outbound credentials. For example, run the calling application under `runas /netonly` and resolve the target through the static overload: + +```csharp +var config = new LdapConfig { + Server = "dc.child.example.test", + ForceSSL = true +}; + +if (LdapUtils.GetDomain("child.example.test", config, out var domain)) { + // domain.Name and domain.DefaultNamingContext come from the target's LDAP response. + var searchBase = domain.DefaultNamingContext; +} +``` + +`UserDomain` defaults to null and appears in configuration logging. It guides endpoint selection without changing credentials or the Windows authentication context. For reliable `/netonly` targeting, supply `Server` or a domain argument to `GetDomain`; `USERDNSDOMAIN` is only a last-resort target hint. + +Controlled resolution tries SSL first, using `SSLPort` (default 636). An SSL operation failure permits a retry on the same endpoint using `Port` (default 389) only when `ForceSSL` is false. Invalid credentials or inappropriate authentication fail resolution without a transport retry. It preserves `AuthType`, enables signing and sealing on plaintext connections unless `DisableSigning` is set, and validates certificates unless `DisableCertVerification` is set. A configured `Username` supplies explicit credentials instead of ambient credentials. + +With `Server` configured, every resolver read stays on that host with referrals and automatic reconnection disabled. Discovered PDC and controller names are returned as metadata. A supplied DNS domain or NetBIOS alias must match the target's advertised identity; a mismatch fails controlled resolution. DNS matches, including single-label names, are accepted before requiring a NetBIOS cross-reference. Unavailable SID, PDC, controller, or trust metadata preserves successful core resolution without changing endpoints or invoking legacy enrichment. + ## Relationship to SharpHoundRPC `SharpHoundCommon` depends on `SharpHoundRPC` and is intended to be the higher-level entry point. Most consumers should not reference `SharpHoundRPC` directly unless they need its lower-level SAM, LSA, NetAPI, or registry APIs. @@ -43,4 +77,4 @@ You may optionally provide an `ILogger` and a pre-created `Cache` instance to `C - Source: https://github.com/SpecterOps/SharpHoundCommon - Issues: https://github.com/SpecterOps/SharpHoundCommon/issues -- License: GPL-3.0-only \ No newline at end of file +- License: GPL-3.0-only diff --git a/test/unit/CommonLibHelperTests.cs b/test/unit/CommonLibHelperTests.cs index 8ba4b0ed2..ff9511e7c 100644 --- a/test/unit/CommonLibHelperTests.cs +++ b/test/unit/CommonLibHelperTests.cs @@ -1,4 +1,6 @@ using System; +using System.Globalization; +using System.Linq; using System.Runtime.Versioning; using System.Text; using System.Threading.Tasks; @@ -26,53 +28,52 @@ public void RemoveDistinguishedNamePrefix_ExpectedResult() { [Fact] public void SplitGPLinkProperty_ValidPropFilterEnabled_ExpectedResult() { - var isPropFilterEnabled = false; - //TODO: Ari, proper test string? - var testGPLinkProperty = - "[LDAP:/o=foo/ou=foo Group (ABC123)/cn=foouser (blah)123; SIP:foouser@example.co.uk; smtp:foouser@sub1.example.co.uk; smtp:foouser@sub2.example.co.uk; SMTP:foouser@example.co.uk][]"; + // Filter is ON (filterDisabled = true): an enabled link (status 0) must pass through. + const string testGPLinkProperty = + "[LDAP://CN={94DD0260-38B5-497E-8876-10E7A96E80D0},CN=Policies,CN=System,DC=testlab,DC=local;0]"; - var res = Helpers.SplitGPLinkProperty(testGPLinkProperty, isPropFilterEnabled); + var res = Helpers.SplitGPLinkProperty(testGPLinkProperty, filterDisabled: true).ToList(); - foreach (var parsedGPLink in res) - Assert.Equal("cn=foouser (blah)123", parsedGPLink.DistinguishedName); - // TODO: issue here with test data? Assert.Equal("1", parsedGPLink.Status); + Assert.Single(res); + Assert.Equal("CN={94DD0260-38B5-497E-8876-10E7A96E80D0},CN=Policies,CN=System,DC=testlab,DC=local", + res[0].DistinguishedName); + Assert.Equal("0", res[0].Status); } [Fact] public void SplitGPLinkProperty_ValidPropFilterDisabled_ExpectedResult() { - var isPropFilterEnabled = false; - //TODO: Ari, proper test string? - var testGPLinkProperty = - "[LDAP:/o=foo/ou=foo Group (ABC123)/cn=foouser (blah)123; SIP:foouser@example.co.uk; smtp:foouser@sub1.example.co.uk; smtp:foouser@sub2.example.co.uk; SMTP:foouser@example.co.uk][]"; + // Filter is OFF (filterDisabled = false): a disabled link (status 1) must still come through. + const string testGPLinkProperty = + "[LDAP://CN={94DD0260-38B5-497E-8876-10E7A96E80D0},CN=Policies,CN=System,DC=testlab,DC=local;1]"; - var res = Helpers.SplitGPLinkProperty(testGPLinkProperty, isPropFilterEnabled); + var res = Helpers.SplitGPLinkProperty(testGPLinkProperty, filterDisabled: false).ToList(); - foreach (var parsedGPLink in res) - Assert.Equal("cn=foouser (blah)123", parsedGPLink.DistinguishedName); - // TODO: issue here with test data? Assert.Equal("1", parsedGPLink.Status); + Assert.Single(res); + Assert.Equal("CN={94DD0260-38B5-497E-8876-10E7A96E80D0},CN=Policies,CN=System,DC=testlab,DC=local", + res[0].DistinguishedName); + Assert.Equal("1", res[0].Status); } - /// [Fact] - public void SplitGPLinkProperty_PropWithUnsupportedDelimiter_FilterEnabled_ExpectedResult() { - var isPropFilterEnabled = true; - //TODO: Ari, proper test string? - var testGPLinkProperty = - "[LDAP:/o=foo/ou=foo Group (ABC123)/cn=foouser (blah)123; DC=somedomainName; SIP:foouser@example.co.uk; smtp:foouser@sub1.example.co.uk; smtp:foouser@sub2.example.co.uk; SMTP:foouser@example.co.uk][]"; + public void SplitGPLinkProperty_MixedStatuses_FilterEnabled_OnlyEnabledLinksReturned() { + // Filter is ON: a multi-link property containing both enabled (status 0) and disabled + // (status 1) links must yield only the enabled one. + const string testGPLinkProperty = + "[LDAP://CN={94DD0260-38B5-497E-8876-10E7A96E80D0},CN=Policies,CN=System,DC=testlab,DC=local;0]" + + "[LDAP://CN={C52F168C-CD05-4487-B405-564934DA8EFF},CN=Policies,CN=System,DC=testlab,DC=local;1]"; - var res = Helpers.SplitGPLinkProperty(testGPLinkProperty, isPropFilterEnabled); + var res = Helpers.SplitGPLinkProperty(testGPLinkProperty, filterDisabled: true).ToList(); - foreach (var parsedGPLink in res) - Assert.Equal("cn=foouser (blah)123", parsedGPLink.DistinguishedName); - // TODO: issue here with test data? Assert.Equal("1", parsedGPLink.Status); + Assert.Single(res); + Assert.Equal("CN={94DD0260-38B5-497E-8876-10E7A96E80D0},CN=Policies,CN=System,DC=testlab,DC=local", + res[0].DistinguishedName); + Assert.Equal("0", res[0].Status); } [Fact] public void SplitGPLinkProperty_InValidPropFilterDisabled_ExpectedResult() { - var isPropFilterEnabled = false; - //TODO: Ari, proper test string? - var testGPLinkProperty = "/*obviously wrong data*/"; - var res = Helpers.SplitGPLinkProperty(testGPLinkProperty, isPropFilterEnabled); + const string testGPLinkProperty = "/*obviously wrong data*/"; + var res = Helpers.SplitGPLinkProperty(testGPLinkProperty, filterDisabled: false); Assert.Empty(res); } @@ -217,6 +218,21 @@ public void ConvertFileTimeToUnixEpoch_InvalidTimestamp_FormatException() { Assert.Equal("The input string '-201adsfasf12180244' was not in a correct format.", ex.Message); } + [Fact] + public void DistinguishedNameToDomain_TurkishCulture_UsesInvariantCasing() { + var originalCulture = CultureInfo.CurrentCulture; + try { + CultureInfo.CurrentCulture = CultureInfo.GetCultureInfo("tr-TR"); + + var result = Helpers.DistinguishedNameToDomain("DC=child,DC=example,DC=test"); + + Assert.Equal("CHILD.EXAMPLE.TEST", result); + } + finally { + CultureInfo.CurrentCulture = originalCulture; + } + } + [Fact] public void DistinguishedNameToDomain_RegularObject_CorrectDomain() { var result = Helpers.DistinguishedNameToDomain( @@ -234,6 +250,46 @@ public void DistinguishedNameToDomain_DeletedObjects_CorrectDomain() { Assert.Equal("TESTLAB.LOCAL", result); } + [Fact] + public void DistinguishedNameToDomain_CNValueStartsWithDCEquals_ReturnsTrailingDCComponents() { + // A naive IndexOf("DC=") would hit the "DC=" inside the CN value and return + // the wrong domain. The correct result uses only the trailing DC= RDNs. + var result = Helpers.DistinguishedNameToDomain("CN=DC=proxy,CN=Users,DC=corp,DC=com"); + Assert.Equal("CORP.COM", result); + } + + [Fact] + public void DistinguishedNameToDomain_DeletedDNSZoneWithLeadingDCRdn_ReturnsCorrectDomain() { + // Deleted DNS zone objects have a DC= attribute type on the first RDN. + // That leading DC= must not be included in the domain result. + var result = Helpers.DistinguishedNameToDomain( + @"DC=_msdcs.corp.com\0ADEL:guid,CN=Deleted Objects,DC=corp,DC=com"); + Assert.Equal("CORP.COM", result); + } + + [Fact] + public void DistinguishedNameToDomain_EscapedCommaInCN_ReturnsCorrectDomain() { + // The escaped comma inside the CN value must not be used as an RDN separator. + var result = Helpers.DistinguishedNameToDomain(@"CN=Smith\, John,OU=Sales,DC=corp,DC=com"); + Assert.Equal("CORP.COM", result); + } + + [Fact] + public void DistinguishedNameToDomain_EscapedCommaInOUValueFollowedByDCEquals_ReturnsCorrectDomain() { + // The fragment after the escaped comma starts with "DC=" which would cause a naive + // Split(',') to misidentify it as a domain component. The unescaped-comma split + // must keep the whole OU value together so only the true trailing DC= RDNs are used. + var result = Helpers.DistinguishedNameToDomain(@"CN=User,OU=Dept\, DC=Proxy,DC=corp,DC=com"); + Assert.Equal("CORP.COM", result); + } + + [Fact] + public void DistinguishedNameToDomain_DomainOnlyDN_ReturnsCorrectDomain() { + // A DN that consists solely of DC= components (e.g. as stored in RootDSE). + var result = Helpers.DistinguishedNameToDomain("DC=corp,DC=com"); + Assert.Equal("CORP.COM", result); + } + [Fact] public void ConvertTimestampToUnixEpoch_ValidTimestamp() { var d = DateTime.Parse("2025-04-07T00:00:00.0000000-07:00"); diff --git a/test/unit/CommonLibTest.csproj b/test/unit/CommonLibTest.csproj index 8910d6da3..1679bd4d0 100644 --- a/test/unit/CommonLibTest.csproj +++ b/test/unit/CommonLibTest.csproj @@ -25,6 +25,7 @@ + diff --git a/test/unit/DomainTrustProcessorTest.cs b/test/unit/DomainTrustProcessorTest.cs index 2ada97fc2..3c0cb99e1 100644 --- a/test/unit/DomainTrustProcessorTest.cs +++ b/test/unit/DomainTrustProcessorTest.cs @@ -1,6 +1,7 @@ -using System; +using System; using System.Collections.Generic; using System.DirectoryServices.Protocols; +using System.Globalization; using System.Linq; using System.Runtime.Versioning; using System.Threading; @@ -9,6 +10,7 @@ using Moq; using SharpHoundCommonLib; using SharpHoundCommonLib.Enums; +using SharpHoundCommonLib.Models; using SharpHoundCommonLib.Processors; using Xunit; using Xunit.Abstractions; @@ -51,7 +53,7 @@ public async Task DomainTrustProcessor_EnumerateDomainTrusts_HappyPath() Assert.Equal("EXTERNAL.LOCAL", trust.TargetDomainName); Assert.Equal("S-1-5-21-3084884204-958224920-2707782874", trust.TargetDomainSid); Assert.True(trust.IsTransitive); - Assert.Equal(TrustType.ParentChild, trust.TrustType); + Assert.Equal(TrustType.Unknown, trust.TrustType); Assert.True(trust.SidFilteringEnabled); } @@ -108,7 +110,7 @@ public void DomainTrustProcessor_TrustAttributesToType() { var attrib = TrustAttributes.WithinForest; var test = DomainTrustProcessor.TrustAttributesToType(attrib); - Assert.Equal(TrustType.ParentChild, test); + Assert.Equal(TrustType.Unknown, test); attrib = TrustAttributes.ForestTransitive; test = DomainTrustProcessor.TrustAttributesToType(attrib); @@ -126,5 +128,88 @@ public void DomainTrustProcessor_TrustAttributesToType() test = DomainTrustProcessor.TrustAttributesToType(attrib); Assert.Equal(TrustType.External, test); } + + [Theory] + [InlineData(TrustType.ParentChild)] + [InlineData(TrustType.TreeRoot)] + [InlineData(TrustType.CrossLink)] + [InlineData(TrustType.External)] + [InlineData(TrustType.Forest)] + [InlineData(TrustType.Kerberos)] + [InlineData(TrustType.Unknown)] + public async Task EnumerateDomainTrusts_PreservesResolvedClassification(TrustType classification) { + var domain = new LdapDomainInfo { Name = "testlab.local", DefaultNamingContext = "DC=testlab,DC=local" }; + domain.TrustTypes["external.local"] = classification; + var utils = new Mock(); + utils.Setup(x => x.GetDomain("testlab.local", out domain)).Returns(true); + utils.Setup(x => x.Query(It.IsAny(), It.IsAny())) + .Returns(new[] { CreateTrustEntry(2) }.ToAsyncEnumerable); + var trusts = await new DomainTrustProcessor(utils.Object).EnumerateDomainTrusts("testlab.local").ToArrayAsync(); + Assert.Equal(classification, Assert.Single(trusts).TrustType); + } + + [SupportedOSPlatform("windows")] + [WindowsOnlyFact] + public async Task EnumerateDomainTrusts_TurkishCulture_PreservesResolvedClassification() { + var originalCulture = CultureInfo.CurrentCulture; + try { + CultureInfo.CurrentCulture = CultureInfo.GetCultureInfo("tr-TR"); + var domain = new LdapDomainInfo { Name = "testlab.local", DefaultNamingContext = "DC=testlab,DC=local" }; + domain.TrustTypes["child.test"] = TrustType.ParentChild; + var utils = new Mock(); + utils.Setup(x => x.GetDomain("testlab.local", out domain)).Returns(true); + utils.Setup(x => x.Query(It.IsAny(), It.IsAny())) + .Returns(new[] { CreateTrustEntry(2, targetDomainName: "child.test") }.ToAsyncEnumerable); + + var trusts = await new DomainTrustProcessor(utils.Object).EnumerateDomainTrusts("testlab.local").ToArrayAsync(); + + var trust = Assert.Single(trusts); + Assert.Equal(TrustType.ParentChild, trust.TrustType); + Assert.Equal("CHILD.TEST", trust.TargetDomainName); + } + finally { + CultureInfo.CurrentCulture = originalCulture; + } + } + + [Theory] + [InlineData(2, TrustAttributes.ForestTransitive, TrustType.Forest)] + [InlineData(2, (TrustAttributes)0, TrustType.External)] + [InlineData(1, TrustAttributes.QuarantinedDomain, TrustType.External)] + [InlineData(3, TrustAttributes.WithinForest, TrustType.Kerberos)] + [InlineData(2, TrustAttributes.WithinForest, TrustType.Unknown)] + public async Task EnumerateDomainTrusts_UnknownClassificationUsesReadableTrustFields( + int ldapType, TrustAttributes attributes, TrustType expected) { + var domain = new LdapDomainInfo { Name = "testlab.local", DefaultNamingContext = "DC=testlab,DC=local" }; + domain.TrustTypes["external.local"] = TrustType.Unknown; + var utils = new Mock(); + utils.Setup(x => x.GetDomain("testlab.local", out domain)).Returns(true); + utils.Setup(x => x.Query(It.IsAny(), It.IsAny())) + .Returns(new[] { CreateTrustEntry(ldapType, attributes) }.ToAsyncEnumerable); + var trusts = await new DomainTrustProcessor(utils.Object).EnumerateDomainTrusts("testlab.local").ToArrayAsync(); + Assert.Equal(expected, Assert.Single(trusts).TrustType); + } + + [Theory] + [InlineData(1, TrustType.Unknown)] + [InlineData(2, TrustType.Unknown)] + [InlineData(3, TrustType.Kerberos)] + public async Task EnumerateDomainTrusts_MissingTopologyDoesNotGuessParentChild(int ldapType, TrustType expected) { + var utils = new Mock(); + utils.Setup(x => x.Query(It.IsAny(), It.IsAny())) + .Returns(new[] { CreateTrustEntry(ldapType) }.ToAsyncEnumerable); + var trusts = await new DomainTrustProcessor(utils.Object).EnumerateDomainTrusts("testlab.local").ToArrayAsync(); + Assert.Equal(expected, Assert.Single(trusts).TrustType); + } + + private static LdapResult CreateTrustEntry(int ldapType, + TrustAttributes attributes = TrustAttributes.WithinForest, string targetDomainName = "EXTERNAL.LOCAL") => + LdapResult.Ok(new MockDirectoryObject("", new Dictionary { + ["trustdirection"] = "3", + ["trusttype"] = ldapType.ToString(), + ["trustattributes"] = ((int)attributes).ToString(), + ["cn"] = targetDomainName, + ["securityidentifier"] = Utils.B64ToBytes("AQQAAAAAAAUVAAAA7JjftxhaHTnafGWh") + })); } -} \ No newline at end of file +} diff --git a/test/unit/Facades/MockLdapUtils.cs b/test/unit/Facades/MockLdapUtils.cs index 1e5f1194d..b092ad5e7 100644 --- a/test/unit/Facades/MockLdapUtils.cs +++ b/test/unit/Facades/MockLdapUtils.cs @@ -1,4 +1,4 @@ -using System; +using System; using System.Collections.Concurrent; using System.Collections.Generic; using System.Diagnostics.CodeAnalysis; @@ -12,7 +12,7 @@ using SharpHoundCommonLib; using SharpHoundCommonLib.Enums; using SharpHoundCommonLib.OutputTypes; -using Domain = System.DirectoryServices.ActiveDirectory.Domain; +using SharpHoundCommonLib.Models; #pragma warning disable CS1998 // Async method lacks 'await' operators and will run synchronously namespace CommonLibTest.Facades @@ -693,12 +693,12 @@ public virtual IAsyncEnumerable> RangedRetrieval(string distingui return (false, default); } - public bool GetDomain(string domainName, out Domain domain) { + public bool GetDomain(string domainName, out LdapDomainInfo domain) { domain = null; return false; } - public bool GetDomain(out Domain domain) { + public bool GetDomain(out LdapDomainInfo domain) { domain = null; return false; } diff --git a/test/unit/Facades/MockableDomain.cs b/test/unit/Facades/MockableDomain.cs deleted file mode 100644 index c28bd093c..000000000 --- a/test/unit/Facades/MockableDomain.cs +++ /dev/null @@ -1,17 +0,0 @@ -using System.Diagnostics.CodeAnalysis; -using System.DirectoryServices.ActiveDirectory; - -namespace CommonLibTest.Facades -{ - [SuppressMessage("Interoperability", "CA1416:Validate platform compatibility")] - public class MockableDomain - { - public static Domain Construct(string domainName) - { - var domain = FacadeHelpers.GetUninitializedObject(); - FacadeHelpers.SetField(domain, "partitionName", domainName); - - return domain; - } - } -} \ No newline at end of file diff --git a/test/unit/GPOLocalGroupProcessorTest.cs b/test/unit/GPOLocalGroupProcessorTest.cs index 07107a12e..c1207a44b 100644 --- a/test/unit/GPOLocalGroupProcessorTest.cs +++ b/test/unit/GPOLocalGroupProcessorTest.cs @@ -147,7 +147,7 @@ public async Task GPOLocalGroupProcessor_ReadGPOLocalGroups_Null_Gpcfilesyspath( var processor = new GPOLocalGroupProcessor(mockLDAPUtils.Object); var testGPLinkProperty = - "[LDAP:/o=foo/ou=foo Group (ABC123)/cn=foouser (blah)123/dc=somedomain;0;][LDAP:/o=foo/ou=foo Group (ABC123)/cn=foouser (blah)123/dc=someotherdomain;2;]"; + "[LDAP://CN={ECAD920E-8EB1-4E31-A80E-DD36367F81F4},CN=Policies,CN=System,DC=somedomain;0][LDAP://CN={ECAD920E-8EB1-4E31-A80E-DD36367F81F4},CN=Policies,CN=System,DC=someotherdomain;2]"; var result = await processor.ReadGPOLocalGroups(testGPLinkProperty, "DC=Testlab,DC=Local"); Assert.Single(result.AffectedComputers); @@ -234,7 +234,7 @@ public async Task GPOLocalGroupProcessor_ReadGPOLocalGroups_Does_Not_Skip_Enable y.SearchScope.Equals(SearchScope.Base) && y.Attributes.Contains(LDAPProperties.GPCFileSYSPath) && y.Attributes.Contains(LDAPProperties.Flags) && - y.SearchBase.Equals("cn=foouser (blah)123/dc=somedomain", StringComparison.OrdinalIgnoreCase) && + y.SearchBase.Equals("CN={ECAD920E-8EB1-4E31-A80E-DD36367F81F4},CN=Policies,CN=System,DC=somedomain", StringComparison.OrdinalIgnoreCase) && y.DomainName.Equals("somedomain", StringComparison.OrdinalIgnoreCase)), It.IsAny())) .Returns(result0.ToAsyncEnumerable); @@ -245,7 +245,7 @@ public async Task GPOLocalGroupProcessor_ReadGPOLocalGroups_Does_Not_Skip_Enable y.SearchScope.Equals(SearchScope.Base) && y.Attributes.Contains(LDAPProperties.GPCFileSYSPath) && y.Attributes.Contains(LDAPProperties.Flags) && - y.SearchBase.Equals("cn=foouser (blah)123/dc=someotherdomain", StringComparison.OrdinalIgnoreCase) && + y.SearchBase.Equals("CN={ECAD920E-8EB1-4E31-A80E-DD36367F81F4},CN=Policies,CN=System,DC=someotherdomain", StringComparison.OrdinalIgnoreCase) && y.DomainName.Equals("someotherdomain", StringComparison.OrdinalIgnoreCase)), It.IsAny())) .Returns(result1.ToAsyncEnumerable); @@ -256,7 +256,7 @@ public async Task GPOLocalGroupProcessor_ReadGPOLocalGroups_Does_Not_Skip_Enable y.SearchScope.Equals(SearchScope.Base) && y.Attributes.Contains(LDAPProperties.GPCFileSYSPath) && y.Attributes.Contains(LDAPProperties.Flags) && - y.SearchBase.Equals("cn=foouser (blah)123/dc=somethirddomain", StringComparison.OrdinalIgnoreCase) && + y.SearchBase.Equals("CN={ECAD920E-8EB1-4E31-A80E-DD36367F81F4},CN=Policies,CN=System,DC=somethirddomain", StringComparison.OrdinalIgnoreCase) && y.DomainName.Equals("somethirddomain", StringComparison.OrdinalIgnoreCase)), It.IsAny())) .Returns(result2.ToAsyncEnumerable); @@ -267,16 +267,16 @@ public async Task GPOLocalGroupProcessor_ReadGPOLocalGroups_Does_Not_Skip_Enable y.SearchScope.Equals(SearchScope.Base) && y.Attributes.Contains(LDAPProperties.GPCFileSYSPath) && y.Attributes.Contains(LDAPProperties.Flags) && - y.SearchBase.Equals("cn=foouser (blah)123/dc=somefourthdomain", StringComparison.OrdinalIgnoreCase) && + y.SearchBase.Equals("CN={ECAD920E-8EB1-4E31-A80E-DD36367F81F4},CN=Policies,CN=System,DC=somefourthdomain", StringComparison.OrdinalIgnoreCase) && y.DomainName.Equals("somefourthdomain", StringComparison.OrdinalIgnoreCase)), It.IsAny())) .Returns(result3.ToAsyncEnumerable); - + var processor = new GPOLocalGroupProcessor(mockLDAPUtils.Object); - var testGPLinkProperty0 = "[LDAP:/o=foo/ou=foo Group (ABC123)/cn=foouser (blah)123/dc=somedomain;0;]"; - var testGPLinkProperty1 = "[LDAP:/o=foo/ou=foo Group (ABC123)/cn=foouser (blah)123/dc=someotherdomain;0;]"; - var testGPLinkProperty2 = "[LDAP:/o=foo/ou=foo Group (ABC123)/cn=foouser (blah)123/dc=somethirddomain;0;]"; - var testGPLinkProperty3 = "[LDAP:/o=foo/ou=foo Group (ABC123)/cn=foouser (blah)123/dc=somefourthdomain;0;]"; + var testGPLinkProperty0 = "[LDAP://CN={ECAD920E-8EB1-4E31-A80E-DD36367F81F4},CN=Policies,CN=System,DC=somedomain;0]"; + var testGPLinkProperty1 = "[LDAP://CN={ECAD920E-8EB1-4E31-A80E-DD36367F81F4},CN=Policies,CN=System,DC=someotherdomain;0]"; + var testGPLinkProperty2 = "[LDAP://CN={ECAD920E-8EB1-4E31-A80E-DD36367F81F4},CN=Policies,CN=System,DC=somethirddomain;0]"; + var testGPLinkProperty3 = "[LDAP://CN={ECAD920E-8EB1-4E31-A80E-DD36367F81F4},CN=Policies,CN=System,DC=somefourthdomain;0]"; // Act var act0 = await processor.ReadGPOLocalGroups(testGPLinkProperty0, "DC=Testlab,DC=Local"); @@ -323,14 +323,16 @@ public async Task GPOLocalGroupProcessor_ReadGPOLocalGroups() { .Returns(mockComputerResults.ToAsyncEnumerable) .Returns(mockGCPFileSysPathResults.ToAsyncEnumerable) .Returns(Array.Empty>().ToAsyncEnumerable); - var domain = MockableDomain.Construct("TESTLAB.LOCAL"); + var domain = new SharpHoundCommonLib.Models.LdapDomainInfo { + Name = "TESTLAB.LOCAL", DefaultNamingContext = "DC=TESTLAB,DC=LOCAL" + }; mockLDAPUtils.Setup(x => x.GetDomain(out domain)).Returns(true); var processor = new GPOLocalGroupProcessor(mockLDAPUtils.Object); var testGPLinkProperty = - "[LDAP:/o=foo/ou=foo Group (ABC123)/cn=foouser (blah)123/dc=somedomain;0;][LDAP:/o=foo/ou=foo Group (ABC123)/cn=foouser (blah)123/dc=someotherdomain;2;]"; + "[LDAP://CN={ECAD920E-8EB1-4E31-A80E-DD36367F81F4},CN=Policies,CN=System,DC=somedomain;0][LDAP://CN={ECAD920E-8EB1-4E31-A80E-DD36367F81F4},CN=Policies,CN=System,DC=someotherdomain;2]"; var result = await processor.ReadGPOLocalGroups(testGPLinkProperty, null); //mockLDAPUtils.VerifyAll(); @@ -338,6 +340,17 @@ public async Task GPOLocalGroupProcessor_ReadGPOLocalGroups() { var actual = result.AffectedComputers.First(); Assert.Equal(Label.Computer, actual.ObjectType); Assert.Equal("teapot", actual.ObjectIdentifier); + mockLDAPUtils.Verify(x => x.Query(It.Is(q => q.DomainName == "TESTLAB.LOCAL"), + It.IsAny()), Times.AtLeastOnce); + } + + [Fact] + public async Task GPOLocalGroupProcessor_MissingDefaultDomainNameSkipsQueries() { + var utils = new Mock(MockBehavior.Strict); + var domain = new SharpHoundCommonLib.Models.LdapDomainInfo(); + utils.Setup(x => x.GetDomain(out domain)).Returns(true); + var result = await new GPOLocalGroupProcessor(utils.Object).ReadGPOLocalGroups("[LDAP://CN=Policy;0]", null); + Assert.Empty(result.AffectedComputers); } [Fact] diff --git a/test/unit/HttpNtlmAuthenticationServiceTest.cs b/test/unit/HttpNtlmAuthenticationServiceTest.cs index e78d654df..f5825e138 100644 --- a/test/unit/HttpNtlmAuthenticationServiceTest.cs +++ b/test/unit/HttpNtlmAuthenticationServiceTest.cs @@ -25,7 +25,7 @@ public void Dispose() { [Fact] public void HttpNtlmAuthenticationService_ExtractAuthSchemes_AuthNotRequiredException() { - var service = new HttpNtlmAuthenticationService(new HttpClientFactory(), null); + var service = new HttpNtlmAuthenticationService(new NtlmHttpClientFactory(), null); var httpResponseMessage = new HttpResponseMessage { StatusCode = HttpStatusCode.OK, }; @@ -37,7 +37,7 @@ public void HttpNtlmAuthenticationService_ExtractAuthSchemes_AuthNotRequiredExce [Fact] public void HttpNtlmAuthenticationService_ExtractAuthSchemes_HttpForbiddenException() { - var service = new HttpNtlmAuthenticationService(new HttpClientFactory(), null); + var service = new HttpNtlmAuthenticationService(new NtlmHttpClientFactory(), null); var httpResponseMessage = new HttpResponseMessage { StatusCode = HttpStatusCode.Forbidden, }; @@ -49,7 +49,7 @@ public void HttpNtlmAuthenticationService_ExtractAuthSchemes_HttpForbiddenExcept [Fact] public void HttpNtlmAuthenticationService_ExtractAuthSchemes_HttpServerErrorException() { - var service = new HttpNtlmAuthenticationService(new HttpClientFactory(), null); + var service = new HttpNtlmAuthenticationService(new NtlmHttpClientFactory(), null); var httpResponseMessage = new HttpResponseMessage { StatusCode = HttpStatusCode.InternalServerError, }; @@ -61,7 +61,7 @@ public void HttpNtlmAuthenticationService_ExtractAuthSchemes_HttpServerErrorExce [Fact] public void HttpNtlmAuthenticationService_ExtractAuthSchemes_Success() { - var service = new HttpNtlmAuthenticationService(new HttpClientFactory(), null); + var service = new HttpNtlmAuthenticationService(new NtlmHttpClientFactory(), null); var httpResponseMessage = new HttpResponseMessage(); httpResponseMessage.StatusCode = HttpStatusCode.Accepted; httpResponseMessage.Headers.WwwAuthenticate.Add( diff --git a/test/unit/LDAPUtilsTest.cs b/test/unit/LDAPUtilsTest.cs index d6d6514ec..149f52449 100644 --- a/test/unit/LDAPUtilsTest.cs +++ b/test/unit/LDAPUtilsTest.cs @@ -1,6 +1,9 @@ using System; using System.Collections.Generic; +using System.DirectoryServices; +using System.DirectoryServices.AccountManagement; using System.DirectoryServices.ActiveDirectory; +using System.Reflection; using System.Runtime.Versioning; using System.Threading.Tasks; using CommonLibTest.Facades; @@ -250,6 +253,363 @@ public async Task Test_ResolveSearchResult_TrustAccount() { Assert.False(result.Deleted); } + #region CreateDirectoryEntry Tests + + // --------------------------------------------------------------------------- + // Helpers + // --------------------------------------------------------------------------- + + /// + /// Invokes the static CreateDirectoryEntry method via reflection. + /// DirectoryEntry does not connect to the server until properties are accessed, + /// so the call succeeds even with a fake path. + /// + private static IDirectoryObject InvokeCreateDirectoryEntry(LdapUtils utils, string path) { + // Extract the LdapConfig from the LdapUtils instance via reflection. + var configField = typeof(LdapUtils).GetField("_ldapConfig", + BindingFlags.NonPublic | BindingFlags.Instance); + Assert.NotNull(configField); + var config = (LdapConfig)configField.GetValue(utils); + + return TestPrivateMethod.StaticMethod(typeof(Helpers), + "CreateDirectoryEntry", new object[] { path, config }); + } + + /// + /// Extracts the underlying DirectoryEntry from its DirectoryEntryWrapper so that + /// Path and AuthenticationType can be inspected without triggering a network call. + /// + private static DirectoryEntry ExtractDirectoryEntry(IDirectoryObject directoryObject) { + var entryField = directoryObject.GetType() + .GetField("_entry", BindingFlags.NonPublic | BindingFlags.Instance); + Assert.NotNull(entryField); + return (DirectoryEntry)entryField.GetValue(directoryObject); + } + + // --------------------------------------------------------------------------- + // Path construction – no server configured + // --------------------------------------------------------------------------- + + [Fact] + public void CreateDirectoryEntry_NoServer_PathIsUnchanged() { + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig()); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://domain.com"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal("LDAP://domain.com", entry.Path); + } + + // --------------------------------------------------------------------------- + // Path construction – server configured + // --------------------------------------------------------------------------- + + [Fact] + public void CreateDirectoryEntry_ServerSet_DefaultPort_InjectsServerWithoutPort() { + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { Server = "dc01.corp.com" }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://domain.com"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal("LDAP://dc01.corp.com/DC=domain,DC=com", entry.Path); + } + + [Fact] + public void CreateDirectoryEntry_ServerSet_CustomPort_InjectsServerWithPort() { + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { Server = "dc01.corp.com", Port = 3636 }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://domain.com"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal("LDAP://dc01.corp.com:3636/DC=domain,DC=com", entry.Path); + } + + [Fact] + public void CreateDirectoryEntry_ServerSet_ForceSSL_DefaultSSLPort_InjectsServerWithoutPort() { + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { Server = "dc01.corp.com", ForceSSL = true }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://domain.com"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal("LDAP://dc01.corp.com/DC=domain,DC=com", entry.Path); + } + + [Fact] + public void CreateDirectoryEntry_ServerSet_ForceSSL_CustomSSLPort_InjectsServerWithPort() { + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { Server = "dc01.corp.com", ForceSSL = true, SSLPort = 1636 }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://domain.com"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal("LDAP://dc01.corp.com:1636/DC=domain,DC=com", entry.Path); + } + + [Fact] + public void CreateDirectoryEntry_ServerSet_SIDPath_InjectsServerBeforeSIDMoniiker() { + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { Server = "dc01.corp.com" }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal("LDAP://dc01.corp.com/", entry.Path); + } + + [Fact] + public void CreateDirectoryEntry_ServerSet_RootDSEPath_InjectsServerBeforeSuffix() { + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { Server = "dc01.corp.com" }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://domain.com/RootDSE"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal("LDAP://dc01.corp.com/RootDSE", entry.Path); + } + + // --------------------------------------------------------------------------- + // Path construction – domain-shortcut edge cases + // --------------------------------------------------------------------------- + + [Fact] + public void CreateDirectoryEntry_ServerSet_SingleLabelDomain_ConvertedToSingleDCPart() { + // A domain name with no dots (e.g. a NetBIOS-style name) should become "DC=". + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { Server = "dc01.corp.com" }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://corp"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal("LDAP://dc01.corp.com/DC=corp", entry.Path); + } + + [Fact] + public void CreateDirectoryEntry_ServerSet_ExistingDNPath_ServerInjectedWithoutConversion() { + // When the path already carries a proper DN (contains '=') it must be forwarded + // verbatim after the server – no DC= conversion should be applied. + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { Server = "dc01.corp.com" }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://DC=domain,DC=com"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal("LDAP://dc01.corp.com/DC=domain,DC=com", entry.Path); + } + + // --------------------------------------------------------------------------- + // Path construction – double injection guard + // --------------------------------------------------------------------------- + + [Fact] + public void CreateDirectoryEntry_ServerSet_PathAlreadyHasServer_NotInjectedAgain() { + // If the path already begins with LDAP:/// the server must not be + // prepended a second time. + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { Server = "dc01.corp.com" }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://dc01.corp.com/DC=domain,DC=com"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal("LDAP://dc01.corp.com/DC=domain,DC=com", entry.Path); + } + + [Fact] + public void CreateDirectoryEntry_ServerSet_CustomPort_PathAlreadyHasServerWithPort_NotInjectedAgain() { + // Same guard when the path already carries the server with a non-default port. + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { Server = "dc01.corp.com", Port = 3636 }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://dc01.corp.com:3636/DC=domain,DC=com"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal("LDAP://dc01.corp.com:3636/DC=domain,DC=com", entry.Path); + } + + [Fact] + public void CreateDirectoryEntry_ServerSet_PathAlreadyHasServerWithRootDSE_NotInjectedAgain() { + // The guard must fire even when the path component after the server is not a DN + // (e.g. the special "RootDSE" target used by GetNamingContextPath). + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { Server = "dc01.corp.com" }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://dc01.corp.com/RootDSE"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal("LDAP://dc01.corp.com/RootDSE", entry.Path); + } + + // --------------------------------------------------------------------------- + // AuthenticationTypes – matching connection pool logic + // --------------------------------------------------------------------------- + + [Fact] + public void CreateDirectoryEntry_Default_HasSecureSigningAndSealing() { + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig()); // ForceSSL=false, DisableSigning=false + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://domain.com"); + var entry = ExtractDirectoryEntry(result); + + var expected = AuthenticationTypes.Secure | AuthenticationTypes.Signing | AuthenticationTypes.Sealing; + Assert.Equal(expected, entry.AuthenticationType); + } + + [Fact] + public void CreateDirectoryEntry_ForceSSL_HasSecureSSLOnly_NoSigningOrSealing() { + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { ForceSSL = true }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://domain.com"); + var entry = ExtractDirectoryEntry(result); + + var expected = AuthenticationTypes.Secure | AuthenticationTypes.SecureSocketsLayer; + Assert.Equal(expected, entry.AuthenticationType); + Assert.Equal(AuthenticationTypes.None, entry.AuthenticationType & AuthenticationTypes.Signing); + Assert.Equal(AuthenticationTypes.None, entry.AuthenticationType & AuthenticationTypes.Sealing); + } + + [Fact] + public void CreateDirectoryEntry_DisableSigning_HasSecureOnly() { + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { DisableSigning = true }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://domain.com"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal(AuthenticationTypes.Secure, entry.AuthenticationType); + } + + [Fact] + public void CreateDirectoryEntry_ForceSSLAndDisableSigning_HasSecureSSLOnly() { + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { ForceSSL = true, DisableSigning = true }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://domain.com"); + var entry = ExtractDirectoryEntry(result); + + var expected = AuthenticationTypes.Secure | AuthenticationTypes.SecureSocketsLayer; + Assert.Equal(expected, entry.AuthenticationType); + } + + // --------------------------------------------------------------------------- + // Credentials + // --------------------------------------------------------------------------- + + [Fact] + public void CreateDirectoryEntry_WithCredentials_UsernameAndPasswordApplied() { + var utils = new LdapUtils(); + utils.SetLdapConfig(new LdapConfig { Username = "testuser", Password = "testpass" }); + + var result = InvokeCreateDirectoryEntry(utils, "LDAP://domain.com"); + var entry = ExtractDirectoryEntry(result); + + Assert.Equal("testuser", entry.Username); + Assert.Equal("LDAP://domain.com", entry.Path); + Assert.Equal(AuthenticationTypes.Secure | AuthenticationTypes.Signing | AuthenticationTypes.Sealing, entry.AuthenticationType); + } + + #endregion + + #region BuildPrincipalContextParameters Tests + + // --------------------------------------------------------------------------- + // contextName – no server configured + // --------------------------------------------------------------------------- + + [Fact] + public void BuildPrincipalContextParameters_NoServer_NoDomain_ContextNameIsNull() { + var (contextName, _) = LdapUtils.BuildPrincipalContextParameters(new LdapConfig()); + Assert.Null(contextName); + } + + [Fact] + public void BuildPrincipalContextParameters_NoServer_DomainProvided_ContextNameIsDomain() { + var (contextName, _) = LdapUtils.BuildPrincipalContextParameters(new LdapConfig(), "testlab.local"); + Assert.Equal("testlab.local", contextName); + } + + // --------------------------------------------------------------------------- + // contextName – server configured + // --------------------------------------------------------------------------- + + [Fact] + public void BuildPrincipalContextParameters_ServerSet_DefaultPort_ContextNameIsServerOnly() { + var config = new LdapConfig { Server = "dc01.corp.com" }; + var (contextName, _) = LdapUtils.BuildPrincipalContextParameters(config); + Assert.Equal("dc01.corp.com", contextName); + } + + [Fact] + public void BuildPrincipalContextParameters_ServerSet_CustomPort_ContextNameIncludesPort() { + var config = new LdapConfig { Server = "dc01.corp.com", Port = 3636 }; + var (contextName, _) = LdapUtils.BuildPrincipalContextParameters(config); + Assert.Equal("dc01.corp.com:3636", contextName); + } + + [Fact] + public void BuildPrincipalContextParameters_ServerSet_DomainIgnoredWhenServerPresent() { + // The domain name should be ignored when a server is explicitly configured. + var config = new LdapConfig { Server = "dc01.corp.com" }; + var (contextName, _) = LdapUtils.BuildPrincipalContextParameters(config, "testlab.local"); + Assert.Equal("dc01.corp.com", contextName); + } + + [Fact] + public void BuildPrincipalContextParameters_ServerSet_ForceSSL_DefaultSSLPort_ContextNameIsServerOnly() { + var config = new LdapConfig { Server = "dc01.corp.com", ForceSSL = true }; + var (contextName, _) = LdapUtils.BuildPrincipalContextParameters(config); + Assert.Equal("dc01.corp.com", contextName); + } + + [Fact] + public void BuildPrincipalContextParameters_ServerSet_ForceSSL_CustomSSLPort_ContextNameIncludesPort() { + var config = new LdapConfig { Server = "dc01.corp.com", ForceSSL = true, SSLPort = 1636 }; + var (contextName, _) = LdapUtils.BuildPrincipalContextParameters(config); + Assert.Equal("dc01.corp.com:1636", contextName); + } + + // --------------------------------------------------------------------------- + // ContextOptions – mirroring connection pool's mutual-exclusion rule + // --------------------------------------------------------------------------- + + [Fact] + public void BuildPrincipalContextParameters_Default_HasNegotiateSigningAndSealing() { + var (_, options) = LdapUtils.BuildPrincipalContextParameters(new LdapConfig()); + var expected = ContextOptions.Negotiate | ContextOptions.Signing | ContextOptions.Sealing; + Assert.Equal(expected, options); + } + + [Fact] + public void BuildPrincipalContextParameters_ForceSSL_HasNegotiateAndSSLOnly_NoSigningOrSealing() { + var config = new LdapConfig { ForceSSL = true }; + var (_, options) = LdapUtils.BuildPrincipalContextParameters(config); + var expected = ContextOptions.Negotiate | ContextOptions.SecureSocketLayer; + Assert.Equal(expected, options); + Assert.Equal((ContextOptions)0, options & ContextOptions.Signing); + Assert.Equal((ContextOptions)0, options & ContextOptions.Sealing); + } + + [Fact] + public void BuildPrincipalContextParameters_DisableSigning_HasNegotiateOnly() { + var config = new LdapConfig { DisableSigning = true }; + var (_, options) = LdapUtils.BuildPrincipalContextParameters(config); + Assert.Equal(ContextOptions.Negotiate, options); + } + + [Fact] + public void BuildPrincipalContextParameters_ForceSSLAndDisableSigning_HasNegotiateAndSSLOnly() { + var config = new LdapConfig { ForceSSL = true, DisableSigning = true }; + var (_, options) = LdapUtils.BuildPrincipalContextParameters(config); + var expected = ContextOptions.Negotiate | ContextOptions.SecureSocketLayer; + Assert.Equal(expected, options); + } + + #endregion + [Fact] public async Task Test_ResolveHostToSid_BlankHost() { var spn = "MSSQLSvc/:1433"; diff --git a/test/unit/LdapConfigTests.cs b/test/unit/LdapConfigTests.cs new file mode 100644 index 000000000..2e728e7c4 --- /dev/null +++ b/test/unit/LdapConfigTests.cs @@ -0,0 +1,49 @@ +using System; +using CommonLibTest.Facades; +using Microsoft.Extensions.Logging; +using Moq; +using SharpHoundCommonLib; +using Xunit; + +namespace CommonLibTest; + +public class LdapConfigTests { + [Fact] + public void UserDomain_IsNullByDefault() { + Assert.Null(new LdapConfig().UserDomain); + } + + [Theory] + [InlineData("child.example.test")] + [InlineData("CHILD")] + public void SetLdapConfig_LogsUserDomain(string userDomain) { + var logger = new Mock>(); + var utils = new LdapUtils(log: logger.Object); + var config = new LdapConfig { UserDomain = userDomain }; + + utils.SetLdapConfig(config); + + logger.VerifyLogContains(LogLevel.Information, "New LDAP Config Set:", $"UserDomain: {userDomain}"); + Assert.Contains($"UserDomain: {userDomain}{Environment.NewLine}", config.ToString()); + } + + [Fact] + public void AllowUncontrolledDomainFallback_IsDisabledByDefault() { + Assert.False(new LdapConfig().AllowUncontrolledDomainFallback); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public void SetLdapConfig_LogsAllowUncontrolledDomainFallback(bool enabled) { + var logger = new Mock>(); + var utils = new LdapUtils(log: logger.Object); + var config = new LdapConfig { AllowUncontrolledDomainFallback = enabled }; + + utils.SetLdapConfig(config); + + logger.VerifyLogContains(LogLevel.Information, + "New LDAP Config Set:", $"AllowUncontrolledDomainFallback: {enabled}"); + Assert.Contains($"AllowUncontrolledDomainFallback: {enabled}{Environment.NewLine}", config.ToString()); + } +} diff --git a/test/unit/LdapConnectionFactoryTests.cs b/test/unit/LdapConnectionFactoryTests.cs new file mode 100644 index 000000000..05eac8888 --- /dev/null +++ b/test/unit/LdapConnectionFactoryTests.cs @@ -0,0 +1,143 @@ +using System; +using System.DirectoryServices.Protocols; +using System.Linq; +using System.Net; +using System.Reflection; +using SharpHoundCommonLib; +using Xunit; + +namespace CommonLibTest; + +public class LdapConnectionFactoryTests { + // Credential is write-only; inspect its stored value without binding to a server. + private static NetworkCredential GetCredential(LdapConnection connection) { + var field = typeof(DirectoryConnection).GetFields(BindingFlags.Instance | BindingFlags.NonPublic) + .Single(x => x.FieldType == typeof(NetworkCredential)); + return (NetworkCredential)field.GetValue(connection); + } + + [Theory] + [InlineData(false, false, 0, 0, 389)] + [InlineData(true, false, 0, 0, 636)] + [InlineData(false, false, 1389, 1636, 1389)] + [InlineData(true, false, 1389, 1636, 1636)] + [InlineData(false, true, 1389, 1636, 3268)] + [InlineData(true, true, 1389, 1636, 3269)] + public void Create_UsesConfiguredPortsAndSessionOptions(bool ssl, bool globalCatalog, + int port, int sslPort, int expectedPort) { + var config = new LdapConfig { Port = port, SSLPort = sslPort }; + + using var connection = LdapConnectionFactory.Create(config, "dc.example.test", ssl, globalCatalog); + + var identifier = Assert.IsType(connection.Directory); + Assert.Equal(new[] { "dc.example.test" }, identifier.Servers); + Assert.Equal(expectedPort, identifier.PortNumber); + Assert.False(identifier.FullyQualifiedDnsHostName); + Assert.False(identifier.Connectionless); + Assert.Equal(TimeSpan.FromMinutes(5), connection.Timeout); + Assert.Equal(3, connection.SessionOptions.ProtocolVersion); + Assert.Equal(ReferralChasingOptions.None, connection.SessionOptions.ReferralChasing); + // On Windows this getter reports connection SSL status; an unbound connection + // does not report an established SSL session. Compare with a configured baseline. + using var baseline = new LdapConnection(new LdapDirectoryIdentifier("dc.example.test", expectedPort, false, false)); + baseline.SessionOptions.SecureSocketLayer = ssl; + Assert.Equal(baseline.SessionOptions.SecureSocketLayer, connection.SessionOptions.SecureSocketLayer); + } + + [Theory] + [InlineData(false, false, true)] + [InlineData(false, true, false)] + [InlineData(true, false, false)] + [InlineData(true, true, false)] + public void Create_PreservesSigningAndSealing(bool ssl, bool disableSigning, bool expected) { + var config = new LdapConfig { DisableSigning = disableSigning }; + + using var connection = LdapConnectionFactory.Create(config, "dc.example.test", ssl); + + Assert.Equal(expected, connection.SessionOptions.Signing); + Assert.Equal(expected, connection.SessionOptions.Sealing); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public void Create_OnlyBypassesCertificateVerificationWhenConfigured(bool disableVerification) { + var config = new LdapConfig { DisableCertVerification = disableVerification }; + + using var connection = LdapConnectionFactory.Create(config, "dc.example.test", true); + + var callback = connection.SessionOptions.VerifyServerCertificate; + if (disableVerification) { + Assert.NotNull(callback); + Assert.True(callback(connection, null)); + } + else { + Assert.Null(callback); + } + } + + [Theory] + [InlineData(AuthType.Kerberos)] + [InlineData(AuthType.Negotiate)] + [InlineData(AuthType.Basic)] + public void Create_UsesConfiguredAuthenticationType(AuthType authType) { + var config = new LdapConfig { AuthType = authType }; + + using var connection = LdapConnectionFactory.Create(config, "dc.example.test", true); + + Assert.Equal(authType, connection.AuthType); + } + + [Fact] + public void Create_LeavesCredentialsUnsetWithoutUsername() { + var config = new LdapConfig { Password = "unused-test-password", UserDomain = "child.example.test" }; + + using var connection = LdapConnectionFactory.Create(config, "dc.example.test", true); + + Assert.Null(GetCredential(connection)); + } + + [Theory] + [InlineData("test-user", "test-password")] + [InlineData("test-user", null)] + [InlineData("", "test-password")] + public void Create_PreservesExplicitCredentials(string username, string password) { + var config = new LdapConfig { Username = username, Password = password, UserDomain = "OTHER" }; + + using var connection = LdapConnectionFactory.Create(config, "dc.example.test", true); + + var credential = GetCredential(connection); + Assert.NotNull(credential); + Assert.Equal(username, credential.UserName); + Assert.Equal(password ?? "", credential.Password); + Assert.Equal("", credential.Domain); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public void Create_PinnedServerUsesExplicitHostAndDisablesReconnectAndReferrals(bool ssl) { + var config = new LdapConfig { Server = "dc.example.test", Port = 1389, SSLPort = 1636 }; + + using var connection = LdapConnectionFactory.Create(config, config.Server, ssl, pinServer: true); + + var identifier = Assert.IsType(connection.Directory); + Assert.Equal(new[] { config.Server }, identifier.Servers); + Assert.Equal(ssl ? 1636 : 1389, identifier.PortNumber); + Assert.True(identifier.FullyQualifiedDnsHostName); + Assert.False(identifier.Connectionless); + Assert.False(connection.SessionOptions.AutoReconnect); + Assert.Equal(ReferralChasingOptions.None, connection.SessionOptions.ReferralChasing); + } + + [Fact] + public void Create_UnpinnedConnectionPreservesDefaultReconnectBehaviorEvenWithConfiguredServer() { + var config = new LdapConfig { Server = "dc.example.test" }; + using var original = new LdapConnection(new LdapDirectoryIdentifier(config.Server, 636, false, false)); + + using var connection = LdapConnectionFactory.Create(config, config.Server, true); + + Assert.Equal(original.SessionOptions.AutoReconnect, connection.SessionOptions.AutoReconnect); + Assert.False(((LdapDirectoryIdentifier)connection.Directory).FullyQualifiedDnsHostName); + } +} diff --git a/test/unit/LdapConnectionPoolTest.cs b/test/unit/LdapConnectionPoolTest.cs index 5ac532c25..b7b5f183f 100644 --- a/test/unit/LdapConnectionPoolTest.cs +++ b/test/unit/LdapConnectionPoolTest.cs @@ -1,3 +1,6 @@ +using System; +using System.Collections.Concurrent; +using System.DirectoryServices.Protocols; using System.Reflection; using System.Threading.Tasks; using Microsoft.Extensions.Logging; @@ -5,20 +8,46 @@ using SharpHoundCommonLib; using Xunit; -public class LdapConnectionPoolTest +public class LdapConnectionPoolTest : IDisposable { + public void Dispose() { + ResetExclusionDomain(); + } + private static void AddExclusionDomain(string identifier) { var excludedDomainsField = typeof(LdapConnectionPool) - .GetField("_excludedDomains", BindingFlags.Static | BindingFlags.NonPublic); + .GetField("ExcludedDomains", BindingFlags.Static | BindingFlags.NonPublic); var excludedDomains = (ConcurrentHashSet)excludedDomainsField.GetValue(null); excludedDomains.Add(identifier); } + private static void ResetExclusionDomain() { + var excludedDomainsField = typeof(LdapConnectionPool) + .GetField("ExcludedDomains", BindingFlags.Static | BindingFlags.NonPublic); + + var excludedDomains = (ConcurrentHashSet)excludedDomainsField.GetValue(null); + + excludedDomains.Clear(); + } + + private static ConcurrentBag GetConnectionsBag(LdapConnectionPool pool) { + var field = typeof(LdapConnectionPool) + .GetField("_connections", BindingFlags.Instance | BindingFlags.NonPublic); + return (ConcurrentBag)field.GetValue(pool); + } + + private static ConcurrentBag GetGlobalCatalogConnectionsBag(LdapConnectionPool pool) { + var field = typeof(LdapConnectionPool) + .GetField("_globalCatalogConnection", BindingFlags.Instance | BindingFlags.NonPublic); + return (ConcurrentBag)field.GetValue(pool); + } + [Fact] public async Task LdapConnectionPool_ExcludedDomains_ShouldExitEarly() { + ResetExclusionDomain(); var mockLogger = new Mock(); var ldapConfig = new LdapConfig(); var connectionPool = new ConnectionPoolManager(ldapConfig, mockLogger.Object); @@ -33,6 +62,7 @@ public async Task LdapConnectionPool_ExcludedDomains_ShouldExitEarly() [Fact] public async Task LdapConnectionPool_ExcludedDomains_NonExcludedShouldntExit() { + ResetExclusionDomain(); var mockLogger = new Mock(); var ldapConfig = new LdapConfig(); var connectionPool = new ConnectionPoolManager(ldapConfig, mockLogger.Object); @@ -42,4 +72,100 @@ public async Task LdapConnectionPool_ExcludedDomains_NonExcludedShouldntExit() Assert.DoesNotContain("excluded for connection attempt", connectAttempt.Message); } -} \ No newline at end of file + + /// + /// Fix: GetGlobalCatalogConnectionAsync was missing the excluded-domain early-exit check. + /// Verifies that a domain in the exclusion list is rejected even for global catalog connections. + /// + [Fact] + public async Task LdapConnectionPool_ExcludedDomains_GlobalCatalog_ShouldExitEarly() + { + ResetExclusionDomain(); + var mockLogger = new Mock(); + var ldapConfig = new LdapConfig(); + var connectionPool = new ConnectionPoolManager(ldapConfig, mockLogger.Object); + + AddExclusionDomain("excludedGcDomain.com"); + var connectAttempt = await connectionPool.TestDomainConnection("excludedGcDomain.com", true); + + Assert.False(connectAttempt.Success); + Assert.Contains("excluded for connection attempt", connectAttempt.Message); + } + + /// + /// Fix: Dispose() previously only drained the regular connection bag; the global-catalog + /// bag was left untouched, leaking those connections. + /// Verifies that Dispose() empties the global-catalog connection bag. + /// + [Fact] + public void LdapConnectionPool_Dispose_ShouldDisposeGlobalCatalogConnections() + { + var ldapConfig = new LdapConfig(); + var pool = new LdapConnectionPool("gc-dispose-test.local", "gc-dispose-test.local", ldapConfig); + + var gcBag = GetGlobalCatalogConnectionsBag(pool); + + // Inject a real (but unconnected) LdapConnection into the GC bag. + var ldapId = new LdapDirectoryIdentifier("localhost", 3268, false, false); + var conn = new LdapConnection(ldapId); + var wrapper = new LdapConnectionWrapper(conn, null, true, "gc-dispose-test.local"); + gcBag.Add(wrapper); + + Assert.False(gcBag.IsEmpty); + + pool.Dispose(); + + // After Dispose the bag must be drained. + Assert.True(gcBag.IsEmpty); + } + + /// + /// Verifies that ReleaseConnection routes a GlobalCatalog wrapper to the GC bag, + /// not the regular connections bag. + /// + [Fact] + public void LdapConnectionPool_ReleaseConnection_GlobalCatalog_RoutesToGCBag() + { + var ldapConfig = new LdapConfig(); + var pool = new LdapConnectionPool("release-gc-test.local", "release-gc-test.local", ldapConfig); + + var connectionsBag = GetConnectionsBag(pool); + var gcBag = GetGlobalCatalogConnectionsBag(pool); + + var ldapId = new LdapDirectoryIdentifier("localhost", 3268, false, false); + var conn = new LdapConnection(ldapId); + var gcWrapper = new LdapConnectionWrapper(conn, null, true, "release-gc-test.local"); + + pool.ReleaseConnection(gcWrapper); + + Assert.Single(gcBag); + Assert.Empty(connectionsBag); + + pool.Dispose(); + } + + /// + /// Verifies that ReleaseConnection routes a non-GlobalCatalog wrapper to the regular + /// connections bag, not the GC bag. + /// + [Fact] + public void LdapConnectionPool_ReleaseConnection_NonGlobalCatalog_RoutesToConnectionsBag() + { + var ldapConfig = new LdapConfig(); + var pool = new LdapConnectionPool("release-regular-test.local", "release-regular-test.local", ldapConfig); + + var connectionsBag = GetConnectionsBag(pool); + var gcBag = GetGlobalCatalogConnectionsBag(pool); + + var ldapId = new LdapDirectoryIdentifier("localhost", 389, false, false); + var conn = new LdapConnection(ldapId); + var wrapper = new LdapConnectionWrapper(conn, null, false, "release-regular-test.local"); + + pool.ReleaseConnection(wrapper); + + Assert.Single(connectionsBag); + Assert.Empty(gcBag); + + pool.Dispose(); + } +} diff --git a/test/unit/LdapDomainFallbackTests.cs b/test/unit/LdapDomainFallbackTests.cs new file mode 100644 index 000000000..0cb738379 --- /dev/null +++ b/test/unit/LdapDomainFallbackTests.cs @@ -0,0 +1,298 @@ +using System; +using System.Collections.Generic; +using System.DirectoryServices.Protocols; +using CommonLibTest.Facades; +using Microsoft.Extensions.Logging; +using Moq; +using SharpHoundCommonLib; +using SharpHoundCommonLib.Enums; +using Xunit; + +namespace CommonLibTest; + +public class LdapDomainFallbackTests { + private const string DomainName = "example.test"; + private const string DomainDn = "DC=example,DC=test"; + + private sealed class Connection : LdapDomainResolver.IConnection { + internal bool FailCore; + internal bool FailMetadata; + internal bool Disposed; + internal Exception BindFailure; + + public void Bind() { + if (BindFailure != null) throw BindFailure; + } + + public IReadOnlyList Search(SearchRequest request) { + if (FailCore) throw new LdapException(); + if (request.DistinguishedName == "") { + return new[] { new MockDirectoryObject("", new Dictionary { + ["defaultNamingContext"] = DomainDn + }) }; + } + if (FailMetadata) throw new LdapException(); + return Array.Empty(); + } + + public IReadOnlyList SearchPage(SearchRequest request, out byte[] cookie) { + if (FailMetadata) throw new LdapException(); + cookie = Array.Empty(); + return Array.Empty(); + } + + public void Dispose() => Disposed = true; + } + + private sealed class LegacyDomain : LdapDomainResolver.ILegacyDomain { + public string Name { get; set; } = DomainName; + internal string NamingContext = DomainDn; + internal bool FailCore; + internal bool FailMetadata; + internal int DisposeCalls; + internal readonly Dictionary NamingContexts = new(); + internal IReadOnlyList Controllers = Array.Empty(); + internal IReadOnlyDictionary Trusts = new Dictionary(); + internal string Forest; + internal string Sid; + internal string Pdc; + + public string DefaultNamingContext => FailCore ? throw new InvalidOperationException() : NamingContext; + public string ForestName => ReadMetadata(Forest); + public string DomainSid => ReadMetadata(Sid); + public string PdcRoleOwnerName => ReadMetadata(Pdc); + public string ReadNamingContext(string attribute) { + NamingContexts.TryGetValue(attribute, out var value); + return ReadMetadata(value); + } + public IReadOnlyList ReadControllerNames() => ReadMetadata(Controllers); + public IReadOnlyDictionary ReadTrustTypes() => ReadMetadata(Trusts); + private T ReadMetadata(T value) => FailMetadata ? throw new InvalidOperationException() : value; + public void Dispose() => DisposeCalls++; + } + + private static LdapDomainResolver CreateResolverWithoutEndpoint( + Func getLegacyDomain) { + return new LdapDomainResolver(new LdapConfig { AllowUncontrolledDomainFallback = true }, + (_, _, _) => throw new Xunit.Sdk.XunitException("No LDAP endpoint available"), () => null, + getLegacyDomain: getLegacyDomain); + } + + [Theory] + [InlineData(false, false)] + [InlineData(false, true)] + [InlineData(true, false)] + [InlineData(true, true)] + public void ControlledSuccess_NeverInvokesFallback(bool enabled, bool failMetadata) { + var connection = new Connection { FailMetadata = failMetadata }; + var connectionCalls = 0; + var resolver = new LdapDomainResolver(new LdapConfig { + Server = "pinned.example.test", AllowUncontrolledDomainFallback = enabled + }, + (target, ssl, pinServer) => { + Assert.Equal("pinned.example.test", target); + Assert.True(ssl); + Assert.True(pinServer); + connectionCalls++; + return connection; + }, () => null, + getLegacyDomain: _ => throw new Xunit.Sdk.XunitException("Fallback must not be invoked")); + + Assert.True(resolver.TryResolveWithFallback(DomainName, out var domain, out var usedLegacy)); + + Assert.False(usedLegacy); + Assert.Equal("EXAMPLE.TEST", domain.Name); + Assert.Null(domain.DomainSid); + Assert.Empty(domain.DomainControllerNames); + Assert.Empty(domain.TrustTypes); + Assert.Equal(1, connectionCalls); + Assert.True(connection.Disposed); + } + + [Fact] + public void ControlledFailure_WithFlagOffNeverInvokesFallback() { + var connection = new Connection { FailCore = true }; + var resolver = new LdapDomainResolver(new LdapConfig { ForceSSL = true }, + (_, _, _) => connection, () => null, + getLegacyDomain: _ => throw new Xunit.Sdk.XunitException("Fallback must not be invoked")); + + Assert.False(resolver.TryResolveWithFallback(DomainName, out var domain, out var usedLegacy)); + + Assert.Null(domain); + Assert.False(usedLegacy); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData((int)LdapErrorCodes.InvalidCredentials, false, true)] + [InlineData((int)LdapErrorCodes.InvalidCredentials, true, true)] + [InlineData((int)ResultCode.InappropriateAuthentication, false, true)] + [InlineData((int)ResultCode.InappropriateAuthentication, true, true)] + [InlineData((int)LdapErrorCodes.InvalidCredentials, false, false)] + [InlineData((int)ResultCode.InappropriateAuthentication, false, false)] + public void AuthenticationRejection_WithFlagOnNeverInvokesFallback(int errorCode, bool forceSsl, + bool rejectOnSsl) { + var sslConnection = new Connection { + BindFailure = new LdapException(rejectOnSsl ? errorCode : 81) + }; + var plaintextConnection = new Connection { BindFailure = new LdapException(errorCode) }; + var attempts = new List(); + var legacyCalls = 0; + var resolver = new LdapDomainResolver(new LdapConfig { + Server = "pinned.example.test", Username = "test-user", Password = "unused", + ForceSSL = forceSsl, AllowUncontrolledDomainFallback = true + }, + (target, ssl, pinServer) => { + Assert.Equal("pinned.example.test", target); + Assert.True(pinServer); + attempts.Add(ssl); + return ssl ? sslConnection : plaintextConnection; + }, () => null, + getLegacyDomain: _ => { + legacyCalls++; + return new LegacyDomain(); + }); + + Assert.False(resolver.TryResolveWithFallback(DomainName, out var domain, out var usedLegacy, + out var metadata)); + + Assert.Null(domain); + Assert.Null(metadata); + Assert.False(usedLegacy); + Assert.Equal(0, legacyCalls); + Assert.Equal(rejectOnSsl ? new[] { true } : new[] { true, false }, attempts); + Assert.True(sslConnection.Disposed); + Assert.Equal(!rejectOnSsl, plaintextConnection.Disposed); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public void TransportFailure_WithFlagOnStillInvokesFallback(bool forceSsl) { + var connections = new List(); + var legacyCalls = 0; + var resolver = new LdapDomainResolver(new LdapConfig { + ForceSSL = forceSsl, AllowUncontrolledDomainFallback = true + }, + (_, _, _) => { + var connection = new Connection { BindFailure = new LdapException(81) }; + connections.Add(connection); + return connection; + }, () => null, + getLegacyDomain: _ => { + legacyCalls++; + return new LegacyDomain(); + }); + + Assert.True(resolver.TryResolveWithFallback(DomainName, out var domain, out var usedLegacy)); + + Assert.NotNull(domain); + Assert.True(usedLegacy); + Assert.Equal(1, legacyCalls); + Assert.Equal(forceSsl ? 1 : 2, connections.Count); + Assert.All(connections, connection => Assert.True(connection.Disposed)); + } + + [Fact] + public void ControlledFailure_WithFlagOnMaterializesAndDisposesFallback() { + var connection = new Connection { FailCore = true }; + var legacy = new LegacyDomain { + Forest = "forest.test", + Sid = "S-1-5-21-111-222-333", + Pdc = "dc.example.test", + NamingContexts = { + ["configurationNamingContext"] = "CN=Configuration," + DomainDn, + ["schemaNamingContext"] = "CN=Schema,CN=Configuration," + DomainDn + }, + Controllers = new[] { "dc.example.test", "DC.EXAMPLE.TEST", null, " " }, + Trusts = new Dictionary { + ["child.example.test"] = TrustType.ParentChild, + ["other.test"] = TrustType.Forest + } + }; + var log = new Mock>(); + var calls = 0; + var resolver = new LdapDomainResolver(new LdapConfig { AllowUncontrolledDomainFallback = true }, + (_, _, _) => connection, () => null, log.Object, name => { + Assert.True(connection.Disposed); + Assert.Equal(DomainName, name); + calls++; + return legacy; + }); + + Assert.True(resolver.TryResolveWithFallback(DomainName, out var domain, out var usedLegacy)); + + Assert.True(usedLegacy); + Assert.Equal(1, calls); + Assert.Equal(DomainName, domain.Name); + Assert.Equal(DomainDn, domain.DefaultNamingContext); + Assert.Equal("forest.test", domain.ForestName); + Assert.Equal("S-1-5-21-111-222-333", domain.DomainSid); + Assert.Equal("dc.example.test", domain.PdcRoleOwnerName); + Assert.Equal("CN=Configuration," + DomainDn, domain.ConfigurationNamingContext); + Assert.Equal("CN=Schema,CN=Configuration," + DomainDn, domain.SchemaNamingContext); + Assert.Equal("dc.example.test", Assert.Single(domain.DomainControllerNames)); + Assert.Equal(TrustType.ParentChild, domain.TrustTypes["CHILD.EXAMPLE.TEST"]); + Assert.Equal(TrustType.Forest, domain.TrustTypes["OTHER.TEST"]); + Assert.Equal(1, legacy.DisposeCalls); + log.VerifyLogContains(LogLevel.Warning, "Using uncontrolled framework domain fallback"); + } + + [Fact] + public void FallbackMetadataFailures_PreserveCoreAndDisposeResource() { + var legacy = new LegacyDomain { FailMetadata = true }; + var resolver = CreateResolverWithoutEndpoint(_ => legacy); + + Assert.True(resolver.TryResolveWithFallback(null, out var domain, out var usedLegacy)); + + Assert.True(usedLegacy); + Assert.Equal(DomainName, domain.Name); + Assert.Null(domain.ForestName); + Assert.Null(domain.DomainSid); + Assert.Null(domain.PdcRoleOwnerName); + Assert.Null(domain.ConfigurationNamingContext); + Assert.Null(domain.SchemaNamingContext); + Assert.Empty(domain.DomainControllerNames); + Assert.Empty(domain.TrustTypes); + Assert.Equal(1, legacy.DisposeCalls); + } + + [Theory] + [InlineData(null, DomainDn)] + [InlineData(" ", DomainDn)] + [InlineData(DomainName, null)] + [InlineData(DomainName, "CN=Invalid")] + public void FallbackMissingCore_ReturnsFailureAndDisposesResource(string name, string namingContext) { + var legacy = new LegacyDomain { Name = name, NamingContext = namingContext }; + var resolver = CreateResolverWithoutEndpoint(_ => legacy); + + Assert.False(resolver.TryResolveWithFallback(null, out var domain, out var usedLegacy)); + + Assert.True(usedLegacy); + Assert.Null(domain); + Assert.Equal(1, legacy.DisposeCalls); + } + + [Fact] + public void FallbackCoreReadThrows_ReturnsFailureAndDisposesResource() { + var legacy = new LegacyDomain { FailCore = true }; + var resolver = CreateResolverWithoutEndpoint(_ => legacy); + + Assert.False(resolver.TryResolveWithFallback(null, out var domain, out _)); + + Assert.Null(domain); + Assert.Equal(1, legacy.DisposeCalls); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public void FallbackCannotOpen_ReturnsFailure(bool throws) { + var resolver = CreateResolverWithoutEndpoint(_ => throws ? throw new InvalidOperationException() : null); + + Assert.False(resolver.TryResolveWithFallback(null, out var domain, out var usedLegacy)); + + Assert.True(usedLegacy); + Assert.Null(domain); + } +} diff --git a/test/unit/LdapDomainInfoTests.cs b/test/unit/LdapDomainInfoTests.cs new file mode 100644 index 000000000..7e4b5a243 --- /dev/null +++ b/test/unit/LdapDomainInfoTests.cs @@ -0,0 +1,57 @@ +using System.Collections.Generic; +using SharpHoundCommonLib.Enums; +using SharpHoundCommonLib.Models; +using Xunit; + +namespace CommonLibTest; + +public class LdapDomainInfoTests { + [Fact] + public void CoreIdentity_AllowsUnavailableAdditionalMetadata() { + var domain = new LdapDomainInfo { + Name = "example.test", + DefaultNamingContext = "DC=example,DC=test" + }; + + Assert.Equal("example.test", domain.Name); + Assert.Equal("DC=example,DC=test", domain.DefaultNamingContext); + Assert.Null(domain.ForestName); + Assert.Null(domain.DomainSid); + Assert.Null(domain.ConfigurationNamingContext); + Assert.Null(domain.SchemaNamingContext); + Assert.Null(domain.PdcRoleOwnerName); + Assert.Empty(domain.DomainControllerNames); + Assert.Empty(domain.TrustTypes); + } + + [Fact] + public void TrustTypes_TargetNamesAreCaseInsensitive() { + var domain = new LdapDomainInfo(); + domain.TrustTypes.Add("CHILD.EXAMPLE.TEST", TrustType.ParentChild); + + Assert.Equal(TrustType.ParentChild, domain.TrustTypes["child.example.test"]); + domain.TrustTypes["Child.Example.Test"] = TrustType.CrossLink; + Assert.Single(domain.TrustTypes); + Assert.Equal(TrustType.CrossLink, domain.TrustTypes["CHILD.EXAMPLE.TEST"]); + } + + [Fact] + public void Collections_AreIndependentForEachResult() { + var first = new LdapDomainInfo(); + var second = new LdapDomainInfo(); + first.DomainControllerNames.Add("dc.example.test"); + first.TrustTypes.Add("child.example.test", TrustType.ParentChild); + + Assert.Empty(second.DomainControllerNames); + Assert.Empty(second.TrustTypes); + } + + [Fact] + public void PublicProperties_ExposeOnlyPlainMetadata() { + foreach (var property in typeof(LdapDomainInfo).GetProperties()) { + Assert.Contains(property.PropertyType, new[] { + typeof(string), typeof(List), typeof(Dictionary) + }); + } + } +} diff --git a/test/unit/LdapDomainResolverTests.cs b/test/unit/LdapDomainResolverTests.cs new file mode 100644 index 000000000..9444bf728 --- /dev/null +++ b/test/unit/LdapDomainResolverTests.cs @@ -0,0 +1,919 @@ +using System; +using System.Collections.Generic; +using System.DirectoryServices.Protocols; +using System.Globalization; +using System.Linq; +using System.Security.Principal; +using System.Text.RegularExpressions; +using CommonLibTest.Facades; +using SharpHoundCommonLib; +using SharpHoundCommonLib.Enums; +using SharpHoundCommonLib.LDAPQueries; +using Xunit; + +namespace CommonLibTest; + +public class LdapDomainResolverTests { + private const string DomainDn = "DC=child,DC=example,DC=test"; + private const string ConfigDn = "CN=Configuration,DC=example,DC=test"; + private const string DomainSid = "S-1-5-21-111-222-333"; + private const string ServerDn = "CN=DC\\, One,CN=Servers,CN=Site,CN=Sites," + ConfigDn; + private const string RoleOwnerDn = "CN=NTDS Settings," + ServerDn; + + private static IDirectoryObject DomainRoot(byte[] sid = null, string owner = RoleOwnerDn) { + if (sid == null) { + var identifier = new SecurityIdentifier(DomainSid); + sid = new byte[identifier.BinaryLength]; + identifier.GetBinaryForm(sid, 0); + } + var values = new Dictionary { ["objectsid"] = sid }; + if (owner != null) values["fsmoroleowner"] = owner; + return new SearchResultEntryWrapper(MockableSearchResultEntry.Construct(values, DomainDn)); + } + + private static IDirectoryObject Entry(params (string Name, object Value)[] attributes) => + new MockDirectoryObject("", attributes.ToDictionary(x => x.Name, x => x.Value, + StringComparer.OrdinalIgnoreCase)); + + private static IDirectoryObject Root() => Entry( + ("defaultNamingContext", DomainDn), + ("rootDomainNamingContext", "DC=example,DC=test"), + ("configurationNamingContext", ConfigDn), + ("schemaNamingContext", "CN=Schema," + ConfigDn)); + + private sealed class FakeConnection : LdapDomainResolver.IConnection { + internal readonly List Requests = new(); + internal Func> OnSearch = _ => new[] { Root() }; + internal Func Entries, byte[] Cookie)> OnPage = + _ => (Array.Empty(), Array.Empty()); + internal Func Entries, byte[] Cookie)> OnTopologyPage = + _ => (Array.Empty(), Array.Empty()); + internal Func Entries, byte[] Cookie)> OnTrustPage = + _ => (Array.Empty(), Array.Empty()); + internal Exception BindFailure; + internal bool Bound; + internal bool Disposed; + + public void Bind() { + Bound = true; + if (BindFailure != null) throw BindFailure; + } + + public IReadOnlyList Search(SearchRequest request) { + Assert.True(Bound); + Requests.Add(request); + return OnSearch(request); + } + + public IReadOnlyList SearchPage(SearchRequest request, out byte[] cookie) { + Assert.True(Bound); + Requests.Add(request); + var page = request.Filter.Equals(CommonFilters.TrustedDomains) ? OnTrustPage(request) : + request.Scope == SearchScope.OneLevel ? OnTopologyPage(request) : OnPage(request); + cookie = page.Cookie; + return page.Entries; + } + + public void Dispose() => Disposed = true; + } + + private sealed class Harness { + internal readonly Queue Connections = new(); + internal readonly List<(string Target, bool Ssl, bool Pinned)> Attempts = new(); + internal int EnvironmentReads; + internal readonly LdapDomainResolver Resolver; + + internal Harness(LdapConfig config = null, string environmentDomain = null) { + Resolver = new LdapDomainResolver(config ?? new LdapConfig(), (target, ssl, pinned) => { + Attempts.Add((target, ssl, pinned)); + return Connections.Dequeue(); + }, () => { + EnvironmentReads++; + return environmentDomain; + }); + } + } + + [Theory] + [InlineData("dc.example.test", "child.example.test", "credential.test", "other.test", "dc.example.test", true)] + [InlineData("dc.example.test", null, "credential.test", "other.test", "dc.example.test", true)] + [InlineData(null, "child.example.test", "credential.test", "other.test", "child.example.test", false)] + [InlineData(" ", "child.example.test", null, "other.test", "child.example.test", false)] + [InlineData(null, null, "child.example.test", "other.test", "child.example.test", false)] + [InlineData(null, " ", " child.example.test ", "other.test", "child.example.test", false)] + [InlineData(null, null, null, "child.example.test", "child.example.test", false)] + [InlineData(null, null, "", "child.example.test", "child.example.test", false)] + [InlineData(null, " ", " ", "other.test", "other.test", false)] + public void TryResolve_SelectsEndpointInOrder(string server, string suppliedDomain, string userDomain, + string environmentDomain, string expectedTarget, bool expectedPinned) { + var harness = new Harness(new LdapConfig { Server = server, UserDomain = userDomain }, environmentDomain); + var connection = new FakeConnection(); + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve(suppliedDomain, out var domain)); + + Assert.Equal("CHILD.EXAMPLE.TEST", domain.Name); + Assert.Equal((expectedTarget, true, expectedPinned), Assert.Single(harness.Attempts)); + Assert.Equal(string.IsNullOrWhiteSpace(server) && string.IsNullOrWhiteSpace(suppliedDomain) && + string.IsNullOrWhiteSpace(userDomain) ? 1 : 0, + harness.EnvironmentReads); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData("child.example.test")] + [InlineData("child.example.test.")] + [InlineData("CHILD")] + public void TryResolve_UserDomainSupportsDnsAndNetBiosEndpointHints(string userDomain) { + var config = new LdapConfig { UserDomain = userDomain }; + var harness = new Harness(config, "local.logon.test"); + var connection = new FakeConnection(); + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve(null, out var domain)); + + Assert.Equal("CHILD.EXAMPLE.TEST", domain.Name); + Assert.Equal((userDomain, true, false), Assert.Single(harness.Attempts)); + Assert.Equal(0, harness.EnvironmentReads); + Assert.Equal(5, connection.Requests.Count); + Assert.Null(config.Username); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData(null)] + [InlineData("")] + [InlineData(" ")] + public void TryResolve_NoTargetDoesNotCreateConnection(string environmentDomain) { + var harness = new Harness(environmentDomain: environmentDomain); + + Assert.False(harness.Resolver.TryResolve(null, out var domain)); + Assert.Null(domain); + Assert.Empty(harness.Attempts); + } + + [Fact] + public void TryResolve_TurkishCulturePreservesDomainAndForestDnsNames() { + var originalCulture = CultureInfo.CurrentCulture; + try { + CultureInfo.CurrentCulture = CultureInfo.GetCultureInfo("tr-TR"); + var harness = new Harness(); + var connection = new FakeConnection { + OnSearch = _ => new[] { Entry(("defaultNamingContext", DomainDn), + ("rootDomainNamingContext", "DC=initial,DC=test")) } + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Equal("CHILD.EXAMPLE.TEST", domain.Name); + Assert.Equal("INITIAL.TEST", domain.ForestName); + Assert.True(connection.Disposed); + } + finally { + CultureInfo.CurrentCulture = originalCulture; + } + } + + [Theory] + [InlineData(null, "child.example.test.", true)] + [InlineData("dc.example.test", "child.example.test.", true)] + [InlineData("dc.example.test", "ChIlD.ExAmPlE.TeSt.", true)] + [InlineData("dc.example.test", "other.test.", false)] + [InlineData("dc.example.test", "child.example.test..", false)] + public void TryResolve_NormalizesOnlyTerminalDnsRootDot(string server, string suppliedDomain, bool expected) { + var harness = new Harness(new LdapConfig { Server = server }); + var connection = new FakeConnection(); + harness.Connections.Enqueue(connection); + + Assert.Equal(expected, harness.Resolver.TryResolve(suppliedDomain, out var domain)); + if (expected) { + Assert.Equal("CHILD.EXAMPLE.TEST", domain.Name); + } + else { + Assert.Null(domain); + } + Assert.Equal((server ?? suppliedDomain, true, server != null), Assert.Single(harness.Attempts)); + Assert.Equal(expected ? 5 : 1, connection.Requests.Count); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData("single", false)] + [InlineData("SiNgLe", false)] + [InlineData("single.", false)] + [InlineData("single", true)] + [InlineData("SiNgLe", true)] + [InlineData("single.", true)] + public void TryResolve_SingleLabelDnsMatchDoesNotRequireNetBiosAlias(string suppliedDomain, + bool missingConfiguration) { + var harness = new Harness(new LdapConfig { Server = "dc.example.test" }); + var connection = new FakeConnection { + OnSearch = request => request.DistinguishedName == "" + ? new[] { missingConfiguration ? Entry(("defaultNamingContext", "DC=single")) : + Entry(("defaultNamingContext", "DC=single"), ("configurationNamingContext", ConfigDn)) } + : request.Scope == SearchScope.OneLevel + ? new[] { Entry(("nCName", "DC=single"), ("nETBIOSName", "OTHER")) } + : Array.Empty() + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve(suppliedDomain, out var domain)); + Assert.Equal("SINGLE", domain.Name); + Assert.Equal("DC=single", domain.DefaultNamingContext); + Assert.DoesNotContain(connection.Requests, request => request.Attributes.Contains("nETBIOSName")); + Assert.Equal(("dc.example.test", true, true), Assert.Single(harness.Attempts)); + Assert.True(connection.Disposed); + } + + [Fact] + public void TryResolve_ReadsRootDseAndMaterializesNamingContexts() { + var harness = new Harness(); + var connection = new FakeConnection(); + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("ChIlD.ExAmPlE.TeSt", out var domain)); + + var request = connection.Requests[0]; + Assert.Equal("", request.DistinguishedName); + Assert.Equal(SearchScope.Base, request.Scope); + Assert.Equal("(objectClass=*)", request.Filter); + Assert.Equal(new[] { "defaultNamingContext", "rootDomainNamingContext", "configurationNamingContext", + "schemaNamingContext" }, request.Attributes.Cast()); + Assert.Empty(request.Controls); + Assert.Equal(DomainDn, domain.DefaultNamingContext); + Assert.Equal("EXAMPLE.TEST", domain.ForestName); + Assert.Equal(ConfigDn, domain.ConfigurationNamingContext); + Assert.Equal("CN=Schema," + ConfigDn, domain.SchemaNamingContext); + Assert.Null(domain.DomainSid); + Assert.Null(domain.PdcRoleOwnerName); + Assert.Empty(domain.DomainControllerNames); + Assert.Empty(domain.TrustTypes); + Assert.True(connection.Disposed); + } + + [Fact] + public void TryResolve_MissingOptionalNamingContextsStillSucceeds() { + var harness = new Harness(); + var connection = new FakeConnection { OnSearch = _ => new[] { Entry(("defaultNamingContext", DomainDn)) } }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Null(domain.ForestName); + Assert.Null(domain.ConfigurationNamingContext); + Assert.Null(domain.SchemaNamingContext); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData(null)] + [InlineData("")] + [InlineData(" ")] + [InlineData("CN=Users")] + [InlineData("DC=,DC=test")] + public void TryResolve_MissingOrMalformedCoreDataFailsWithoutRetry(string defaultNamingContext) { + var harness = new Harness(); + var connection = new FakeConnection { + OnSearch = _ => new[] { defaultNamingContext == null ? Entry() : + Entry(("defaultNamingContext", defaultNamingContext)) } + }; + harness.Connections.Enqueue(connection); + + Assert.False(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Null(domain); + Assert.Single(harness.Attempts); + Assert.True(connection.Disposed); + } + + [Fact] + public void TryResolve_MultipleDefaultNamingContextsFailWithoutRetry() { + var harness = new Harness(); + // The real wrapper reads the first value unless the resolver checks the count. + var entry = MockableSearchResultEntry.Construct(new Dictionary { + ["defaultnamingcontext"] = DomainDn + }, ""); + entry.Attributes["defaultNamingContext"].Add("DC=other,DC=test"); + var connection = new FakeConnection { + OnSearch = _ => new IDirectoryObject[] { new SearchResultEntryWrapper(entry) } + }; + harness.Connections.Enqueue(connection); + + Assert.False(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Null(domain); + Assert.Single(harness.Attempts); + Assert.True(connection.Disposed); + } + + [Fact] + public void TryResolve_MultipleNetBiosAliasesFailWithoutRetry() { + var harness = new Harness(); + var entry = MockableSearchResultEntry.Construct(new Dictionary { + ["ncname"] = DomainDn, + ["netbiosname"] = "CHILD" + }, ""); + entry.Attributes["nETBIOSName"].Add("OTHER"); + var connection = new FakeConnection { + OnSearch = request => request.Scope == SearchScope.Base ? new[] { Root() } : + new IDirectoryObject[] { new SearchResultEntryWrapper(entry) } + }; + harness.Connections.Enqueue(connection); + + Assert.False(harness.Resolver.TryResolve("CHILD", out var domain)); + Assert.Null(domain); + Assert.Single(harness.Attempts); + Assert.True(connection.Disposed); + } + + [Fact] + public void TryResolve_EmptyRootDseFailsAndDisposesConnection() { + var harness = new Harness(); + var connection = new FakeConnection { OnSearch = _ => Array.Empty() }; + harness.Connections.Enqueue(connection); + + Assert.False(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Null(domain); + Assert.True(connection.Disposed); + } + + [Fact] + public void TryResolve_ConfiguredServerAdvertisingDifferentDnsDomainFailsWithoutRetry() { + var harness = new Harness(new LdapConfig { Server = "dc.example.test" }, "child.example.test"); + var connection = new FakeConnection(); + harness.Connections.Enqueue(connection); + + Assert.False(harness.Resolver.TryResolve("other.test", out var domain)); + Assert.Null(domain); + Assert.Equal(("dc.example.test", true, true), Assert.Single(harness.Attempts)); + Assert.Single(connection.Requests); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData("CHILD", DomainDn, true)] + [InlineData("child", DomainDn, true)] + [InlineData("OTHER", DomainDn, false)] + [InlineData(null, DomainDn, false)] + [InlineData("CHILD", "DC=other,DC=test", false)] + public void TryResolve_ValidatesNetBiosAliasOnSameConnection(string alias, string crossRefDn, bool expected) { + var harness = new Harness(new LdapConfig { Server = "dc.example.test" }); + var connection = new FakeConnection { + OnSearch = request => request.Scope == SearchScope.Base ? new[] { Root() } : + new[] { alias == null ? Entry(("nCName", crossRefDn)) : + Entry(("nCName", crossRefDn), ("nETBIOSName", alias)) } + }; + harness.Connections.Enqueue(connection); + + Assert.Equal(expected, harness.Resolver.TryResolve("CHILD", out var domain)); + Assert.Equal(expected, domain != null); + Assert.Equal(("dc.example.test", true, true), Assert.Single(harness.Attempts)); + Assert.Equal(expected ? 6 : 2, connection.Requests.Count); + var request = connection.Requests[1]; + Assert.Equal("CN=Partitions," + ConfigDn, request.DistinguishedName); + Assert.Equal(SearchScope.OneLevel, request.Scope); + Assert.Equal("(&(objectClass=crossRef)(nCName=" + DomainDn + "))", request.Filter); + Assert.Equal(new[] { "nCName", "nETBIOSName" }, request.Attributes.Cast()); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public void TryResolve_NetBiosRequiresReadableCrossReference(bool missingConfiguration) { + var harness = new Harness(); + var connection = new FakeConnection { + OnSearch = request => request.Scope != SearchScope.Base ? Array.Empty() : + new[] { missingConfiguration ? Entry(("defaultNamingContext", DomainDn)) : Root() } + }; + harness.Connections.Enqueue(connection); + + Assert.False(harness.Resolver.TryResolve("CHILD", out var domain)); + Assert.Null(domain); + Assert.Single(harness.Attempts); + Assert.True(connection.Disposed); + } + + [Fact] + public void TryResolve_EscapesNamingContextInCrossReferenceFilter() { + var harness = new Harness(); + var connection = new FakeConnection { + OnSearch = request => request.Scope == SearchScope.Base ? new[] { Entry( + ("defaultNamingContext", "DC=a*(b)\\c\0,DC=test"), ("configurationNamingContext", ConfigDn)) } : + Array.Empty() + }; + harness.Connections.Enqueue(connection); + + Assert.False(harness.Resolver.TryResolve("ALIAS", out _)); + Assert.Equal("(&(objectClass=crossRef)(nCName=DC=a\\2a\\28b\\29\\5cc\\00,DC=test))", + connection.Requests[1].Filter); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public void TryResolve_LdapFailureRetriesSameEndpointAndDisposesBothConnections(bool failDuringBind) { + var harness = new Harness(new LdapConfig { Server = "dc.example.test" }); + var failed = new FakeConnection(); + if (failDuringBind) failed.BindFailure = new LdapException(81); + else failed.OnSearch = _ => throw new DirectoryOperationException("Test search failure"); + var success = new FakeConnection(); + harness.Connections.Enqueue(failed); + harness.Connections.Enqueue(success); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.NotNull(domain); + Assert.Equal(new[] { ("dc.example.test", true, true), ("dc.example.test", false, true) }, harness.Attempts); + Assert.True(failed.Disposed); + Assert.True(success.Disposed); + } + + [Theory] + [InlineData(false, 2)] + [InlineData(true, 1)] + public void TryResolve_ConnectionFailuresRespectForceSsl(bool forceSsl, int expectedAttempts) { + var harness = new Harness(new LdapConfig { ForceSSL = forceSsl }); + var first = new FakeConnection { BindFailure = new LdapException(81) }; + var second = new FakeConnection { BindFailure = new LdapException(81) }; + harness.Connections.Enqueue(first); + harness.Connections.Enqueue(second); + + Assert.False(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Null(domain); + Assert.Equal(expectedAttempts, harness.Attempts.Count); + Assert.True(first.Disposed); + Assert.Equal(!forceSsl, second.Disposed); + } + + [Fact] + public void TryResolve_ConnectionConstructionFailureReturnsFailure() { + var resolver = new LdapDomainResolver(new LdapConfig { ForceSSL = true }, + (_, _, _) => throw new LdapException(81), () => null); + + Assert.False(resolver.TryResolve("child.example.test", out var domain)); + Assert.Null(domain); + } + + [Theory] + [InlineData((int)LdapErrorCodes.InvalidCredentials, false)] + [InlineData((int)LdapErrorCodes.InvalidCredentials, true)] + [InlineData((int)ResultCode.InappropriateAuthentication, false)] + [InlineData((int)ResultCode.InappropriateAuthentication, true)] + public void TryResolve_AuthenticationRejectionFailsWithoutTransportRetry(int errorCode, bool forceSsl) { + var harness = new Harness(new LdapConfig { ForceSSL = forceSsl }); + var rejected = new FakeConnection { BindFailure = new LdapException(errorCode) }; + var second = new FakeConnection(); + harness.Connections.Enqueue(rejected); + harness.Connections.Enqueue(second); + + Assert.False(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Null(domain); + Assert.Equal(("child.example.test", true, false), Assert.Single(harness.Attempts)); + Assert.True(rejected.Bound); + Assert.True(rejected.Disposed); + Assert.Empty(rejected.Requests); + Assert.False(second.Bound); + Assert.False(second.Disposed); + } + + [Fact] + public void TryResolve_ReadsSidPdcAndControllerPagesOnConfiguredConnection() { + // Convert a binary SID, resolve the PDC's parent server DN (including an escaped comma), + // and collect controller pages without connecting to any discovered hostname. + var harness = new Harness(new LdapConfig { Server = "pinned.example.test" }); + var page = 0; + var connection = new FakeConnection { + OnSearch = request => { + if (request.DistinguishedName == "") return new[] { Root() }; + if (request.DistinguishedName == DomainDn) { + Assert.Equal(SearchScope.Base, request.Scope); + Assert.Equal(new[] { "objectSid", "fSMORoleOwner" }, request.Attributes.Cast()); + return new[] { DomainRoot() }; + } + Assert.Equal(ServerDn, request.DistinguishedName); + Assert.Equal(SearchScope.Base, request.Scope); + Assert.Equal("(objectClass=server)", request.Filter); + Assert.Equal(new[] { "dNSHostName" }, request.Attributes.Cast()); + return new[] { Entry(("dNSHostName", "pdc.child.example.test")) }; + }, + OnPage = request => { + Assert.Equal(DomainDn, request.DistinguishedName); + Assert.Equal(SearchScope.Subtree, request.Scope); + Assert.Equal("(|(userAccountControl:1.2.840.113556.1.4.803:=8192)" + + "(userAccountControl:1.2.840.113556.1.4.803:=67108864))", request.Filter); + Assert.Equal(new[] { "dNSHostName" }, request.Attributes.Cast()); + var control = Assert.IsType(Assert.Single(request.Controls.Cast())); + Assert.Equal(500, control.PageSize); + if (page++ == 0) { + Assert.Empty(control.Cookie); + return (new[] { Entry(("dNSHostName", "dc1.child.example.test")), Entry() }, new byte[] { 7, 9 }); + } + // Carry the first page's cookie forward; ignore duplicates, missing names, and invalid hostnames. + Assert.Equal(new byte[] { 7, 9 }, control.Cookie); + return (new[] { Entry(("dNSHostName", "DC1.CHILD.EXAMPLE.TEST")), + Entry(("dNSHostName", " dc2.child.example.test ")), Entry(("dNSHostName", "bad host")) }, + Array.Empty()); + } + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + + Assert.Equal(DomainSid, domain.DomainSid); + Assert.Equal("pdc.child.example.test", domain.PdcRoleOwnerName); + Assert.Equal(new[] { "dc1.child.example.test", "dc2.child.example.test" }, domain.DomainControllerNames); + Assert.Equal(2, page); + Assert.Equal(("pinned.example.test", true, true), Assert.Single(harness.Attempts)); + Assert.True(connection.Disposed); + } + + [Fact] + public void TryResolve_ControllerDiscoveryIncludesWritableAndReadOnlyControllersWithoutWorkstations() { + var accounts = new[] { + (Hostname: "dc.child.example.test", UserAccountControl: 0x82000), + (Hostname: "rodc.child.example.test", UserAccountControl: 0x5001000), + (Hostname: "workstation.child.example.test", UserAccountControl: 0x1000) + }; + var harness = new Harness(); + var connection = new FakeConnection { + OnSearch = request => request.DistinguishedName == "" ? new[] { Root() } : + Array.Empty(), + OnPage = request => { + // Evaluate the two LDAP bitwise-AND clauses against a mixed account set, + // rather than returning an RODC regardless of the requested filter. + var clause = @"\(userAccountControl:1\.2\.840\.113556\.1\.4\.803:=(\d+)\)"; + var filter = Regex.Match(Assert.IsType(request.Filter), @"^\(\|" + clause + clause + @"\)$"); + Assert.True(filter.Success); + var masks = new[] { int.Parse(filter.Groups[1].Value), int.Parse(filter.Groups[2].Value) }; + var entries = accounts.Where(account => masks.Any(mask => + (account.UserAccountControl & mask) == mask)) + .Select(account => Entry(("dNSHostName", account.Hostname))).ToArray(); + return (entries, Array.Empty()); + } + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + + Assert.Equal(new[] { "dc.child.example.test", "rodc.child.example.test" }, domain.DomainControllerNames); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData("root")] + [InlineData("pdc")] + [InlineData("controllers")] + public void TryResolve_MetadataSearchFailuresPreserveOtherMetadataWithoutRetry(string failingRead) { + // Fail each optional LDAP read independently: keep core identity and other available + // metadata, dispose the connection, and avoid a new connection or plaintext retry. + var harness = new Harness(new LdapConfig { Server = "pinned.example.test" }); + var connection = new FakeConnection { + OnSearch = request => { + if (request.DistinguishedName == "") return new[] { Root() }; + if (request.DistinguishedName == DomainDn) { + if (failingRead == "root") throw new LdapException(81); + return new[] { DomainRoot() }; + } + if (failingRead == "pdc") throw new DirectoryOperationException("PDC lookup unavailable"); + return new[] { Entry(("dNSHostName", "pdc.child.example.test")) }; + }, + OnPage = _ => { + if (failingRead == "controllers") throw new LdapException(81); + return (new[] { Entry(("dNSHostName", "dc.child.example.test")) }, Array.Empty()); + } + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + + Assert.Equal("CHILD.EXAMPLE.TEST", domain.Name); + Assert.Equal(DomainDn, domain.DefaultNamingContext); + Assert.Equal(failingRead == "root" ? null : DomainSid, domain.DomainSid); + Assert.Equal(failingRead == "root" || failingRead == "pdc" ? null : "pdc.child.example.test", domain.PdcRoleOwnerName); + Assert.Equal(failingRead == "controllers" ? Array.Empty() : new[] { "dc.child.example.test" }, + domain.DomainControllerNames); + Assert.Equal(("pinned.example.test", true, true), Assert.Single(harness.Attempts)); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public void TryResolve_MalformedSidPreservesPdcAndCoreIdentity(bool emptySid) { + // Empty or truncated SID bytes leave only the SID unavailable; the PDC lookup still succeeds. + var harness = new Harness(); + var connection = new FakeConnection { + OnSearch = request => { + if (request.DistinguishedName == "") return new[] { Root() }; + if (request.DistinguishedName == DomainDn) { + var sid = emptySid ? Array.Empty() : new byte[] { 1, 2 }; + return new[] { DomainRoot(sid) }; + } + return new[] { Entry(("dNSHostName", "pdc.child.example.test")) }; + } + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Null(domain.DomainSid); + Assert.Equal("pdc.child.example.test", domain.PdcRoleOwnerName); + Assert.Single(harness.Attempts); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData(null)] + [InlineData("")] + [InlineData("CN=NTDS Settings,")] + [InlineData("not a distinguished name")] + public void TryResolve_MissingOrMalformedRoleOwnerSkipsPdcLookup(string owner) { + // Without a usable NTDS Settings parent DN, skip the PDC search while retaining the SID. + var harness = new Harness(); + var connection = new FakeConnection { + OnSearch = request => request.DistinguishedName == "" ? new[] { Root() } : new[] { DomainRoot(owner: owner) } + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Equal(DomainSid, domain.DomainSid); + Assert.Null(domain.PdcRoleOwnerName); + Assert.Equal(5, connection.Requests.Count); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData(0)] + [InlineData(2)] + public void TryResolve_UnavailableOrAmbiguousDomainRootLeavesMetadataEmpty(int count) { + // Zero or multiple domain-root entries cannot supply reliable SID/PDC metadata, + // but the identity already resolved from RootDSE remains valid. + var harness = new Harness(); + var connection = new FakeConnection { + OnSearch = request => request.DistinguishedName == "" ? new[] { Root() } : + Enumerable.Range(0, count).Select(_ => DomainRoot()).ToArray() + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Null(domain.DomainSid); + Assert.Null(domain.PdcRoleOwnerName); + Assert.Empty(domain.DomainControllerNames); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData(null)] + [InlineData("")] + [InlineData("bad host")] + public void TryResolve_UnavailableOrMalformedPdcHostnamePreservesSid(string hostname) { + // A server object with no usable hostname leaves the PDC null without discarding the SID. + var harness = new Harness(); + var connection = new FakeConnection { + OnSearch = request => { + if (request.DistinguishedName == "") return new[] { Root() }; + if (request.DistinguishedName == DomainDn) return new[] { DomainRoot() }; + if (hostname == null) return new[] { Entry() }; + return new[] { Entry(("dNSHostName", hostname)) }; + } + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Equal(DomainSid, domain.DomainSid); + Assert.Null(domain.PdcRoleOwnerName); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public void TryResolve_IncompletePagingLeavesControllersEmptyWithoutRetry(bool missingControl) { + // After a successful first page, a failed read or missing paging control must discard + // partial controller results and preserve core success on the configured connection. + var harness = new Harness(new LdapConfig { Server = "pinned.example.test" }); + var pages = 0; + var connection = new FakeConnection { + OnPage = _ => { + if (pages++ == 0) return (new[] { Entry(("dNSHostName", "dc.child.example.test")) }, new byte[] { 1 }); + if (missingControl) return (new[] { Entry(("dNSHostName", "dc2.child.example.test")) }, null); + throw new DirectoryOperationException("Later page unavailable"); + } + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Empty(domain.DomainControllerNames); + Assert.Equal(2, pages); + Assert.Equal(("pinned.example.test", true, true), Assert.Single(harness.Attempts)); + Assert.True(connection.Disposed); + } + + private static IDirectoryObject CrossRef(string name, string parent = null) { + var entry = (MockDirectoryObject)Entry(("nCName", "DC=" + name.Replace(".", ",DC="))); + entry.DistinguishedName = "CN=" + name + ",CN=Partitions," + ConfigDn; + if (parent != null) entry.Properties["trustParent"] = "CN=" + parent + ",CN=Partitions," + ConfigDn; + return entry; + } + + private static IDirectoryObject Trust(string target, int type = 2, + TrustAttributes attributes = TrustAttributes.WithinForest) => + Entry(("trustPartner", target), ("trustType", type), ("trustAttributes", (int)attributes)); + + private static IDirectoryObject[] ForestTopology() => new[] { + CrossRef("example.test"), + CrossRef("child.example.test", "EXAMPLE.TEST"), + CrossRef("grandchild.child.example.test", "child.example.test"), + CrossRef("sibling.example.test", "example.test"), + CrossRef("alternate.test"), + CrossRef("other.test"), + CrossRef("child.alternate.test", "alternate.test") + }; + + [Theory] + [InlineData("child.example.test", "EXAMPLE.TEST", TrustType.ParentChild)] + [InlineData("example.test", "CHILD.EXAMPLE.TEST", TrustType.ParentChild)] + [InlineData("example.test", "alternate.test", TrustType.TreeRoot)] + [InlineData("alternate.test", "example.test", TrustType.TreeRoot)] + [InlineData("alternate.test", "other.test", TrustType.CrossLink)] + [InlineData("child.example.test", "alternate.test", TrustType.CrossLink)] + [InlineData("alternate.test", "child.example.test", TrustType.CrossLink)] + [InlineData("child.example.test", "sibling.example.test", TrustType.CrossLink)] + [InlineData("child.example.test", "grandchild.child.example.test", TrustType.ParentChild)] + [InlineData("example.test", "grandchild.child.example.test", TrustType.CrossLink)] + [InlineData("alternate.test", "child.alternate.test", TrustType.ParentChild)] + public void TryResolve_ClassifiesWithinForestTrustsFromCrossRefLinks(string source, string target, TrustType expected) { + var harness = new Harness(new LdapConfig { Server = "pinned.example.test" }); + var connection = new FakeConnection { + OnSearch = request => request.DistinguishedName == "" ? new[] { Entry( + ("defaultNamingContext", "DC=" + source.Replace(".", ",DC=")), + ("rootDomainNamingContext", "DC=example,DC=test"), ("configurationNamingContext", ConfigDn)) } : + Array.Empty(), + OnTopologyPage = _ => (ForestTopology(), Array.Empty()), + OnTrustPage = _ => (new[] { Trust(target) }, Array.Empty()) + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve(source, out var domain)); + Assert.Equal(expected, domain.TrustTypes[target.ToLowerInvariant()]); + Assert.Single(domain.TrustTypes); + Assert.Equal(("pinned.example.test", true, true), Assert.Single(harness.Attempts)); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData(3, TrustAttributes.WithinForest, TrustType.Kerberos)] + [InlineData(2, TrustAttributes.ForestTransitive, TrustType.Forest)] + [InlineData(2, TrustAttributes.ForestTransitive | TrustAttributes.TreatAsExternal, TrustType.Forest)] + [InlineData(2, TrustAttributes.NonTransitive, TrustType.External)] + [InlineData(1, TrustAttributes.QuarantinedDomain, TrustType.External)] + [InlineData(2, (TrustAttributes)0, TrustType.External)] + [InlineData(2, TrustAttributes.WithinForest, TrustType.Unknown)] + [InlineData(99, (TrustAttributes)0, TrustType.Unknown)] + public void TryResolve_ClassifiesTrustRecordsWithoutTopology(int type, TrustAttributes attributes, TrustType expected) { + var harness = new Harness(); + var connection = new FakeConnection { + OnTopologyPage = _ => throw new LdapException(81), + OnTrustPage = _ => (new[] { Trust("target.test", type, attributes) }, Array.Empty()) + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Equal(expected, domain.TrustTypes["TARGET.TEST"]); + Assert.Single(harness.Attempts); + Assert.True(connection.Disposed); + } + + [Fact] + public void TryResolve_PagesTopologyAndTrustRecordsOnConfiguredConnection() { + var harness = new Harness(new LdapConfig { Server = "pinned.example.test" }); + var topologyPages = 0; + var trustPages = 0; + var connection = new FakeConnection { + OnTopologyPage = request => { + Assert.Equal("CN=Partitions," + ConfigDn, request.DistinguishedName); + Assert.Equal(SearchScope.OneLevel, request.Scope); + Assert.Equal("(&(objectClass=crossRef)(systemFlags:1.2.840.113556.1.4.803:=2))", request.Filter); + Assert.Equal(new[] { "nCName", "trustParent", "distinguishedName" }, request.Attributes.Cast()); + var control = Assert.IsType(Assert.Single(request.Controls.Cast())); + Assert.Equal(500, control.PageSize); + if (topologyPages++ == 0) { + Assert.Empty(control.Cookie); + return (new[] { CrossRef("child.example.test", "example.test") }, new byte[] { 1 }); + } + Assert.Equal(new byte[] { 1 }, control.Cookie); + return (new[] { CrossRef("example.test"), CrossRef("alternate.test") }, Array.Empty()); + }, + OnTrustPage = request => { + Assert.Equal(DomainDn, request.DistinguishedName); + Assert.Equal(SearchScope.Subtree, request.Scope); + Assert.Equal(new[] { "trustPartner", "trustType", "trustAttributes" }, request.Attributes.Cast()); + var control = Assert.IsType(Assert.Single(request.Controls.Cast())); + Assert.Equal(500, control.PageSize); + if (trustPages++ == 0) { + Assert.Empty(control.Cookie); + return (new[] { Trust("example.test") }, new byte[] { 2 }); + } + Assert.Equal(new byte[] { 2 }, control.Cookie); + return (new[] { Trust("EXAMPLE.TEST"), Trust("alternate.test"), Entry() }, Array.Empty()); + } + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Equal(TrustType.ParentChild, domain.TrustTypes["Example.Test"]); + Assert.Equal(TrustType.CrossLink, domain.TrustTypes["ALTERNATE.TEST"]); + Assert.Equal(2, domain.TrustTypes.Count); + Assert.Equal(2, topologyPages); + Assert.Equal(2, trustPages); + Assert.Equal(("pinned.example.test", true, true), Assert.Single(harness.Attempts)); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData("configuration")] + [InlineData("source")] + [InlineData("target")] + [InlineData("forest")] + [InlineData("duplicate")] + public void TryResolve_MissingOrAmbiguousTopologyLeavesWithinForestTrustUnknown(string missing) { + var topology = ForestTopology().ToList(); + if (missing == "source") topology.RemoveAt(1); + if (missing == "target") topology.RemoveAt(4); + if (missing == "duplicate") topology.Add(CrossRef("CHILD.EXAMPLE.TEST", "example.test")); + var source = missing == "forest" ? "example.test" : "child.example.test"; + var harness = new Harness(); + var connection = new FakeConnection { + OnSearch = request => request.DistinguishedName == "" ? new[] { Entry( + ("defaultNamingContext", "DC=" + source.Replace(".", ",DC=")), + ("rootDomainNamingContext", missing == "forest" ? "" : "DC=example,DC=test"), + ("configurationNamingContext", missing == "configuration" ? "" : ConfigDn)) } : Array.Empty(), + OnTopologyPage = _ => (topology, Array.Empty()), + OnTrustPage = _ => (new[] { Trust("alternate.test") }, Array.Empty()) + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve(source, out var domain)); + Assert.Equal(TrustType.Unknown, domain.TrustTypes["alternate.test"]); + Assert.Single(harness.Attempts); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData(true, true)] + [InlineData(true, false)] + [InlineData(false, true)] + [InlineData(false, false)] + public void TryResolve_IncompleteTrustMetadataPreservesCoreWithoutRetry(bool failTopology, bool missingControl) { + var harness = new Harness(new LdapConfig { Server = "pinned.example.test", AllowUncontrolledDomainFallback = true }); + var pages = 0; + (IReadOnlyList, byte[]) ReadIncompletePage(SearchRequest _) { + if (pages++ == 0) return (failTopology ? ForestTopology() : new[] { Trust("example.test") }, new byte[] { 1 }); + if (missingControl) return (Array.Empty(), null); + throw new DirectoryOperationException("Later page unavailable"); + } + var connection = new FakeConnection { + OnTopologyPage = failTopology ? ReadIncompletePage : _ => (ForestTopology(), Array.Empty()), + OnTrustPage = failTopology ? _ => (new[] { Trust("example.test"), + Trust("external.test", attributes: TrustAttributes.NonTransitive) }, Array.Empty()) : ReadIncompletePage + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Equal("CHILD.EXAMPLE.TEST", domain.Name); + Assert.Equal(DomainDn, domain.DefaultNamingContext); + if (failTopology) { + Assert.Equal(TrustType.Unknown, domain.TrustTypes["example.test"]); + Assert.Equal(TrustType.External, domain.TrustTypes["external.test"]); + } + else Assert.Empty(domain.TrustTypes); + Assert.Equal(2, pages); + Assert.Equal(("pinned.example.test", true, true), Assert.Single(harness.Attempts)); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData("trustType")] + [InlineData("trustAttributes")] + [InlineData("sourceDn")] + [InlineData("targetDn")] + [InlineData("emptyParent")] + [InlineData("multipleParents")] + public void TryResolve_UnreadableTrustFieldsOrTopologyLeaveClassificationUnknown(string malformed) { + var source = (MockDirectoryObject)CrossRef("child.example.test", "example.test"); + var target = (MockDirectoryObject)CrossRef("example.test"); + var trust = (MockDirectoryObject)Trust("example.test"); + if (malformed == "trustType" || malformed == "trustAttributes") trust.Properties[malformed] = "invalid"; + if (malformed == "sourceDn") source.DistinguishedName = ""; + if (malformed == "targetDn") target.DistinguishedName = ""; + if (malformed == "emptyParent") source.Properties["trustParent"] = " "; + if (malformed == "multipleParents") source.Properties["trustParent"] = new[] { target.DistinguishedName, "CN=other" }; + var harness = new Harness(); + var connection = new FakeConnection { + OnTopologyPage = _ => (new[] { source, target }, Array.Empty()), + OnTrustPage = _ => (new[] { trust }, Array.Empty()) + }; + harness.Connections.Enqueue(connection); + + Assert.True(harness.Resolver.TryResolve("child.example.test", out var domain)); + Assert.Equal(TrustType.Unknown, domain.TrustTypes["example.test"]); + Assert.Single(harness.Attempts); + Assert.True(connection.Disposed); + } +} diff --git a/test/unit/LdapUtilsDomainTests.cs b/test/unit/LdapUtilsDomainTests.cs new file mode 100644 index 000000000..7ad5060ba --- /dev/null +++ b/test/unit/LdapUtilsDomainTests.cs @@ -0,0 +1,676 @@ +using System; +using System.Collections.Generic; +using System.DirectoryServices.Protocols; +using System.Reflection; +using System.Linq; +using System.Threading; +using CommonLibTest.Facades; +using CommonLibTest.CollectionDefinitions; +using Moq; +using SharpHoundCommonLib; +using SharpHoundCommonLib.Enums; +using SharpHoundCommonLib.Models; +using SharpHoundCommonLib.LDAPQueries; +using SharpHoundRPC.NetAPINative; +using System.Threading.Tasks; +using Xunit; + +namespace CommonLibTest; + +[Collection(nameof(CacheTestCollectionDefinition))] +public class LdapUtilsDomainTests { + private const string DomainName = "child.example.test"; + private const string DomainDn = "DC=child,DC=example,DC=test"; + private const string ConfigurationDn = "CN=Configuration,DC=example,DC=test"; + + private sealed class Connection : LdapDomainResolver.IConnection { + internal bool FailCore; + internal bool IncludeMetadata; + internal string ForestDn; + internal bool Disposed; + internal Func> OnSearch; + internal Func Entries, byte[] Cookie)> OnPage; + internal readonly List Requests = new(); + + public void Bind() { + if (FailCore) throw new LdapException(); + } + + public IReadOnlyList Search(SearchRequest request) { + Requests.Add(request); + if (OnSearch != null) return OnSearch(request); + if (request.DistinguishedName != "") return Array.Empty(); + var attributes = new Dictionary { ["defaultNamingContext"] = DomainDn }; + if (IncludeMetadata) { + attributes["rootDomainNamingContext"] = "DC=example,DC=test"; + attributes["configurationNamingContext"] = ConfigurationDn; + attributes["schemaNamingContext"] = "CN=Schema," + ConfigurationDn; + } + if (ForestDn != null) attributes["rootDomainNamingContext"] = ForestDn; + return new[] { new MockDirectoryObject("", attributes) }; + } + + public IReadOnlyList SearchPage(SearchRequest request, out byte[] cookie) { + Requests.Add(request); + if (OnPage != null) { + var page = OnPage(request); + cookie = page.Cookie; + return page.Entries; + } + cookie = Array.Empty(); + return Array.Empty(); + } + + public void Dispose() => Disposed = true; + } + + private sealed class Harness { + internal int ConnectionAttempts; + internal bool FailCore; + internal string ForestDn; + internal Action ConfigureConnection; + internal DateTime UtcNow = new(2026, 1, 1, 0, 0, 0, DateTimeKind.Utc); + internal readonly List Connections = new(); + internal readonly List Targets = new(); + internal readonly LegacyDomain Legacy = new(); + + internal LdapUtils CreateUtils() => new(CreateResolver, () => UtcNow); + + internal LdapDomainResolver CreateResolver(LdapConfig config) => new(config, (target, _, _) => { + ConnectionAttempts++; + Targets.Add(target); + var connection = new Connection { FailCore = FailCore, ForestDn = ForestDn }; + ConfigureConnection?.Invoke(connection); + Connections.Add(connection); + return connection; + }, () => DomainName, getLegacyDomain: _ => Legacy); + } + + private sealed class LegacyDomain : LdapDomainResolver.ILegacyDomain { + internal int DisposeCalls; + internal bool FailCore; + internal string CoreName = DomainName; + internal string NamingContext = DomainDn; + internal string FailingRead; + internal string Forest; + internal string Sid; + internal string Pdc; + internal readonly Dictionary NamingContexts = new(); + internal IReadOnlyList Controllers = Array.Empty(); + internal IReadOnlyDictionary Trusts = new Dictionary(); + internal readonly Dictionary ReadCounts = new(); + public string Name => CoreName; + public string DefaultNamingContext => FailCore ? throw new InvalidOperationException() : NamingContext; + public string ForestName => Read("forest", Forest); + public string DomainSid => Read("sid", Sid); + public string PdcRoleOwnerName => Read("pdc", Pdc); + public string ReadNamingContext(string attribute) => + Read(attribute, NamingContexts.TryGetValue(attribute, out var value) ? value : null); + public IReadOnlyList ReadControllerNames() => Read("controllers", Controllers); + public IReadOnlyDictionary ReadTrustTypes() => Read("trusts", Trusts); + private T Read(string metadata, T value) { + ReadCounts.TryGetValue(metadata, out var count); + ReadCounts[metadata] = count + 1; + return FailingRead == metadata ? throw new InvalidOperationException("Metadata unavailable") : value; + } + public void Dispose() => DisposeCalls++; + } + + private static IDirectoryObject Entry(string dn, params (string Name, object Value)[] attributes) => + new MockDirectoryObject(dn, attributes.ToDictionary(x => x.Name, x => x.Value)); + + private static Harness CreateMetadataHarness(Func failingRead = null, + bool emptyMetadata = false, bool emptyTopology = false) { + const string serverDn = "CN=DC,CN=Servers," + ConfigurationDn; + const string forestRef = "CN=example.test,CN=Partitions," + ConfigurationDn; + var harness = new Harness(); + harness.ConfigureConnection = connection => { + var failure = failingRead?.Invoke(); + connection.OnSearch = request => { + if (request.DistinguishedName == "") return new[] { Entry("", + ("defaultNamingContext", DomainDn), ("rootDomainNamingContext", "DC=example,DC=test"), + ("configurationNamingContext", ConfigurationDn)) }; + if (request.DistinguishedName == DomainDn) { + if (failure == "root") throw new LdapException(81); + if (emptyMetadata) return Array.Empty(); + var owner = "CN=NTDS Settings," + serverDn; + if (failure == "sid") { + var root = new Mock(); + root.Setup(x => x.TryGetSecurityIdentifier(out It.Ref.IsAny)) + .Throws(new InvalidOperationException("SID unavailable")); + root.Setup(x => x.PropertyCount("fSMORoleOwner")).Returns(1); + root.Setup(x => x.TryGetProperty("fSMORoleOwner", out owner)).Returns(true); + return new[] { root.Object }; + } + return new[] { new MockDirectoryObject(DomainDn, + new Dictionary { ["fSMORoleOwner"] = owner }, "S-1-5-21-111-222-333") }; + } + if (failure == "pdc") throw new LdapException(81); + return new[] { Entry(serverDn, ("dNSHostName", "pdc.child.example.test")) }; + }; + connection.OnPage = request => { + var read = request.Filter.Equals(CommonFilters.TrustedDomains) ? "trusts" : + request.Scope == SearchScope.OneLevel ? "topology" : "controllers"; + if (failure == read) throw new LdapException(81); + if (failure == "incomplete-" + read && + request.Controls.OfType().Single().Cookie.Length != 0) { + throw new DirectoryOperationException("Later page unavailable"); + } + IDirectoryObject[] entries; + if (emptyMetadata || read == "topology" && emptyTopology) entries = Array.Empty(); + else if (read == "controllers") entries = new[] { Entry("", ("dNSHostName", "dc.child.example.test")) }; + else if (read == "topology") entries = new[] { + Entry("CN=child.example.test,CN=Partitions," + ConfigurationDn, + ("nCName", DomainDn), ("trustParent", forestRef)), + Entry(forestRef, ("nCName", "DC=example,DC=test")) + }; + else entries = new[] { Entry("", ("trustPartner", "example.test"), + ("trustType", 2), ("trustAttributes", (int)TrustAttributes.WithinForest)) }; + var cookie = failure == "missing-" + read ? null : + failure == "incomplete-" + read ? new byte[] { 1 } : Array.Empty(); + return (entries, cookie); + }; + }; + return harness; + } + + [Theory] + [InlineData("root")] + [InlineData("sid")] + [InlineData("pdc")] + [InlineData("controllers")] + [InlineData("topology")] + [InlineData("trusts")] + [InlineData("incomplete-controllers")] + [InlineData("incomplete-topology")] + [InlineData("incomplete-trusts")] + [InlineData("missing-controllers")] + [InlineData("missing-topology")] + [InlineData("missing-trusts")] + public void GetDomain_RetriesFailedMetadataAfterBackoff(string failingRead) { + var failure = failingRead; + var harness = CreateMetadataHarness(() => failure); + using var utils = harness.CreateUtils(); + utils.SetLdapConfig(new LdapConfig { Server = "pinned.example.test", ForceSSL = true }); + Assert.True(utils.GetDomain(DomainName, out var first)); + Assert.Equal(DomainName.ToUpperInvariant(), first.Name); + Assert.Equal(DomainDn, first.DefaultNamingContext); + if (failingRead == "root" || failingRead == "sid") Assert.Null(first.DomainSid); + if (failingRead == "root" || failingRead == "pdc") Assert.Null(first.PdcRoleOwnerName); + if (failingRead.EndsWith("controllers")) Assert.Empty(first.DomainControllerNames); + if (failingRead.EndsWith("topology")) Assert.Equal(TrustType.Unknown, first.TrustTypes["example.test"]); + if (failingRead.EndsWith("trusts")) Assert.Empty(first.TrustTypes); + + harness.UtcNow = harness.UtcNow.AddSeconds(29); + Assert.True(utils.GetDomain(DomainName, out var beforeRetry)); + Assert.Same(first, beforeRetry); + Assert.Equal(1, harness.ConnectionAttempts); + + harness.UtcNow = harness.UtcNow.AddSeconds(1); + Assert.True(utils.GetDomain(DomainName, out var stillUnavailable)); + Assert.Equal(first.DomainSid, stillUnavailable.DomainSid); + Assert.Equal(first.PdcRoleOwnerName, stillUnavailable.PdcRoleOwnerName); + Assert.Equal(first.DomainControllerNames, stillUnavailable.DomainControllerNames); + Assert.Equal(first.TrustTypes, stillUnavailable.TrustTypes); + Assert.Equal(2, harness.ConnectionAttempts); + failure = null; + Assert.True(utils.GetDomain(DomainName, out var waiting)); + Assert.Same(stillUnavailable, waiting); + Assert.Equal(2, harness.ConnectionAttempts); + + harness.UtcNow = harness.UtcNow.AddSeconds(30); + Assert.True(utils.GetDomain(DomainName, out var recovered)); + Assert.NotSame(first, recovered); + Assert.Equal(first.Name, recovered.Name); + Assert.Equal(first.DefaultNamingContext, recovered.DefaultNamingContext); + Assert.Equal("S-1-5-21-111-222-333", recovered.DomainSid); + Assert.Equal("pdc.child.example.test", recovered.PdcRoleOwnerName); + Assert.Equal(new[] { "dc.child.example.test" }, recovered.DomainControllerNames); + Assert.Equal(TrustType.ParentChild, recovered.TrustTypes["example.test"]); + Assert.Equal(Enumerable.Repeat("pinned.example.test", 3), harness.Targets); + Assert.All(harness.Connections, connection => Assert.True(connection.Disposed)); + + // Recovery leaves the earlier snapshot intact and successful reads are not repeated. + var retry = harness.Connections[2].Requests; + Assert.Equal(failingRead == "root" || failingRead == "sid" || failingRead == "pdc", + retry.Any(request => request.DistinguishedName == DomainDn && request.Scope == SearchScope.Base)); + Assert.Equal(failingRead == "root" || failingRead == "pdc", + retry.Any(request => request.Attributes.Contains("dNSHostName") && request.Scope == SearchScope.Base)); + Assert.Equal(failingRead.EndsWith("controllers"), + retry.Any(request => request.Attributes.Contains("dNSHostName") && request.Scope == SearchScope.Subtree)); + Assert.Equal(failingRead.EndsWith("topology"), retry.Any(request => request.Scope == SearchScope.OneLevel)); + Assert.Equal(failingRead.EndsWith("trusts"), retry.Any(request => request.Filter.Equals(CommonFilters.TrustedDomains))); + if (failingRead.EndsWith("topology")) Assert.Equal(TrustType.Unknown, first.TrustTypes["example.test"]); + if (failingRead.EndsWith("controllers")) Assert.Empty(first.DomainControllerNames); + + harness.UtcNow = harness.UtcNow.AddHours(1); + Assert.True(utils.GetDomain(DomainName, out var cached)); + Assert.Same(recovered, cached); + Assert.Equal(3, harness.ConnectionAttempts); + } + + [Theory] + [InlineData(true)] + [InlineData(false)] + public void GetDomain_SuccessfulEmptyOrUnknownMetadataDoesNotRetry(bool emptyMetadata) { + var harness = CreateMetadataHarness(emptyMetadata: emptyMetadata, emptyTopology: true); + using var utils = harness.CreateUtils(); + Assert.True(utils.GetDomain(DomainName, out var first)); + if (emptyMetadata) { + Assert.Null(first.DomainSid); + Assert.Null(first.PdcRoleOwnerName); + Assert.Empty(first.DomainControllerNames); + Assert.Empty(first.TrustTypes); + } + else Assert.Equal(TrustType.Unknown, first.TrustTypes["example.test"]); + harness.UtcNow = harness.UtcNow.AddHours(1); + Assert.True(utils.GetDomain(DomainName, out var cached)); + Assert.Same(first, cached); + Assert.Equal(1, harness.ConnectionAttempts); + } + + [Fact] + public void GetDomain_RefreshConnectionFailurePreservesSnapshotAndBacksOffWithoutLegacyFallback() { + string failure = "controllers"; + var harness = CreateMetadataHarness(() => failure); + using var utils = harness.CreateUtils(); + utils.SetLdapConfig(new LdapConfig { ForceSSL = true, AllowUncontrolledDomainFallback = true }); + Assert.True(utils.GetDomain(DomainName, out var first)); + harness.FailCore = true; + harness.UtcNow = harness.UtcNow.AddSeconds(30); + Assert.True(utils.GetDomain(DomainName, out var failedRefresh)); + Assert.Same(first, failedRefresh); + Assert.Equal(2, harness.ConnectionAttempts); + Assert.Equal(0, harness.Legacy.DisposeCalls); + + harness.FailCore = false; + failure = null; + Assert.True(utils.GetDomain(DomainName, out var waiting)); + Assert.Same(first, waiting); + Assert.Equal(2, harness.ConnectionAttempts); + harness.UtcNow = harness.UtcNow.AddSeconds(30); + Assert.True(utils.GetDomain(DomainName, out var recovered)); + Assert.Single(recovered.DomainControllerNames); + Assert.Equal(first.DomainSid, recovered.DomainSid); + Assert.Equal(first.PdcRoleOwnerName, recovered.PdcRoleOwnerName); + Assert.Equal(first.TrustTypes, recovered.TrustTypes); + Assert.Equal(3, harness.ConnectionAttempts); + } + + [Fact] + public void GetDomain_RefreshRejectsChangedIdentityAndPreservesSuccessfulMetadata() { + string failure = "controllers"; + var harness = CreateMetadataHarness(() => failure); + using var utils = harness.CreateUtils(); + Assert.True(utils.GetDomain(out var first)); + var configure = harness.ConfigureConnection; + harness.ConfigureConnection = connection => { + configure(connection); + connection.OnSearch = _ => new[] { Entry("", ("defaultNamingContext", "DC=other,DC=test")) }; + }; + harness.UtcNow = harness.UtcNow.AddSeconds(30); + Assert.True(utils.GetDomain(out var cached)); + Assert.Same(first, cached); + Assert.Equal(DomainName.ToUpperInvariant(), cached.Name); + Assert.Equal(2, harness.ConnectionAttempts); + Assert.Single(harness.Connections[1].Requests); + } + + [Fact] + public void GetDomain_AcceptsSingleLabelNameReturnedByDefaultResolution() { + var harness = new Harness { + ConfigureConnection = connection => connection.OnSearch = request => request.DistinguishedName == "" + ? new[] { Entry("", ("defaultNamingContext", "DC=single")) } + : Array.Empty() + }; + using var utils = harness.CreateUtils(); + + Assert.True(utils.GetDomain(out var first)); + Assert.Equal("SINGLE", first.Name); + Assert.True(utils.GetDomain(first.Name, out var resolved)); + Assert.Equal(first.Name, resolved.Name); + Assert.Equal(first.DefaultNamingContext, resolved.DefaultNamingContext); + Assert.Equal(new[] { DomainName, "SINGLE" }, harness.Targets); + Assert.All(harness.Connections, connection => { + Assert.DoesNotContain(connection.Requests, request => request.Attributes.Contains("nETBIOSName")); + Assert.True(connection.Disposed); + }); + } + + [Fact] + public void GetDomain_RefreshValidatesSingleLabelDnsIdentityWithoutNetBiosLookup() { + string failure = "controllers"; + var harness = CreateMetadataHarness(() => failure); + var configure = harness.ConfigureConnection; + harness.ConfigureConnection = connection => { + configure(connection); + var search = connection.OnSearch; + connection.OnSearch = request => request.DistinguishedName == "" + ? new[] { Entry("", ("defaultNamingContext", "DC=single")) } + : search(request); + }; + using var utils = harness.CreateUtils(); + Assert.True(utils.GetDomain(out var first)); + failure = null; + harness.UtcNow = harness.UtcNow.AddSeconds(30); + Assert.True(utils.GetDomain(out var recovered)); + Assert.Equal("SINGLE", recovered.Name); + Assert.Equal(first.DefaultNamingContext, recovered.DefaultNamingContext); + Assert.Single(recovered.DomainControllerNames); + Assert.Equal(new[] { DomainName, DomainName }, harness.Targets); + Assert.DoesNotContain(harness.Connections[1].Requests, request => request.Scope == SearchScope.OneLevel); + } + + [Fact] + public async Task GetDomain_ConcurrentCallsShareOneMetadataRefresh() { + string failure = "controllers"; + var harness = CreateMetadataHarness(() => failure); + using var utils = harness.CreateUtils(); + Assert.True(utils.GetDomain(DomainName, out var first)); + failure = null; + harness.UtcNow = harness.UtcNow.AddSeconds(30); + using var entered = new ManualResetEventSlim(); + using var release = new ManualResetEventSlim(); + var configure = harness.ConfigureConnection; + harness.ConfigureConnection = connection => { + configure(connection); + var readPage = connection.OnPage; + connection.OnPage = request => { + entered.Set(); + Assert.True(release.Wait(TimeSpan.FromSeconds(10))); + return readPage(request); + }; + }; + var refresh = Task.Run(() => { + Assert.True(utils.GetDomain(DomainName, out var domain)); + return domain; + }); + Task[] callers; + try { + Assert.True(entered.Wait(TimeSpan.FromSeconds(10))); + callers = Enumerable.Range(0, 8).Select(_ => Task.Run(() => { + Assert.True(utils.GetDomain(DomainName, out var domain)); + return domain; + })).ToArray(); + } + finally { + release.Set(); + } + var recovered = await refresh; + Assert.NotSame(first, recovered); + Assert.Empty(first.DomainControllerNames); + Assert.Single(recovered.DomainControllerNames); + foreach (var domain in await Task.WhenAll(callers)) Assert.Same(recovered, domain); + Assert.Equal(2, harness.ConnectionAttempts); + } + + [Fact] + public void GetDomain_CachesControlledSuccessCaseInsensitively() { + var harness = new Harness(); + using var utils = harness.CreateUtils(); + Assert.True(utils.GetDomain(DomainName, out var first)); + Assert.True(utils.GetDomain(" CHILD.EXAMPLE.TEST ", out var second)); + Assert.Same(first, second); + Assert.Equal(1, harness.ConnectionAttempts); + Assert.True(Assert.Single(harness.Connections).Disposed); + Assert.Null(first.DomainSid); + Assert.Null(first.PdcRoleOwnerName); + Assert.Empty(first.DomainControllerNames); + Assert.Equal(0, harness.Legacy.DisposeCalls); + } + + [Fact] + public void GetDomain_DefaultOverloadsShareControlledCache() { + var harness = new Harness(); + using var utils = harness.CreateUtils(); + Assert.True(utils.GetDomain(out var first)); + Assert.True(utils.GetDomain(null, out var second)); + Assert.True(utils.GetDomain(" ", out var third)); + Assert.Same(first, second); + Assert.Same(first, third); + Assert.Equal(1, harness.ConnectionAttempts); + } + + [Fact] + public void GetDomain_ConfigAndUtilsResetInvalidateCache() { + var harness = new Harness(); + using var utils = harness.CreateUtils(); + utils.SetLdapConfig(new LdapConfig { Server = "first.example.test" }); + Assert.True(utils.GetDomain(out var first)); + utils.SetLdapConfig(new LdapConfig { Server = "second.example.test" }); + Assert.True(utils.GetDomain(out var second)); + utils.ResetUtils(); + Assert.True(utils.GetDomain(out var third)); + Assert.NotSame(first, second); + Assert.NotSame(second, third); + Assert.Equal(new[] { "first.example.test", "second.example.test", "second.example.test" }, harness.Targets); + } + + [Fact] + public void GetDomain_InstancesHaveIndependentCachesAndResets() { + var harness = new Harness(); + using var firstUtils = harness.CreateUtils(); + using var secondUtils = harness.CreateUtils(); + Assert.True(firstUtils.GetDomain(DomainName, out var first)); + Assert.True(secondUtils.GetDomain(DomainName, out var second)); + Assert.NotSame(first, second); + firstUtils.ResetUtils(); + Assert.True(secondUtils.GetDomain(DomainName, out var cached)); + Assert.Same(second, cached); + Assert.Equal(2, harness.ConnectionAttempts); + } + + [Fact] + public void GetDomain_CachesLegacySuccessCaseInsensitivelyUntilReset() { + var harness = new Harness { FailCore = true }; + using var utils = harness.CreateUtils(); + utils.SetLdapConfig(new LdapConfig { AllowUncontrolledDomainFallback = true, ForceSSL = true }); + Assert.True(utils.GetDomain(DomainName, out var first)); + harness.UtcNow = harness.UtcNow.AddHours(1); + Assert.True(utils.GetDomain(" CHILD.EXAMPLE.TEST ", out var second)); + Assert.Same(first, second); + Assert.Equal(1, harness.ConnectionAttempts); + Assert.Equal(1, harness.Legacy.DisposeCalls); + harness.FailCore = false; + Assert.True(utils.GetDomain(DomainName, out var cached)); + Assert.Same(first, cached); + Assert.Equal(1, harness.ConnectionAttempts); + utils.ResetUtils(); + Assert.True(utils.GetDomain(DomainName, out var controlled)); + Assert.NotSame(first, controlled); + Assert.Equal(2, harness.ConnectionAttempts); + Assert.Equal(1, harness.Legacy.DisposeCalls); + } + + [Theory] + [InlineData("forest")] + [InlineData("configurationNamingContext")] + [InlineData("schemaNamingContext")] + [InlineData("sid")] + [InlineData("pdc")] + [InlineData("controllers")] + [InlineData("trusts")] + public void GetDomain_CachedLegacyMetadataRetriesOnlyFailedReads(string failingRead) { + var harness = new Harness { FailCore = true }; + var legacy = harness.Legacy; + legacy.FailingRead = failingRead; + legacy.Forest = "example.test"; + legacy.Sid = "S-1-5-21-111-222-333"; + legacy.Pdc = "pdc.child.example.test"; + legacy.NamingContexts["configurationNamingContext"] = ConfigurationDn; + legacy.NamingContexts["schemaNamingContext"] = "CN=Schema," + ConfigurationDn; + legacy.Controllers = new[] { "dc.child.example.test", "DC.CHILD.EXAMPLE.TEST" }; + legacy.Trusts = new Dictionary { ["example.test"] = TrustType.ParentChild }; + using var utils = harness.CreateUtils(); + utils.SetLdapConfig(new LdapConfig { AllowUncontrolledDomainFallback = true, ForceSSL = true }); + Assert.True(utils.GetDomain(DomainName, out var first)); + legacy.FailingRead = null; + // Successful values remain cached even if the framework would now return other data. + if (failingRead != "forest") legacy.Forest = "changed.test"; + harness.UtcNow = harness.UtcNow.AddSeconds(29); + Assert.True(utils.GetDomain(DomainName, out var waiting)); + Assert.Same(first, waiting); + Assert.Equal(1, legacy.DisposeCalls); + harness.UtcNow = harness.UtcNow.AddSeconds(1); + Assert.True(utils.GetDomain(DomainName, out var recovered)); + Assert.NotSame(first, recovered); + Assert.Equal(first.Name, recovered.Name); + Assert.Equal(first.DefaultNamingContext, recovered.DefaultNamingContext); + Assert.Equal("example.test", recovered.ForestName); + Assert.Equal(ConfigurationDn, recovered.ConfigurationNamingContext); + Assert.Equal("CN=Schema," + ConfigurationDn, recovered.SchemaNamingContext); + Assert.Equal("S-1-5-21-111-222-333", recovered.DomainSid); + Assert.Equal("pdc.child.example.test", recovered.PdcRoleOwnerName); + Assert.Equal("dc.child.example.test", Assert.Single(recovered.DomainControllerNames)); + Assert.Equal(TrustType.ParentChild, recovered.TrustTypes["example.test"]); + if (failingRead == "forest") Assert.Null(first.ForestName); + if (failingRead == "controllers") Assert.Empty(first.DomainControllerNames); + if (failingRead == "trusts") Assert.Empty(first.TrustTypes); + foreach (var read in legacy.ReadCounts) Assert.Equal(read.Key == failingRead ? 2 : 1, read.Value); + Assert.Equal(2, legacy.DisposeCalls); + Assert.Equal(1, harness.ConnectionAttempts); + harness.UtcNow = harness.UtcNow.AddHours(1); + Assert.True(utils.GetDomain(DomainName, out var cached)); + Assert.Same(recovered, cached); + Assert.Equal(2, legacy.DisposeCalls); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public void GetDomain_LegacyRefreshFailurePreservesSnapshotAndBacksOff(bool changedIdentity) { + var harness = new Harness { FailCore = true }; + harness.Legacy.Forest = "example.test"; + harness.Legacy.FailingRead = "controllers"; + using var utils = harness.CreateUtils(); + utils.SetLdapConfig(new LdapConfig { AllowUncontrolledDomainFallback = true, ForceSSL = true }); + Assert.True(utils.GetDomain(DomainName, out var first)); + harness.Legacy.FailCore = !changedIdentity; + if (changedIdentity) harness.Legacy.CoreName = "other.test"; + harness.UtcNow = harness.UtcNow.AddSeconds(30); + Assert.True(utils.GetDomain(DomainName, out var failedRefresh)); + Assert.Same(first, failedRefresh); + Assert.Equal(2, harness.Legacy.DisposeCalls); + harness.Legacy.FailCore = false; + harness.Legacy.CoreName = DomainName; + harness.Legacy.FailingRead = null; + harness.Legacy.Controllers = new[] { "dc.child.example.test" }; + Assert.True(utils.GetDomain(DomainName, out var waiting)); + Assert.Same(first, waiting); + Assert.Equal(2, harness.Legacy.DisposeCalls); + harness.UtcNow = harness.UtcNow.AddSeconds(30); + Assert.True(utils.GetDomain(DomainName, out var recovered)); + Assert.Single(recovered.DomainControllerNames); + Assert.Equal(first.ForestName, recovered.ForestName); + Assert.Empty(first.DomainControllerNames); + Assert.Equal(3, harness.Legacy.DisposeCalls); + Assert.Equal(1, harness.ConnectionAttempts); + } + + [Fact] + public void GetDomain_LegacyCoreFailureIsRetriedAndDisablingFallbackClearsCachedSuccess() { + var harness = new Harness { FailCore = true }; + harness.Legacy.FailCore = true; + using var utils = harness.CreateUtils(); + utils.SetLdapConfig(new LdapConfig { AllowUncontrolledDomainFallback = true, ForceSSL = true }); + Assert.False(utils.GetDomain(DomainName, out _)); + Assert.False(utils.GetDomain(DomainName, out _)); + harness.Legacy.FailCore = false; + Assert.True(utils.GetDomain(DomainName, out _)); + Assert.Equal(3, harness.Legacy.DisposeCalls); + Assert.Equal(3, harness.ConnectionAttempts); + utils.SetLdapConfig(new LdapConfig { ForceSSL = true }); + Assert.False(utils.GetDomain(DomainName, out _)); + Assert.Equal(3, harness.Legacy.DisposeCalls); + } + + [Fact] + public void GetDomain_CoreFailureIsNotCachedOrPassedToLegacyByDefault() { + var harness = new Harness { FailCore = true }; + using var utils = harness.CreateUtils(); + utils.SetLdapConfig(new LdapConfig { ForceSSL = true }); + Assert.False(utils.GetDomain(DomainName, out var failed)); + Assert.Null(failed); + harness.FailCore = false; + Assert.True(utils.GetDomain(DomainName, out _)); + Assert.Equal(2, harness.ConnectionAttempts); + Assert.Equal(0, harness.Legacy.DisposeCalls); + } + + [Fact] + public async Task GetForest_ControlledMetadataRespectsInstanceAndConfiguration() { + var firstHarness = new Harness { ForestDn = "DC=first,DC=test" }; + var secondHarness = new Harness { ForestDn = "DC=second,DC=test" }; + using var first = firstHarness.CreateUtils(); + using var second = secondHarness.CreateUtils(); + Assert.Equal((true, "FIRST.TEST"), await first.GetForest(DomainName)); + Assert.Equal((true, "SECOND.TEST"), await second.GetForest(DomainName)); + firstHarness.ForestDn = "DC=updated,DC=test"; + first.SetLdapConfig(new LdapConfig()); + Assert.Equal((true, "UPDATED.TEST"), await first.GetForest(DomainName)); + } + + [Fact] + public async Task GetForest_LegacyMetadataIsCachedUntilReset() { + var harness = new Harness { FailCore = true }; + harness.Legacy.Forest = "first.test"; + using var utils = harness.CreateUtils(); + utils.SetLdapConfig(new LdapConfig { AllowUncontrolledDomainFallback = true, ForceSSL = true }); + Assert.Equal((true, "FIRST.TEST"), await utils.GetForest(DomainName)); + harness.Legacy.Forest = "second.test"; + Assert.Equal((true, "FIRST.TEST"), await utils.GetForest(DomainName)); + Assert.Equal(1, harness.Legacy.DisposeCalls); + utils.ResetUtils(); + Assert.Equal((true, "SECOND.TEST"), await utils.GetForest(DomainName)); + Assert.Equal(2, harness.Legacy.DisposeCalls); + } + + [Theory] + [InlineData(NamingContext.Default, DomainDn)] + [InlineData(NamingContext.Configuration, ConfigurationDn)] + [InlineData(NamingContext.Schema, "CN=Schema," + ConfigurationDn)] + public void PoolSearchBaseUsesAdvertisedForestRoot(NamingContext context, string expected) { + var connection = new Connection { IncludeMetadata = true }; + var resolver = new LdapDomainResolver(new LdapConfig(), (_, _, _) => connection, () => null); + // Force native discovery to fail so the default context also exercises LDAP resolution. + var native = new Mock(); + native.Setup(x => x.CallDsGetDcName(It.IsAny(), It.IsAny(), It.IsAny())) + .Returns(NetAPIResult.Fail("No discovery result")); + using var pool = new LdapConnectionPool(DomainName, DomainName, new LdapConfig(), + nativeMethods: native.Object, domainResolver: resolver); + var wrapper = new LdapConnectionWrapper(null, new MockDirectoryObject("", new Dictionary()), false, DomainName); + var parameters = new LdapQueryParameters { + DomainName = DomainName, NamingContext = context, LDAPFilter = "(objectClass=*)", + RelativeSearchBase = "CN=Container" + }; + var result = CreatePoolSearchRequest(pool, parameters, wrapper); + Assert.True(result.Success); + Assert.Equal("CN=Container," + expected, result.Request.DistinguishedName); + Assert.True(wrapper.GetSearchBase(context, out var saved)); + Assert.Equal(expected, saved); + Assert.True(connection.Disposed); + } + + [Theory] + [InlineData(NamingContext.Configuration)] + [InlineData(NamingContext.Schema)] + public void PoolSearchBaseFailsWhenNamingContextIsUnavailable(NamingContext context) { + var connection = new Connection(); + var resolver = new LdapDomainResolver(new LdapConfig(), (_, _, _) => connection, () => null); + using var pool = new LdapConnectionPool(DomainName, DomainName, new LdapConfig(), domainResolver: resolver); + var wrapper = new LdapConnectionWrapper(null, new MockDirectoryObject("", new Dictionary()), false, DomainName); + var result = CreatePoolSearchRequest(pool, new LdapQueryParameters { + DomainName = DomainName, NamingContext = context, LDAPFilter = "(objectClass=*)" + }, wrapper); + Assert.False(result.Success); + Assert.Null(result.Request); + Assert.False(wrapper.GetSearchBase(context, out _)); + } + + private static (bool Success, SearchRequest Request) CreatePoolSearchRequest(LdapConnectionPool pool, + LdapQueryParameters parameters, LdapConnectionWrapper wrapper) { + var method = typeof(LdapConnectionPool).GetMethod("CreateSearchRequest", BindingFlags.Instance | BindingFlags.NonPublic, + null, new[] { typeof(LdapQueryParameters), typeof(LdapConnectionWrapper) }, null); + return ((bool, SearchRequest))method.Invoke(pool, new object[] { parameters, wrapper }); + } +}