diff --git a/sdk/keyvault/azure-security-keyvault-jca/CHANGELOG.md b/sdk/keyvault/azure-security-keyvault-jca/CHANGELOG.md index a83d3dab37e4..64c04e53f02b 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/CHANGELOG.md +++ b/sdk/keyvault/azure-security-keyvault-jca/CHANGELOG.md @@ -3,6 +3,8 @@ ## 2.13.0-beta.1 (Unreleased) ### Features Added +- Added lazy loading for Key Vault certificate details in the JCA keystore. Certificate details are now loaded by alias when requested, avoiding unnecessary reads for unconfigured certificates. ([#49774](https://github.com/Azure/azure-sdk-for-java/pull/49774)) +- Added support for `azure.keyvault.jca.certificate-alias-filter-pattern` to filter Key Vault certificate aliases with include/exclude regex patterns. Include patterns are configured directly and exclude patterns are prefixed with `!`. Configure more than one filter by appending a suffix to the property name, such as `azure.keyvault.jca.certificate-alias-filter-pattern.1`, so that a pattern can contain any character. If no such property is configured, alias filtering is disabled and all discovered Key Vault aliases remain eligible for lazy loading. ([#39487](https://github.com/Azure/azure-sdk-for-java/issues/39487)) ### Breaking Changes diff --git a/sdk/keyvault/azure-security-keyvault-jca/README.md b/sdk/keyvault/azure-security-keyvault-jca/README.md index 5940351269aa..23e9039fedd4 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/README.md +++ b/sdk/keyvault/azure-security-keyvault-jca/README.md @@ -141,6 +141,7 @@ The JCA library supports configuring the following options: * `azure.keyvault.jca.refresh-certificates-when-have-un-trust-certificate`: Indicates whether to refresh certificates when have untrusted certificate. * `azure.keyvault.jca.certificates-refresh-interval`: The refresh interval time. * `azure.keyvault.jca.certificates-refresh-interval-in-ms`: The refresh interval time. +* `azure.keyvault.jca.certificate-alias-filter-pattern`: A regex that filters which Key Vault certificate aliases are eligible for lazy loading. Append a suffix to the property name to configure more than one filter, for example `azure.keyvault.jca.certificate-alias-filter-pattern.1` or `azure.keyvault.jca.certificate-alias-filter-pattern.prod`. If no such property is configured, all discovered Key Vault aliases are eligible for lazy loading. See "Filtering Key Vault certificate aliases" below. * `azure.keyvault.disable-challenge-resource-verification`: Indicates whether to disable verification that the authentication challenge resource matches the Key Vault or Managed HSM domain. You can configure these properties using: @@ -152,6 +153,30 @@ or as a JVM argument: -Dazure.keyvault.uri= ``` +#### Filtering Key Vault certificate aliases + +Each filter is configured as its own property, so no delimiter is required and a pattern may contain any character: + +```shell +-Dazure.keyvault.jca.certificate-alias-filter-pattern.1='^prod-.*' +-Dazure.keyvault.jca.certificate-alias-filter-pattern.2='^cert-\d{1,5}$' +-Dazure.keyvault.jca.certificate-alias-filter-pattern.exclude-old='!.*-old$' +``` + +* Use an include pattern directly and an exclude pattern with a `!` prefix. +* A suffix can be a number or a string. It only keeps the property names unique and does not affect evaluation, so the filters are unordered. Property names are case-sensitive, which means `.prod` and `.PROD` are two different filters. +* Patterns use full-alias matching (`Pattern.matcher(alias).matches()`). +* An alias is loaded only if it matches at least one include pattern, or if no include pattern is configured, and matches no exclude pattern. +* An invalid pattern fails fast with an `IllegalArgumentException` that names the offending pattern. + +Quote the value as required by your shell, otherwise characters such as `^` and `\` can be altered before the JVM receives them: + +| Shell | Example | +| --- | --- | +| Bash, including Git Bash | `-Dazure.keyvault.jca.certificate-alias-filter-pattern.1='^prod-.*'` | +| PowerShell | `'-Dazure.keyvault.jca.certificate-alias-filter-pattern.1=^prod-.*'` | +| Windows `cmd.exe` | `"-Dazure.keyvault.jca.certificate-alias-filter-pattern.1=^^prod-.*"` | + ### SSL/TLS #### Server side SSL If you are looking to integrate the JCA provider to create an SSLServerSocket see the example below. diff --git a/sdk/keyvault/azure-security-keyvault-jca/checkstyle-suppressions.xml b/sdk/keyvault/azure-security-keyvault-jca/checkstyle-suppressions.xml index 1796f495a42b..12bdda26dccf 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/checkstyle-suppressions.xml +++ b/sdk/keyvault/azure-security-keyvault-jca/checkstyle-suppressions.xml @@ -13,12 +13,14 @@ + + diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/KeyVaultKeyStore.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/KeyVaultKeyStore.java index a87b1ef8ef78..f477dbf2de4d 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/KeyVaultKeyStore.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/KeyVaultKeyStore.java @@ -29,7 +29,10 @@ import java.util.Map; import java.util.Objects; import java.util.Optional; +import java.util.Properties; +import java.util.Set; import java.util.logging.Logger; +import java.util.stream.Collectors; import java.util.stream.Stream; import static java.util.logging.Level.FINE; @@ -56,6 +59,9 @@ public final class KeyVaultKeyStore extends KeyStoreSpi { */ private static final Logger LOGGER = Logger.getLogger(KeyVaultKeyStore.class.getName()); + static final String CERTIFICATE_ALIAS_FILTER_PATTERN_PROPERTY + = "azure.keyvault.jca.certificate-alias-filter-pattern"; + /** * Stores the Jre key store certificates. */ @@ -146,9 +152,10 @@ public KeyVaultKeyStore() { customCertificates = SpecificPathCertificates.getSpecificPathCertificates(customPath); LOGGER.log(FINE, String.format("Loaded custom certificates: %s.", customCertificates.getAliases())); - keyVaultCertificates = new KeyVaultCertificates(refreshInterval, keyVaultUri, tenantId, clientId, clientSecret, - managedIdentity, accessToken, disableChallengeResourceVerification); - LOGGER.log(FINE, String.format("Loaded Key Vault certificates: %s.", keyVaultCertificates.getAliases())); + keyVaultCertificates + = new KeyVaultCertificates(refreshInterval, keyVaultUri, tenantId, clientId, clientSecret, managedIdentity, + accessToken, disableChallengeResourceVerification, getKeyVaultCertificateAliasFilterPatterns()); + LOGGER.log(FINE, () -> String.format("Loaded Key Vault certificates: %s.", keyVaultCertificates.getAliases())); classpathCertificates = new ClasspathCertificates(); LOGGER.log(FINE, String.format("Loaded classpath certificates: %s.", classpathCertificates.getAliases())); @@ -168,6 +175,22 @@ Long getRefreshInterval() { .orElse(0L); } + Set getKeyVaultCertificateAliasFilterPatterns() { + // Each pattern gets its own property because any delimiter character can be part of a regex. + Properties properties = System.getProperties(); + String suffixedPropertyPrefix = CERTIFICATE_ALIAS_FILTER_PATTERN_PROPERTY + "."; + + return properties.stringPropertyNames() + .stream() + .filter(name -> name.equals(CERTIFICATE_ALIAS_FILTER_PATTERN_PROPERTY) + || name.startsWith(suffixedPropertyPrefix)) + .map(properties::getProperty) + .filter(Objects::nonNull) + .map(String::trim) + .filter(pattern -> !pattern.isEmpty()) + .collect(Collectors.toSet()); + } + /** * get key vault key store by system property * @@ -254,16 +277,22 @@ public boolean engineEntryInstanceOf(String alias, Class a.containsKey(alias)) - .findFirst() - .map(certificates -> certificates.get(alias)) - .orElse(null); + Certificate certificate = null; + for (AzureCertificates certificatesSource : allCertificates) { + if (certificatesSource instanceof KeyVaultCertificates) { + certificate = ((KeyVaultCertificates) certificatesSource).getCertificate(alias); + } else { + certificate = certificatesSource.getCertificates().get(alias); + } + + if (certificate != null) { + break; + } + } if (refreshCertificatesWhenHaveUnTrustCertificate && certificate == null) { keyVaultCertificates.refreshCertificates(); - certificate = keyVaultCertificates.getCertificates().get(alias); + certificate = keyVaultCertificates.getCertificate(alias); } return certificate; @@ -283,7 +312,7 @@ public String engineGetCertificateAlias(Certificate cert) { List aliasList = getAllAliases(); for (String candidateAlias : aliasList) { Certificate certificate = engineGetCertificate(candidateAlias); - if (certificate.equals(cert)) { + if (certificate != null && certificate.equals(cert)) { alias = candidateAlias; break; } @@ -307,16 +336,23 @@ public String engineGetCertificateAlias(Certificate cert) { */ @Override public Certificate[] engineGetCertificateChain(String alias) { - Certificate[] certificates = allCertificates.stream() - .map(AzureCertificates::getCertificateChains) - .filter(Objects::nonNull) - .filter(a -> a.containsKey(alias)) - .findFirst() - .map(m -> m.get(alias)) - .orElse(null); + Certificate[] certificates = null; + for (AzureCertificates certificatesSource : allCertificates) { + if (certificatesSource instanceof KeyVaultCertificates) { + certificates = ((KeyVaultCertificates) certificatesSource).getCertificateChain(alias); + } else { + Certificate[] certificateChain = certificatesSource.getCertificateChains().get(alias); + certificates = certificateChain == null ? null : certificateChain.clone(); + } + + if (certificates != null) { + break; + } + } + if (refreshCertificatesWhenHaveUnTrustCertificate && certificates == null) { keyVaultCertificates.refreshCertificates(); - return keyVaultCertificates.getCertificateChains().get(alias); + return keyVaultCertificates.getCertificateChain(alias); } return certificates; } @@ -358,12 +394,20 @@ public KeyStore.Entry engineGetEntry(String alias, KeyStore.ProtectionParameter */ @Override public Key engineGetKey(String alias, char[] password) { - return allCertificates.stream() - .map(AzureCertificates::getCertificateKeys) - .filter(a -> a.containsKey(alias)) - .findFirst() - .map(certificateKeys -> certificateKeys.get(alias)) - .orElse(null); + Key key = null; + for (AzureCertificates certificatesSource : allCertificates) { + if (certificatesSource instanceof KeyVaultCertificates) { + key = ((KeyVaultCertificates) certificatesSource).getCertificateKey(alias); + } else { + key = certificatesSource.getCertificateKeys().get(alias); + } + + if (key != null) { + break; + } + } + + return key; } /** diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/certificates/KeyVaultCertificates.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/certificates/KeyVaultCertificates.java index f4636abb2e0b..5885c9744554 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/certificates/KeyVaultCertificates.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/certificates/KeyVaultCertificates.java @@ -11,20 +11,47 @@ import java.util.Collections; import java.util.Date; import java.util.HashMap; +import java.util.HashSet; import java.util.List; import java.util.Map; import java.util.Objects; import java.util.Optional; +import java.util.Set; +import java.util.logging.Logger; +import java.util.stream.Collectors; +import java.util.regex.Pattern; +import java.util.regex.PatternSyntaxException; + +import static java.util.logging.Level.WARNING; /** * Store certificates loaded from KeyVault. */ public final class KeyVaultCertificates implements AzureCertificates { + private static final Logger LOGGER = Logger.getLogger(KeyVaultCertificates.class.getName()); + private static final String CERTIFICATE_ALIAS_FILTER_PATTERN_PROPERTY + = "azure.keyvault.jca.certificate-alias-filter-pattern"; + /** * Stores the list of aliases. */ private List aliases = new ArrayList<>(); + /** + * Stores aliases whose certificate has already been loaded. + */ + private final Set loadedCertificateAliases = new HashSet<>(); + + /** + * Stores aliases whose certificate chain has already been loaded. + */ + private final Set loadedCertificateChainAliases = new HashSet<>(); + + /** + * Stores aliases whose private key has already been loaded. + */ + private final Set loadedCertificateKeyAliases = new HashSet<>(); + /** * Stores the certificates by alias. */ @@ -49,17 +76,100 @@ public final class KeyVaultCertificates implements AzureCertificates { private final long refreshInterval; + private final List includeAliasPatterns; + + private final List excludeAliasPatterns; + public KeyVaultCertificates(long refreshInterval, String keyVaultUri, String tenantId, String clientId, String clientSecret, String managedIdentity, String accessToken, boolean disableChallengeResourceVerification) { + this(refreshInterval, keyVaultUri, tenantId, clientId, clientSecret, managedIdentity, accessToken, + disableChallengeResourceVerification, Collections.emptySet()); + } + + public KeyVaultCertificates(long refreshInterval, String keyVaultUri, String tenantId, String clientId, + String clientSecret, String managedIdentity, String accessToken, boolean disableChallengeResourceVerification, + Set certificateFilterPatterns) { this.refreshInterval = refreshInterval; + Set normalizedFilterPatterns = normalizeFilterPatterns(certificateFilterPatterns); + this.includeAliasPatterns = getAliasPatterns(normalizedFilterPatterns, false); + this.excludeAliasPatterns = getAliasPatterns(normalizedFilterPatterns, true); updateKeyVaultClient(keyVaultUri, tenantId, clientId, clientSecret, managedIdentity, accessToken, disableChallengeResourceVerification); } public KeyVaultCertificates(long refreshInterval, KeyVaultClient keyVaultClient) { + this(refreshInterval, keyVaultClient, Collections.emptySet()); + } + + public KeyVaultCertificates(long refreshInterval, KeyVaultClient keyVaultClient, + Set certificateFilterPatterns) { this.refreshInterval = refreshInterval; + setKeyVaultClient(keyVaultClient); + Set normalizedFilterPatterns = normalizeFilterPatterns(certificateFilterPatterns); + this.includeAliasPatterns = getAliasPatterns(normalizedFilterPatterns, false); + this.excludeAliasPatterns = getAliasPatterns(normalizedFilterPatterns, true); + } + + private Set normalizeFilterPatterns(Set filterPatterns) { + return Optional.ofNullable(filterPatterns) + .orElse(Collections.emptySet()) + .stream() + .filter(Objects::nonNull) + .map(String::trim) + .filter(pattern -> !pattern.isEmpty()) + .collect(Collectors.toCollection(HashSet::new)); + } + + private List getAliasPatterns(Set filterPatterns, boolean excludePatterns) { + return Optional.ofNullable(filterPatterns) + .orElse(Collections.emptySet()) + .stream() + .filter(Objects::nonNull) + .map(String::trim) + .filter(pattern -> !pattern.isEmpty()) + .filter(pattern -> pattern.startsWith("!") == excludePatterns) + .map(pattern -> { + if (excludePatterns) { + return pattern.substring(1); + } + return pattern; + }) + .filter(pattern -> !pattern.isEmpty()) + .map(this::compileRegexPattern) + .collect(Collectors.toList()); + } + + private Pattern compileRegexPattern(String regexPattern) { + try { + return Pattern.compile(regexPattern); + } catch (PatternSyntaxException exception) { + throw new IllegalArgumentException("Invalid certificate alias filter regex pattern: " + regexPattern + + ". If configured via system property, check '" + CERTIFICATE_ALIAS_FILTER_PATTERN_PROPERTY + "' and '" + + CERTIFICATE_ALIAS_FILTER_PATTERN_PROPERTY + ".'.", exception); + } + } + + private boolean shouldIncludeAlias(String alias) { + if (alias == null) { + return false; + } + + boolean included = includeAliasPatterns.isEmpty() + || includeAliasPatterns.stream().anyMatch(pattern -> pattern.matcher(alias).matches()); + if (!included) { + return false; + } + + return excludeAliasPatterns.stream().noneMatch(pattern -> pattern.matcher(alias).matches()); + } + + private synchronized KeyVaultClient getKeyVaultClient() { + return keyVaultClient; + } + + private synchronized void setKeyVaultClient(KeyVaultClient keyVaultClient) { this.keyVaultClient = keyVaultClient; } @@ -74,19 +184,32 @@ public KeyVaultCertificates(long refreshInterval, KeyVaultClient keyVaultClient) * @param accessToken Access token. * @param disableChallengeResourceVerification Indicates if the challenge resource verification should be disabled. */ - public void updateKeyVaultClient(String keyVaultUri, String tenantId, String clientId, String clientSecret, - String managedIdentity, String accessToken, boolean disableChallengeResourceVerification) { + public synchronized void updateKeyVaultClient(String keyVaultUri, String tenantId, String clientId, + String clientSecret, String managedIdentity, String accessToken, boolean disableChallengeResourceVerification) { if (keyVaultUri != null) { - keyVaultClient = new KeyVaultClient(keyVaultUri, tenantId, clientId, clientSecret, managedIdentity, - accessToken, disableChallengeResourceVerification); + setKeyVaultClient(new KeyVaultClient(keyVaultUri, tenantId, clientId, clientSecret, managedIdentity, + accessToken, disableChallengeResourceVerification)); } else { - keyVaultClient = null; + setKeyVaultClient(null); } + + clearCachedState(); } - boolean certificatesNeedRefresh() { - if (keyVaultClient == null) { + private synchronized void clearCachedState() { + aliases = new ArrayList<>(); + loadedCertificateAliases.clear(); + loadedCertificateChainAliases.clear(); + loadedCertificateKeyAliases.clear(); + certificateKeys.clear(); + certificates.clear(); + certificateChains.clear(); + lastRefreshTime = null; + } + + synchronized boolean certificatesNeedRefresh() { + if (getKeyVaultClient() == null) { return false; } if (lastRefreshTime == null) { @@ -104,8 +227,9 @@ boolean certificatesNeedRefresh() { @Override public List getAliases() { refreshCertificatesIfNeeded(); - - return aliases; + synchronized (this) { + return new ArrayList<>(aliases); + } } /** @@ -116,7 +240,9 @@ public List getAliases() { @Override public Map getCertificates() { refreshCertificatesIfNeeded(); - return certificates; + synchronized (this) { + return new HashMap<>(certificates); + } } /** @@ -126,7 +252,9 @@ public Map getCertificates() { @Override public Map getCertificateChains() { refreshCertificatesIfNeeded(); - return certificateChains; + synchronized (this) { + return copyCertificateChains(certificateChains); + } } /** @@ -137,45 +265,198 @@ public Map getCertificateChains() { @Override public Map getCertificateKeys() { refreshCertificatesIfNeeded(); - return certificateKeys; + synchronized (this) { + return new HashMap<>(certificateKeys); + } + } + + /** + * Get key by alias. + * + * @param alias The alias. + * @return The key, or {@code null}. + */ + public Key getCertificateKey(String alias) { + loadCertificateKeyIfNeeded(alias); + synchronized (this) { + return certificateKeys.get(alias); + } + } + + /** + * Get certificate by alias. + * + * @param alias The alias. + * @return The certificate, or {@code null}. + */ + public Certificate getCertificate(String alias) { + loadCertificateIfNeeded(alias); + synchronized (this) { + return certificates.get(alias); + } + } + + /** + * Get certificate chain by alias. + * + * @param alias The alias. + * @return The certificate chain, or {@code null}. + */ + public Certificate[] getCertificateChain(String alias) { + loadCertificateChainIfNeeded(alias); + synchronized (this) { + Certificate[] chain = certificateChains.get(alias); + return chain == null ? null : chain.clone(); + } + } + + private Map copyCertificateChains(Map source) { + Map copiedChains = new HashMap<>(); + source.forEach((alias, chain) -> copiedChains.put(alias, chain == null ? null : chain.clone())); + return copiedChains; } private void refreshCertificatesIfNeeded() { if (certificatesNeedRefresh()) { // Avoid acquiring the lock as much as possible. + refreshCertificates(false); + } + } + + /** + * Refresh aliases and invalidate cached certificate details. + */ + public void refreshCertificates() { + refreshCertificates(true); + } + + private void refreshCertificates(boolean forceRefresh) { + // Listing aliases keeps the lock so concurrent refreshes cannot apply their results out of order. + synchronized (this) { + if (keyVaultClient == null) { + clearCachedState(); + return; + } + + if (!forceRefresh && !certificatesNeedRefresh()) { + return; + } + + // Discover aliases from Key Vault and apply include/exclude regex filters. + aliases = Optional.ofNullable(keyVaultClient.getAliases()) + .orElse(Collections.emptyList()) + .stream() + .filter(this::shouldIncludeAlias) + .sorted() + .collect(Collectors.toCollection(ArrayList::new)); + + loadedCertificateAliases.clear(); + loadedCertificateChainAliases.clear(); + loadedCertificateKeyAliases.clear(); + certificateKeys.clear(); + certificates.clear(); + certificateChains.clear(); + + lastRefreshTime = new Date(); + } + } + + private void loadCertificateIfNeeded(String alias) { + refreshCertificatesIfNeeded(); + KeyVaultClient currentKeyVaultClient = getKeyVaultClient(); + + if (alias == null || currentKeyVaultClient == null) { + return; + } + + synchronized (this) { + if (loadedCertificateAliases.contains(alias) || !aliases.contains(alias)) { + return; + } + } + + try { + Certificate loadedCertificate = currentKeyVaultClient.getCertificate(alias); synchronized (this) { - if (certificatesNeedRefresh()) { // After obtaining the lock, avoid doing too many operations. - refreshCertificates(); + if (currentKeyVaultClient != keyVaultClient + || loadedCertificateAliases.contains(alias) + || !aliases.contains(alias)) { + return; + } + + if (loadedCertificate != null) { + certificates.put(alias, loadedCertificate); + loadedCertificateAliases.add(alias); } } + } catch (RuntimeException exception) { + LOGGER.log(WARNING, exception, () -> "Failed to load certificate for alias: " + alias); } } - /** - * Refresh certificates. Including certificates, aliases, certificate keys, certificate chains. - */ - public synchronized void refreshCertificates() { - // When refreshing certificates, the update of the 3 variables should be an atomic operation. - aliases = keyVaultClient.getAliases(); - certificateKeys.clear(); - certificates.clear(); - certificateChains.clear(); + private void loadCertificateChainIfNeeded(String alias) { + refreshCertificatesIfNeeded(); + KeyVaultClient currentKeyVaultClient = getKeyVaultClient(); + + if (alias == null || currentKeyVaultClient == null) { + return; + } - Optional.ofNullable(aliases).orElse(Collections.emptyList()).forEach(alias -> { - Key key = keyVaultClient.getKey(alias, null); - if (!Objects.isNull(key)) { - certificateKeys.put(alias, key); + synchronized (this) { + if (loadedCertificateChainAliases.contains(alias) || !aliases.contains(alias)) { + return; } - Certificate certificate = keyVaultClient.getCertificate(alias); - if (!Objects.isNull(certificate)) { - certificates.put(alias, certificate); + } + + try { + Certificate[] loadedCertificateChain = currentKeyVaultClient.getCertificateChain(alias); + synchronized (this) { + if (currentKeyVaultClient != keyVaultClient + || loadedCertificateChainAliases.contains(alias) + || !aliases.contains(alias)) { + return; + } + + if (loadedCertificateChain != null && loadedCertificateChain.length > 0) { + certificateChains.put(alias, loadedCertificateChain); + loadedCertificateChainAliases.add(alias); + } } - Certificate[] certificateChain = keyVaultClient.getCertificateChain(alias); - if (!Objects.isNull(certificateChain)) { - certificateChains.put(alias, certificateChain); + } catch (RuntimeException exception) { + LOGGER.log(WARNING, exception, () -> "Failed to load certificate chain for alias: " + alias); + } + } + + private void loadCertificateKeyIfNeeded(String alias) { + refreshCertificatesIfNeeded(); + KeyVaultClient currentKeyVaultClient = getKeyVaultClient(); + + if (alias == null || currentKeyVaultClient == null) { + return; + } + + synchronized (this) { + if (loadedCertificateKeyAliases.contains(alias) || !aliases.contains(alias)) { + return; } - }); + } + + try { + Key loadedKey = currentKeyVaultClient.getKey(alias, null); + synchronized (this) { + if (currentKeyVaultClient != keyVaultClient + || loadedCertificateKeyAliases.contains(alias) + || !aliases.contains(alias)) { + return; + } - lastRefreshTime = new Date(); + if (loadedKey != null) { + certificateKeys.put(alias, loadedKey); + loadedCertificateKeyAliases.add(alias); + } + } + } catch (RuntimeException exception) { + LOGGER.log(WARNING, exception, () -> "Failed to load certificate key for alias: " + alias); + } } /** @@ -186,8 +467,25 @@ public synchronized void refreshCertificates() { * @return Certificate alias if it exists. */ public String refreshAndGetAliasByCertificate(Certificate certificate) { + if (certificate == null) { + return null; + } + refreshCertificates(); - return getCertificates().entrySet() + + List aliasesSnapshot; + synchronized (this) { + aliasesSnapshot = new ArrayList<>(aliases); + } + + aliasesSnapshot.forEach(this::loadCertificateIfNeeded); + + Map certificatesSnapshot; + synchronized (this) { + certificatesSnapshot = new HashMap<>(certificates); + } + + return certificatesSnapshot.entrySet() .stream() .filter(entry -> certificate.equals(entry.getValue())) .findFirst() @@ -202,10 +500,13 @@ public String refreshAndGetAliasByCertificate(Certificate certificate) { * @param alias Deleted certificate. */ @Override - public void deleteEntry(String alias) { + public synchronized void deleteEntry(String alias) { if (aliases != null) { aliases.remove(alias); } + loadedCertificateAliases.remove(alias); + loadedCertificateChainAliases.remove(alias); + loadedCertificateKeyAliases.remove(alias); certificates.remove(alias); certificateChains.remove(alias); certificateKeys.remove(alias); diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/KeyVaultKeyStoreUnitTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/KeyVaultKeyStoreUnitTest.java index 9e290eaf45fc..656a731afbe1 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/KeyVaultKeyStoreUnitTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/KeyVaultKeyStoreUnitTest.java @@ -3,18 +3,28 @@ package com.azure.security.keyvault.jca; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.parallel.ResourceLock; +import org.junit.jupiter.api.parallel.Resources; import java.io.ByteArrayInputStream; import java.security.ProviderException; import java.security.cert.CertificateException; import java.security.cert.CertificateFactory; import java.security.cert.X509Certificate; +import java.util.Arrays; import java.util.Base64; +import java.util.Collections; +import java.util.HashSet; +import java.util.Set; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertTrue; +@ResourceLock(Resources.SYSTEM_PROPERTIES) public class KeyVaultKeyStoreUnitTest { /** @@ -91,4 +101,56 @@ public void testEngineSetCertificateEntry() { assertNotNull(keystore.engineGetCertificate("setcert")); } + @Test + public void testGetKeyVaultCertificateAliasFilterPatternsWhenNotConfigured() { + assertTrue(new KeyVaultKeyStore().getKeyVaultCertificateAliasFilterPatterns().isEmpty()); + } + + @Test + public void testGetKeyVaultCertificateAliasFilterPatternsFromBaseProperty() { + System.setProperty(KeyVaultKeyStore.CERTIFICATE_ALIAS_FILTER_PATTERN_PROPERTY, " ^prod-.* "); + + assertEquals(Collections.singleton("^prod-.*"), + new KeyVaultKeyStore().getKeyVaultCertificateAliasFilterPatterns()); + } + + @Test + public void testGetKeyVaultCertificateAliasFilterPatternsFromSuffixedProperties() { + String base = KeyVaultKeyStore.CERTIFICATE_ALIAS_FILTER_PATTERN_PROPERTY; + System.setProperty(base, "myalias"); + System.setProperty(base + ".1", "^prod-.*"); + System.setProperty(base + ".prod", "^prod-a.*"); + System.setProperty(base + ".PROD", "^prod-b.*"); + System.setProperty(base + ".exclude-old", "!.*-old$"); + System.setProperty(base + ".blank", " "); + + Set expected + = new HashSet<>(Arrays.asList("myalias", "^prod-.*", "^prod-a.*", "^prod-b.*", "!.*-old$")); + + assertEquals(expected, new KeyVaultKeyStore().getKeyVaultCertificateAliasFilterPatterns()); + } + + @Test + public void testGetKeyVaultCertificateAliasFilterPatternsKeepsCommas() { + String base = KeyVaultKeyStore.CERTIFICATE_ALIAS_FILTER_PATTERN_PROPERTY; + System.setProperty(base + ".1", "^cert-\\d{1,5}$"); + System.setProperty(base + ".2", "![a-z]{2,}"); + + Set expected = new HashSet<>(Arrays.asList("^cert-\\d{1,5}$", "![a-z]{2,}")); + + assertEquals(expected, new KeyVaultKeyStore().getKeyVaultCertificateAliasFilterPatterns()); + } + + @BeforeEach + @AfterEach + public void clearCertificateAliasFilterPatternProperties() { + String base = KeyVaultKeyStore.CERTIFICATE_ALIAS_FILTER_PATTERN_PROPERTY; + System.clearProperty(base); + System.getProperties() + .stringPropertyNames() + .stream() + .filter(name -> name.startsWith(base + ".")) + .forEach(System::clearProperty); + } + } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/certificates/KeyVaultCertificatesTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/certificates/KeyVaultCertificatesTest.java index 2791d1742e6a..890136104f07 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/certificates/KeyVaultCertificatesTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/certificates/KeyVaultCertificatesTest.java @@ -4,25 +4,40 @@ package com.azure.security.keyvault.jca.implementation.certificates; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; import com.azure.security.keyvault.jca.implementation.KeyVaultClient; import java.security.Key; import java.security.cert.Certificate; import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.HashSet; import java.util.List; +import java.util.Map; +import java.util.Set; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; import org.junit.jupiter.api.Assertions; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; public class KeyVaultCertificatesTest { + private static final long TIMEOUT_MILLIS = 10_000; + private final KeyVaultClient keyVaultClient = mock(KeyVaultClient.class); private final Key key = mock(Key.class); private final Certificate certificate = mock(Certificate.class); + private final Certificate[] certificateChain = new Certificate[] { certificate }; + private KeyVaultCertificates keyVaultCertificates; @BeforeEach @@ -32,6 +47,7 @@ public void beforeEach() { when(keyVaultClient.getAliases()).thenReturn(aliases); when(keyVaultClient.getKey("myalias", null)).thenReturn(key); when(keyVaultClient.getCertificate("myalias")).thenReturn(certificate); + when(keyVaultClient.getCertificateChain("myalias")).thenReturn(certificateChain); keyVaultCertificates = new KeyVaultCertificates(60_000, keyVaultClient); } @@ -40,14 +56,77 @@ public void testGetAliases() { Assertions.assertTrue(keyVaultCertificates.getAliases().contains("myalias")); } + @Test + public void testGetAliasesReturnsSnapshot() { + List aliasSnapshot = keyVaultCertificates.getAliases(); + Assertions.assertTrue(aliasSnapshot.contains("myalias")); + + aliasSnapshot.clear(); + + Assertions.assertTrue(keyVaultCertificates.getAliases().contains("myalias")); + } + @Test public void testGetKey() { - Assertions.assertTrue(keyVaultCertificates.getCertificateKeys().containsValue(key)); + Assertions.assertEquals(key, keyVaultCertificates.getCertificateKey("myalias")); + } + + @Test + public void testGetCertificateKeysReturnsSnapshot() { + Assertions.assertEquals(key, keyVaultCertificates.getCertificateKey("myalias")); + + Map keySnapshot = keyVaultCertificates.getCertificateKeys(); + Assertions.assertEquals(key, keySnapshot.get("myalias")); + + keySnapshot.clear(); + + Assertions.assertEquals(key, keyVaultCertificates.getCertificateKeys().get("myalias")); } @Test public void testGetCertificate() { - Assertions.assertTrue(keyVaultCertificates.getCertificates().containsValue(certificate)); + Assertions.assertEquals(certificate, keyVaultCertificates.getCertificate("myalias")); + } + + @Test + public void testGetCertificatesReturnsSnapshot() { + Assertions.assertEquals(certificate, keyVaultCertificates.getCertificate("myalias")); + + Map certificateSnapshot = keyVaultCertificates.getCertificates(); + Assertions.assertEquals(certificate, certificateSnapshot.get("myalias")); + + certificateSnapshot.clear(); + + Assertions.assertEquals(certificate, keyVaultCertificates.getCertificates().get("myalias")); + } + + @Test + public void testGetCertificateChain() { + Assertions.assertArrayEquals(certificateChain, keyVaultCertificates.getCertificateChain("myalias")); + } + + @Test + public void testGetCertificateChainReturnsClone() { + Certificate[] firstRead = keyVaultCertificates.getCertificateChain("myalias"); + Assertions.assertNotNull(firstRead); + + firstRead[0] = null; + + Certificate[] secondRead = keyVaultCertificates.getCertificateChain("myalias"); + Assertions.assertNotNull(secondRead); + Assertions.assertEquals(certificate, secondRead[0]); + } + + @Test + public void testGetCertificateChainsReturnsSnapshot() { + Assertions.assertArrayEquals(certificateChain, keyVaultCertificates.getCertificateChain("myalias")); + + Map chainSnapshot = keyVaultCertificates.getCertificateChains(); + Assertions.assertArrayEquals(certificateChain, chainSnapshot.get("myalias")); + + chainSnapshot.clear(); + + Assertions.assertArrayEquals(certificateChain, keyVaultCertificates.getCertificateChains().get("myalias")); } @Test @@ -59,6 +138,11 @@ public void testRefreshAndGetAliasByCertificate() { Assertions.assertNull(keyVaultCertificates.getCertificates().get("myalias")); } + @Test + public void testRefreshAndGetAliasByCertificateWithNullCertificate() { + Assertions.assertNull(keyVaultCertificates.refreshAndGetAliasByCertificate(null)); + } + @Test public void testDeleteAlias() { Assertions.assertTrue(keyVaultCertificates.getAliases().contains("myalias")); @@ -66,4 +150,333 @@ public void testDeleteAlias() { Assertions.assertFalse(keyVaultCertificates.getAliases().contains("myalias")); } + @Test + public void testGetAliasesDoesNotLoadCertificateDetailsEagerly() { + keyVaultCertificates.getAliases(); + + verify(keyVaultClient, never()).getKey("myalias", null); + verify(keyVaultClient, never()).getCertificate("myalias"); + verify(keyVaultClient, never()).getCertificateChain("myalias"); + } + + @Test + public void testLoadCertificateDetailsForRequestedAliasOnly() { + List aliases = new ArrayList<>(); + aliases.add("myalias"); + aliases.add("otheralias"); + + Key otherKey = mock(Key.class); + Certificate otherCertificate = mock(Certificate.class); + + when(keyVaultClient.getAliases()).thenReturn(aliases); + when(keyVaultClient.getKey("otheralias", null)).thenReturn(otherKey); + when(keyVaultClient.getCertificate("otheralias")).thenReturn(otherCertificate); + + keyVaultCertificates.getCertificate("myalias"); + + verify(keyVaultClient, times(1)).getCertificate("myalias"); + verify(keyVaultClient, never()).getKey("myalias", null); + verify(keyVaultClient, never()).getCertificateChain("myalias"); + verify(keyVaultClient, never()).getKey("otheralias", null); + verify(keyVaultClient, never()).getCertificate("otheralias"); + verify(keyVaultClient, never()).getCertificateChain("otheralias"); + } + + @Test + public void testGetKeyLoadsOnlyKeyForRequestedAlias() { + keyVaultCertificates.getCertificateKey("myalias"); + + verify(keyVaultClient, times(1)).getKey("myalias", null); + verify(keyVaultClient, never()).getCertificate("myalias"); + verify(keyVaultClient, never()).getCertificateChain("myalias"); + } + + @Test + public void testGetCertificateChainLoadsOnlyChainForRequestedAlias() { + keyVaultCertificates.getCertificateChain("myalias"); + + verify(keyVaultClient, times(1)).getCertificateChain("myalias"); + verify(keyVaultClient, never()).getKey("myalias", null); + verify(keyVaultClient, never()).getCertificate("myalias"); + } + + @Test + public void testConfiguredAliasesFilter() { + List aliases = new ArrayList<>(); + aliases.add("myalias"); + aliases.add("otheralias"); + when(keyVaultClient.getAliases()).thenReturn(aliases); + + keyVaultCertificates = new KeyVaultCertificates(60_000, keyVaultClient, Collections.singleton("myalias")); + + List result = keyVaultCertificates.getAliases(); + Assertions.assertEquals(1, result.size()); + Assertions.assertTrue(result.contains("myalias")); + Assertions.assertFalse(result.contains("otheralias")); + } + + @Test + public void testFilterPatternsIncludeRegex() { + List aliases = new ArrayList<>(); + aliases.add("prod-cert"); + aliases.add("dev-cert"); + when(keyVaultClient.getAliases()).thenReturn(aliases); + + keyVaultCertificates = new KeyVaultCertificates(60_000, keyVaultClient, Collections.singleton("^prod-.*")); + + Assertions.assertEquals(Collections.singletonList("prod-cert"), keyVaultCertificates.getAliases()); + verify(keyVaultClient, times(1)).getAliases(); + } + + @Test + public void testFilterPatternsExcludeRegex() { + List aliases = new ArrayList<>(); + aliases.add("prod-active"); + aliases.add("prod-deprecated"); + when(keyVaultClient.getAliases()).thenReturn(aliases); + + Set filterPatterns = new HashSet<>(Arrays.asList("^prod-.*", "!^prod-deprecated$")); + keyVaultCertificates = new KeyVaultCertificates(60_000, keyVaultClient, filterPatterns); + + Assertions.assertEquals(Collections.singletonList("prod-active"), keyVaultCertificates.getAliases()); + verify(keyVaultClient, times(1)).getAliases(); + } + + @Test + public void testConfiguredAliasesFilterAfterRefresh() { + keyVaultCertificates = new KeyVaultCertificates(60_000, keyVaultClient, Collections.singleton("myalias")); + + Assertions.assertEquals(Collections.singletonList("myalias"), keyVaultCertificates.getAliases()); + + keyVaultCertificates.refreshCertificates(); + + List refreshedAliases = keyVaultCertificates.getAliases(); + Assertions.assertEquals(1, refreshedAliases.size()); + Assertions.assertTrue(refreshedAliases.contains("myalias")); + Assertions.assertFalse(refreshedAliases.contains("otheralias")); + Assertions.assertFalse(refreshedAliases.contains("new")); + verify(keyVaultClient, times(2)).getAliases(); + } + + @Test + public void testConfiguredAliasesFilterUsesListApi() { + when(keyVaultClient.getAliases()).thenReturn(Arrays.asList("configured-alias", "other-alias")); + + keyVaultCertificates + = new KeyVaultCertificates(60_000, keyVaultClient, Collections.singleton("configured-alias")); + + Assertions.assertEquals(Collections.singletonList("configured-alias"), keyVaultCertificates.getAliases()); + verify(keyVaultClient, times(1)).getAliases(); + } + + @Test + public void testConfiguredAliasesIgnoreNullEntries() { + Set configuredAliases = new HashSet<>(Arrays.asList("myalias", null)); + keyVaultCertificates = new KeyVaultCertificates(60_000, keyVaultClient, configuredAliases); + + Assertions.assertEquals(Collections.singletonList("myalias"), keyVaultCertificates.getAliases()); + verify(keyVaultClient, times(1)).getAliases(); + } + + @Test + public void testInvalidFilterPatternThrows() { + Set filterPatterns = new HashSet<>(Collections.singletonList("[invalid")); + + Assertions.assertThrows(IllegalArgumentException.class, + () -> new KeyVaultCertificates(60_000, keyVaultClient, filterPatterns)); + } + + @Test + public void testFilterPatternWithBoundedQuantifier() { + when(keyVaultClient.getAliases()).thenReturn(Arrays.asList("cert-42", "cert-1234567", "cert-abc")); + + keyVaultCertificates + = new KeyVaultCertificates(60_000, keyVaultClient, Collections.singleton("^cert-\\d{1,5}$")); + + Assertions.assertEquals(Collections.singletonList("cert-42"), keyVaultCertificates.getAliases()); + } + + @Test + public void testGetCertificateWithUnconfiguredAliasDoesNotFetchDetails() { + keyVaultCertificates = new KeyVaultCertificates(60_000, keyVaultClient, Collections.singleton("myalias")); + + Assertions.assertNull(keyVaultCertificates.getCertificate("otheralias")); + + verify(keyVaultClient, never()).getKey("otheralias", null); + verify(keyVaultClient, never()).getCertificate("otheralias"); + verify(keyVaultClient, never()).getCertificateChain("otheralias"); + } + + @Test + public void testAliasCertificateLoadFailureIsRetriedOnNextAccess() { + when(keyVaultClient.getCertificate("myalias")).thenThrow(new RuntimeException("transient error")) + .thenReturn(certificate); + + Assertions.assertNull(keyVaultCertificates.getCertificate("myalias")); + Assertions.assertEquals(certificate, keyVaultCertificates.getCertificate("myalias")); + verify(keyVaultClient, times(2)).getCertificate("myalias"); + verify(keyVaultClient, never()).getKey("myalias", null); + verify(keyVaultClient, never()).getCertificateChain("myalias"); + } + + @Test + public void testAliasCertificateNullLoadIsRetriedOnNextAccess() { + when(keyVaultClient.getCertificate("myalias")).thenReturn(null).thenReturn(certificate); + + Assertions.assertNull(keyVaultCertificates.getCertificate("myalias")); + Assertions.assertEquals(certificate, keyVaultCertificates.getCertificate("myalias")); + verify(keyVaultClient, times(2)).getCertificate("myalias"); + verify(keyVaultClient, never()).getKey("myalias", null); + verify(keyVaultClient, never()).getCertificateChain("myalias"); + } + + @Test + public void testAliasKeyLoadFailureIsRetriedOnNextAccess() { + when(keyVaultClient.getKey("myalias", null)).thenThrow(new RuntimeException("transient error")).thenReturn(key); + + Assertions.assertNull(keyVaultCertificates.getCertificateKey("myalias")); + Assertions.assertEquals(key, keyVaultCertificates.getCertificateKey("myalias")); + verify(keyVaultClient, times(2)).getKey("myalias", null); + verify(keyVaultClient, never()).getCertificate("myalias"); + verify(keyVaultClient, never()).getCertificateChain("myalias"); + } + + @Test + public void testAliasKeyNullLoadIsRetriedOnNextAccess() { + when(keyVaultClient.getKey("myalias", null)).thenReturn(null).thenReturn(key); + + Assertions.assertNull(keyVaultCertificates.getCertificateKey("myalias")); + Assertions.assertEquals(key, keyVaultCertificates.getCertificateKey("myalias")); + verify(keyVaultClient, times(2)).getKey("myalias", null); + verify(keyVaultClient, never()).getCertificate("myalias"); + verify(keyVaultClient, never()).getCertificateChain("myalias"); + } + + @Test + public void testAliasChainLoadFailureIsRetriedOnNextAccess() { + when(keyVaultClient.getCertificateChain("myalias")).thenThrow(new RuntimeException("transient error")) + .thenReturn(certificateChain); + + Assertions.assertNull(keyVaultCertificates.getCertificateChain("myalias")); + Assertions.assertArrayEquals(certificateChain, keyVaultCertificates.getCertificateChain("myalias")); + verify(keyVaultClient, times(2)).getCertificateChain("myalias"); + verify(keyVaultClient, never()).getCertificate("myalias"); + verify(keyVaultClient, never()).getKey("myalias", null); + } + + @Test + public void testAliasChainEmptyLoadIsRetriedOnNextAccess() { + when(keyVaultClient.getCertificateChain("myalias")).thenReturn(new Certificate[0]).thenReturn(certificateChain); + + Assertions.assertNull(keyVaultCertificates.getCertificateChain("myalias")); + Assertions.assertArrayEquals(certificateChain, keyVaultCertificates.getCertificateChain("myalias")); + verify(keyVaultClient, times(2)).getCertificateChain("myalias"); + verify(keyVaultClient, never()).getCertificate("myalias"); + verify(keyVaultClient, never()).getKey("myalias", null); + } + + @Test + public void testUpdateKeyVaultClientClearsCachedState() { + Assertions.assertTrue(keyVaultCertificates.getAliases().contains("myalias")); + Assertions.assertEquals(certificate, keyVaultCertificates.getCertificate("myalias")); + + keyVaultCertificates.updateKeyVaultClient(null, null, null, null, null, null, false); + + Assertions.assertTrue(keyVaultCertificates.getAliases().isEmpty()); + Assertions.assertTrue(keyVaultCertificates.getCertificates().isEmpty()); + Assertions.assertTrue(keyVaultCertificates.getCertificateChains().isEmpty()); + Assertions.assertTrue(keyVaultCertificates.getCertificateKeys().isEmpty()); + Assertions.assertNull(keyVaultCertificates.getCertificate("myalias")); + } + + @Test + public void testConcurrentForceRefreshAppliesLatestAliases() throws Exception { + CountDownLatch firstListCallStarted = new CountDownLatch(1); + CountDownLatch firstListCallMayFinish = new CountDownLatch(1); + AtomicInteger listCallCount = new AtomicInteger(); + + when(keyVaultClient.getAliases()).thenAnswer(invocation -> { + if (listCallCount.getAndIncrement() == 0) { + firstListCallStarted.countDown(); + awaitLatch(firstListCallMayFinish); + return Collections.singletonList("stale-alias"); + } + return Collections.singletonList("fresh-alias"); + }); + + keyVaultCertificates = new KeyVaultCertificates(60_000, keyVaultClient); + + Thread slowRefresh = new Thread(keyVaultCertificates::refreshCertificates); + slowRefresh.start(); + awaitLatch(firstListCallStarted); + + Thread fastRefresh = new Thread(keyVaultCertificates::refreshCertificates); + fastRefresh.start(); + awaitThreadsParked(Collections.singletonList(fastRefresh)); + + firstListCallMayFinish.countDown(); + slowRefresh.join(TIMEOUT_MILLIS); + fastRefresh.join(TIMEOUT_MILLIS); + + Assertions.assertEquals(Collections.singletonList("fresh-alias"), keyVaultCertificates.getAliases()); + } + + @Test + public void testConcurrentRefreshIssuesSingleAliasListCall() throws Exception { + CountDownLatch listCallStarted = new CountDownLatch(1); + CountDownLatch listCallMayFinish = new CountDownLatch(1); + + when(keyVaultClient.getAliases()).thenAnswer(invocation -> { + listCallStarted.countDown(); + awaitLatch(listCallMayFinish); + return Collections.singletonList("myalias"); + }); + + keyVaultCertificates = new KeyVaultCertificates(60_000, keyVaultClient); + + List readers = new ArrayList<>(); + for (int i = 0; i < 4; i++) { + Thread reader = new Thread(keyVaultCertificates::getAliases); + readers.add(reader); + reader.start(); + } + + awaitLatch(listCallStarted); + awaitThreadsParked(readers); + listCallMayFinish.countDown(); + + for (Thread reader : readers) { + reader.join(TIMEOUT_MILLIS); + } + + verify(keyVaultClient, times(1)).getAliases(); + } + + private static void awaitLatch(CountDownLatch latch) throws InterruptedException { + if (!latch.await(TIMEOUT_MILLIS, TimeUnit.MILLISECONDS)) { + throw new IllegalStateException("Timed out waiting for the test latch."); + } + } + + /** + * Waits until none of the threads can still reach Key Vault, so the pending call cannot be released too early. + */ + private static void awaitThreadsParked(List threads) throws InterruptedException { + long deadline = System.nanoTime() + TimeUnit.MILLISECONDS.toNanos(TIMEOUT_MILLIS); + + while (System.nanoTime() < deadline) { + boolean parked = threads.stream() + .map(Thread::getState) + .noneMatch(state -> state == Thread.State.NEW || state == Thread.State.RUNNABLE); + + if (parked) { + return; + } + + Thread.sleep(10); + } + + throw new IllegalStateException("Timed out waiting for the threads to park."); + } + }