diff --git a/sdk/keyvault/azure-security-keyvault-jca/CHANGELOG.md b/sdk/keyvault/azure-security-keyvault-jca/CHANGELOG.md index 366ac061f8d7..f6d584c06a75 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/CHANGELOG.md +++ b/sdk/keyvault/azure-security-keyvault-jca/CHANGELOG.md @@ -16,6 +16,7 @@ ### Other Changes - Added system property `azure.keyvault.jca.disable-aia-download` to disable automatic AIA chain completion. AIA chain completion downloads certificates from URLs embedded in certificate extensions, so this allows locked-down environments to prevent those outbound HTTP(S) requests, mitigating potential SSRF-like attack vectors when loading untrusted certificates. The value is captured when each Key Vault client is initialized and retained for lazy certificate-chain loading, so multiple keystores can use different settings without overwriting one another. Set to `true` to disable (defaults to `false`). - Added `KeyVaultJcaPropertyNames` as the central source for the system property names supported by the Azure Key Vault JCA provider. ([#50163](https://github.com/Azure/azure-sdk-for-java/pull/50163)) +- Replaced Apache HttpClient 5 with the JDK `HttpURLConnection`, removing the Apache HttpClient and SLF4J runtime dependencies. ## 2.12.0 (2026-07-24) diff --git a/sdk/keyvault/azure-security-keyvault-jca/README.md b/sdk/keyvault/azure-security-keyvault-jca/README.md index 8109aa6a452f..8d508ca54631 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/README.md +++ b/sdk/keyvault/azure-security-keyvault-jca/README.md @@ -216,13 +216,17 @@ System.setProperty("azure.keyvault.client-id", "" System.setProperty("azure.keyvault.client-secret", ""); KeyVaultJcaProvider provider = new KeyVaultJcaProvider(); +// Register the provider before requesting its KeyStore implementation. Security.addProvider(provider); +// Load the certificate and private key that identify this server to connecting clients. KeyStore keyStore = KeyVaultKeyStore.getKeyVaultKeyStoreBySystemProperty(); +// Key managers select the server certificate and private key during each TLS handshake. KeyManagerFactory managerFactory = KeyManagerFactory.getInstance(KeyManagerFactory.getDefaultAlgorithm()); managerFactory.init(keyStore, "".toCharArray()); +// Configure one-way TLS: clients aren't required to present a certificate. SSLContext context = SSLContext.getInstance("TLS"); context.init(managerFactory.getKeyManagers(), null, null); @@ -230,13 +234,15 @@ SSLServerSocketFactory socketFactory = context.getServerSocketFactory(); SSLServerSocket serverSocket = (SSLServerSocket) socketFactory.createServerSocket(8765); while (true) { + // Accept a TLS connection and write a minimal HTTP response over it. SSLSocket socket = (SSLSocket) serverSocket.accept(); System.out.println("Client connected: " + socket.getInetAddress()); BufferedWriter out = new BufferedWriter(new OutputStreamWriter(socket.getOutputStream())); String body = "Hello, this is server."; - String response = - "HTTP/1.1 200 OK\r\n" + "Content-Type: text/plain\r\n" + "Content-Length: " + body.getBytes("UTF-8").length + "\r\n" + "Connection: close\r\n" + "\r\n" + body; + // Build a minimal HTTP response and calculate Content-Length from the UTF-8 body bytes. + String response = "HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: " + + body.getBytes(StandardCharsets.UTF_8).length + "\r\nConnection: close\r\n\r\n" + body; out.write(response); out.flush(); @@ -247,7 +253,7 @@ while (true) { **Note:** See [Authentication Methods](#authentication-methods) for configuration details. #### Client side SSL -If you are looking to integrate the JCA provider for client side socket connections, see the Apache HTTP client example below. +If you are looking to integrate the JCA provider for client side socket connections, see the HTTPS URL connection example below. ```java readme-sample-clientSSL System.setProperty("azure.keyvault.uri", ""); @@ -256,38 +262,93 @@ System.setProperty("azure.keyvault.client-id", "" System.setProperty("azure.keyvault.client-secret", ""); KeyVaultJcaProvider provider = new KeyVaultJcaProvider(); +// Register the provider before requesting its KeyStore implementation. Security.addProvider(provider); KeyStore keyStore = KeyVaultKeyStore.getKeyVaultKeyStoreBySystemProperty(); -SSLContext sslContext = SSLContexts - .custom() - .loadTrustMaterial(keyStore, new TrustSelfSignedStrategy()) - .build(); - -SSLConnectionSocketFactory sslConnectionSocketFactory = new SSLConnectionSocketFactory( - sslContext, (hostname, session) -> true); +// Create trust managers from the certificates in the Key Vault-backed KeyStore. +TrustManagerFactory trustManagerFactory + = TrustManagerFactory.getInstance(TrustManagerFactory.getDefaultAlgorithm()); +trustManagerFactory.init(keyStore); +TrustManager[] trustManagers = trustManagerFactory.getTrustManagers(); + +// The local server may use a self-signed certificate. Accept a one-certificate server chain while delegating +// validation of all other chains to the platform trust manager. Do not use this behavior in production. +for (int i = 0; i < trustManagers.length; i++) { + if (trustManagers[i] instanceof X509TrustManager) { + X509TrustManager delegate = (X509TrustManager) trustManagers[i]; + trustManagers[i] = new X509TrustManager() { + @Override + public void checkClientTrusted(X509Certificate[] chain, String authType) + throws CertificateException { + delegate.checkClientTrusted(chain, authType); + } + + @Override + public void checkServerTrusted(X509Certificate[] chain, String authType) + throws CertificateException { + if (chain.length != 1) { + delegate.checkServerTrusted(chain, authType); + } + } + + @Override + public X509Certificate[] getAcceptedIssuers() { + return delegate.getAcceptedIssuers(); + } + }; + } +} -PoolingHttpClientConnectionManager manager = new PoolingHttpClientConnectionManager( - RegistryBuilder.create() - .register("https", sslConnectionSocketFactory) - .build()); +// Configure one-way TLS: the client validates the server but doesn't present a client certificate. +SSLContext sslContext = SSLContext.getInstance("TLS"); +sslContext.init(null, trustManagers, null); String result = null; +HttpsURLConnection connection = null; +try { + // openConnection will return HttpsURLConnection when the protocol is 'https'. + connection = (HttpsURLConnection) URI.create("https://localhost:8765").toURL().openConnection(); + + // Apply the custom trust configuration to this HTTPS connection. + connection.setSSLSocketFactory(sslContext.getSocketFactory()); + // Allow the sample certificate to use a hostname other than localhost. Do not do this in production. + connection.setHostnameVerifier((hostname, session) -> true); + + connection.setRequestMethod("GET"); + int status = connection.getResponseCode(); + if (status == 200) { + // Decode the response using its declared charset, or UTF-8 when no charset is present. + Charset responseCharset = StandardCharsets.UTF_8; + String contentType = connection.getContentType(); + if (contentType != null) { + Matcher matcher = Pattern.compile("(?i)\\bcharset\\s*=\\s*\"?([^;\\s\"]+)") + .matcher(contentType); + if (matcher.find()) { + responseCharset = Charset.forName(matcher.group(1)); + } + } -try (CloseableHttpClient client = HttpClients.custom().setConnectionManager(manager).build()) { - HttpGet httpGet = new HttpGet("https://localhost:8765"); - HttpClientResponseHandler responseHandler = (ClassicHttpResponse response) -> { - int status = response.getCode(); - String result1 = "Not success"; - if (status == 200) { - result1 = EntityUtils.toString(response.getEntity()); + // Read the complete body without changing its line endings. + try (Reader reader = new InputStreamReader(connection.getInputStream(), responseCharset)) { + StringBuilder responseBody = new StringBuilder(); + char[] buffer = new char[1024]; + int read; + while ((read = reader.read(buffer)) != -1) { + responseBody.append(buffer, 0, read); + } + result = responseBody.toString(); } - return result1; - }; - result = client.execute(httpGet, responseHandler); + } else { + result = "Not success"; + } } catch (IOException ioe) { ioe.printStackTrace(); +} finally { + if (connection != null) { + connection.disconnect(); + } } System.out.println(result); ``` @@ -300,14 +361,17 @@ If you are looking to integrate the JCA provider to create an SSLServerSocket se ```java readme-sample-serverMTLS KeyVaultJcaProvider provider = new KeyVaultJcaProvider(); +// Register the provider before requesting its KeyStore implementation. Security.addProvider(provider); System.setProperty("azure.keyvault.uri", ""); System.setProperty("azure.keyvault.tenant-id", ""); System.setProperty("azure.keyvault.client-id", ""); System.setProperty("azure.keyvault.client-secret", ""); +// Load the certificate and private key that identify this server to connecting clients. KeyStore keyStore = KeyVaultKeyStore.getKeyVaultKeyStoreBySystemProperty(); +// Key managers select the server certificate and private key during each mTLS handshake. KeyManagerFactory kmf = KeyManagerFactory.getInstance(KeyManagerFactory.getDefaultAlgorithm()); kmf.init(keyStore, "".toCharArray()); @@ -315,26 +379,32 @@ System.setProperty("azure.keyvault.uri", ""); System.setProperty("azure.keyvault.tenant-id", ""); System.setProperty("azure.keyvault.client-id", ""); System.setProperty("azure.keyvault.client-secret", ""); +// Load the client certificates that this server trusts. KeyStore trustStore = KeyVaultKeyStore.getKeyVaultKeyStoreBySystemProperty(); +// Trust managers validate the certificate presented by each client. TrustManagerFactory tmf = TrustManagerFactory.getInstance(TrustManagerFactory.getDefaultAlgorithm()); tmf.init(trustStore); +// Combine the server identity with the client trust configuration. SSLContext context = SSLContext.getInstance("TLS"); context.init(kmf.getKeyManagers(), tmf.getTrustManagers(), null); SSLServerSocketFactory socketFactory = context.getServerSocketFactory(); SSLServerSocket serverSocket = (SSLServerSocket) socketFactory.createServerSocket(8765); +// Require every client to present a trusted certificate during the TLS handshake. serverSocket.setNeedClientAuth(true); while (true) { + // Accept an mTLS connection and write a minimal HTTP response over it. SSLSocket socket = (SSLSocket) serverSocket.accept(); System.out.println("Client connected: " + socket.getInetAddress()); BufferedWriter out = new BufferedWriter(new OutputStreamWriter(socket.getOutputStream())); String body = "Hello, this is server."; - String response = - "HTTP/1.1 200 OK\r\n" + "Content-Type: text/plain\r\n" + "Content-Length: " + body.getBytes("UTF-8").length + "\r\n" + "Connection: close\r\n" + "\r\n" + body; + // Build a minimal HTTP response and calculate Content-Length from the UTF-8 body bytes. + String response = "HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: " + + body.getBytes(StandardCharsets.UTF_8).length + "\r\nConnection: close\r\n\r\n" + body; out.write(response); out.flush(); @@ -345,54 +415,117 @@ while (true) { **Note:** See [Authentication Methods](#authentication-methods) for configuration details. #### Client side mTLS -If you are looking to integrate the JCA provider for client side socket connections, see the Apache HTTP client example below. +If you are looking to integrate the JCA provider for client side socket connections, see the HTTPS URL connection example below. ```java readme-sample-clientMTLS KeyVaultJcaProvider provider = new KeyVaultJcaProvider(); +// Register the provider before requesting its KeyStore implementation. Security.addProvider(provider); System.setProperty("azure.keyvault.uri", ""); System.setProperty("azure.keyvault.tenant-id", ""); System.setProperty("azure.keyvault.client-id", ""); System.setProperty("azure.keyvault.client-secret", ""); +// Load the certificate and private key that identify this client to the server. KeyStore keyStore = KeyVaultKeyStore.getKeyVaultKeyStoreBySystemProperty(); System.setProperty("azure.keyvault.uri", ""); System.setProperty("azure.keyvault.tenant-id", ""); System.setProperty("azure.keyvault.client-id", ""); System.setProperty("azure.keyvault.client-secret", ""); +// Load the server certificates that this client trusts. KeyStore trustStore = KeyVaultKeyStore.getKeyVaultKeyStoreBySystemProperty(); -SSLContext sslContext = SSLContexts - .custom() - .loadTrustMaterial(trustStore, new TrustSelfSignedStrategy()) - .loadKeyMaterial(keyStore, "".toCharArray()) - .build(); +// Create trust managers from the server trust material. +TrustManagerFactory trustManagerFactory + = TrustManagerFactory.getInstance(TrustManagerFactory.getDefaultAlgorithm()); +trustManagerFactory.init(trustStore); +TrustManager[] trustManagers = trustManagerFactory.getTrustManagers(); + +// The local server may use a self-signed certificate. Accept a one-certificate server chain while delegating +// validation of all other chains to the platform trust manager. Do not use this behavior in production. +for (int i = 0; i < trustManagers.length; i++) { + if (trustManagers[i] instanceof X509TrustManager) { + X509TrustManager delegate = (X509TrustManager) trustManagers[i]; + trustManagers[i] = new X509TrustManager() { + @Override + public void checkClientTrusted(X509Certificate[] chain, String authType) + throws CertificateException { + delegate.checkClientTrusted(chain, authType); + } + + @Override + public void checkServerTrusted(X509Certificate[] chain, String authType) + throws CertificateException { + if (chain.length != 1) { + delegate.checkServerTrusted(chain, authType); + } + } + + @Override + public X509Certificate[] getAcceptedIssuers() { + return delegate.getAcceptedIssuers(); + } + }; + } +} -SSLConnectionSocketFactory sslConnectionSocketFactory = new SSLConnectionSocketFactory( - sslContext, (hostname, session) -> true); +// Create key managers that select the client certificate and private key during the mTLS handshake. +KeyManagerFactory keyManagerFactory = KeyManagerFactory.getInstance(KeyManagerFactory.getDefaultAlgorithm()); +keyManagerFactory.init(keyStore, "".toCharArray()); -PoolingHttpClientConnectionManager manager = new PoolingHttpClientConnectionManager( - RegistryBuilder.create() - .register("https", sslConnectionSocketFactory) - .build()); +// Combine the client identity with the server trust configuration. +SSLContext sslContext = SSLContext.getInstance("TLS"); +sslContext.init(keyManagerFactory.getKeyManagers(), trustManagers, null); String result = null; +HttpsURLConnection connection = null; + +try { + // openConnection will return HttpsURLConnection when the protocol is 'https'. + connection = (HttpsURLConnection) URI.create("https://localhost:8765").toURL().openConnection(); + + // Apply the custom identity and trust configuration to this HTTPS connection. + connection.setSSLSocketFactory(sslContext.getSocketFactory()); + // Allow the sample certificate to use a hostname other than localhost. Do not do this in production. + connection.setHostnameVerifier((hostname, session) -> true); + + connection.setRequestMethod("GET"); + int status = connection.getResponseCode(); + + if (status == 200) { + // Decode the response using its declared charset, or UTF-8 when no charset is present. + Charset responseCharset = StandardCharsets.UTF_8; + String contentType = connection.getContentType(); + if (contentType != null) { + Matcher matcher = Pattern.compile("(?i)\\bcharset\\s*=\\s*\"?([^;\\s\"]+)") + .matcher(contentType); + if (matcher.find()) { + responseCharset = Charset.forName(matcher.group(1)); + } + } -try (CloseableHttpClient client = HttpClients.custom().setConnectionManager(manager).build()) { - HttpGet httpGet = new HttpGet("https://localhost:8765"); - HttpClientResponseHandler responseHandler = (ClassicHttpResponse response) -> { - int status = response.getCode(); - String result1 = "Not success"; - if (status == 200) { - result1 = EntityUtils.toString(response.getEntity()); + // Read the complete body without changing its line endings. + try (Reader reader = new InputStreamReader(connection.getInputStream(), responseCharset)) { + StringBuilder responseBody = new StringBuilder(); + char[] buffer = new char[1024]; + int read; + while ((read = reader.read(buffer)) != -1) { + responseBody.append(buffer, 0, read); + } + result = responseBody.toString(); } - return result1; - }; - result = client.execute(httpGet, responseHandler); + } else { + result = "Not success"; + } } catch (IOException ioe) { ioe.printStackTrace(); +} finally { + if (connection != null) { + connection.disconnect(); + } } + System.out.println(result); ``` @@ -735,4 +868,3 @@ This project has adopted the [Microsoft Open Source Code of Conduct][microsoft_c [microsoft_code_of_conduct]: https://opensource.microsoft.com/codeofconduct/ [non-exportable]: https://learn.microsoft.com/azure/key-vault/certificates/about-certificates#exportable-or-non-exportable-key - diff --git a/sdk/keyvault/azure-security-keyvault-jca/pom.xml b/sdk/keyvault/azure-security-keyvault-jca/pom.xml index 9a0034bcdbc0..12672e09e1b3 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/pom.xml +++ b/sdk/keyvault/azure-security-keyvault-jca/pom.xml @@ -1,5 +1,4 @@ - @@ -35,12 +34,6 @@ 2.73.11 true - - - org.apache.httpcomponents.client5 - httpclient5 - 5.4.3 - org.conscrypt @@ -55,40 +48,6 @@ 1.5.1 true - - - org.slf4j - slf4j-nop - 1.7.36 - - - - org.mockito - mockito-inline - 4.11.0 - test - - - - - net.bytebuddy - byte-buddy - 1.18.11 - test - - - net.bytebuddy - byte-buddy-agent - 1.18.11 - test - - - - com.github.spotbugs - spotbugs-annotations - 4.8.3 - test - com.azure azure-core @@ -218,14 +177,6 @@ com.azure.json com.azure.security.keyvault.jca.implementation.shaded.com.azure.json - - org.apache.hc - com.azure.security.keyvault.jca.implementation.shaded.org.apache.hc - - - org.slf4j - com.azure.security.keyvault.jca.implementation.shaded.org.slf4j - @@ -266,8 +217,6 @@ org.bouncycastle:bcpkix-lts8on:[2.73.11] org.conscrypt:conscrypt-openjdk-uber:[2.5.2] - org.apache.httpcomponents.client5:httpclient5:[5.4.3] - org.slf4j:slf4j-nop:[1.7.36] diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/KeyVaultJcaProvider.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/KeyVaultJcaProvider.java index 99e70d955eab..bfd15e14c149 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/KeyVaultJcaProvider.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/KeyVaultJcaProvider.java @@ -3,20 +3,17 @@ package com.azure.security.keyvault.jca; -import com.azure.security.keyvault.jca.implementation.signature.KeyVaultKeylessRsa256Signature; -import com.azure.security.keyvault.jca.implementation.signature.KeyVaultKeylessRsa512Signature; +import com.azure.security.keyvault.jca.implementation.signature.KeyVaultKeylessEcSha256Signature; import com.azure.security.keyvault.jca.implementation.signature.KeyVaultKeylessEcSha384Signature; import com.azure.security.keyvault.jca.implementation.signature.KeyVaultKeylessEcSha512Signature; -import com.azure.security.keyvault.jca.implementation.signature.KeyVaultKeylessEcSha256Signature; -import com.azure.security.keyvault.jca.implementation.signature.AbstractKeyVaultKeylessSignature; +import com.azure.security.keyvault.jca.implementation.signature.KeyVaultKeylessRsa256Signature; +import com.azure.security.keyvault.jca.implementation.signature.KeyVaultKeylessRsa512Signature; import com.azure.security.keyvault.jca.implementation.signature.KeyVaultKeylessRsaSsaPssSignature; -import java.lang.reflect.InvocationTargetException; import java.security.PrivilegedAction; import java.security.Provider; import java.util.Arrays; import java.util.Collections; -import java.util.stream.Stream; /** * The Azure Key Vault security provider. @@ -48,6 +45,7 @@ public final class KeyVaultJcaProvider extends Provider { /** * Constructor. */ + @SuppressWarnings("deprecation") public KeyVaultJcaProvider() { super(PROVIDER_NAME, VERSION, INFO); initialize(); @@ -74,21 +72,20 @@ private void initialize() { Collections.singletonList("DKS"), null)); putService(new Provider.Service(this, "KeyStore", KeyVaultKeyStore.ALGORITHM_NAME, KeyVaultKeyStore.class.getName(), Collections.singletonList(KeyVaultKeyStore.ALGORITHM_NAME), null)); - Stream - .of(KeyVaultKeylessRsaSsaPssSignature.class, KeyVaultKeylessRsa256Signature.class, - KeyVaultKeylessRsa512Signature.class, KeyVaultKeylessEcSha256Signature.class, - KeyVaultKeylessEcSha384Signature.class, KeyVaultKeylessEcSha512Signature.class) - .forEach(c -> putService(new Service(this, "Signature", getAlgorithmName(c), c.getName(), null, null))); + + putService(new Service(this, "Signature", KeyVaultKeylessRsaSsaPssSignature.ALGORITHM_NAME, + KeyVaultKeylessRsaSsaPssSignature.class.getName(), null, null)); + putService(new Service(this, "Signature", KeyVaultKeylessRsa256Signature.ALGORITHM_NAME, + KeyVaultKeylessRsa256Signature.class.getName(), null, null)); + putService(new Service(this, "Signature", KeyVaultKeylessRsa512Signature.ALGORITHM_NAME, + KeyVaultKeylessRsa512Signature.class.getName(), null, null)); + putService(new Service(this, "Signature", KeyVaultKeylessEcSha256Signature.ALGORITHM_NAME, + KeyVaultKeylessEcSha256Signature.class.getName(), null, null)); + putService(new Service(this, "Signature", KeyVaultKeylessEcSha384Signature.ALGORITHM_NAME, + KeyVaultKeylessEcSha384Signature.class.getName(), null, null)); + putService(new Service(this, "Signature", KeyVaultKeylessEcSha512Signature.ALGORITHM_NAME, + KeyVaultKeylessEcSha512Signature.class.getName(), null, null)); return null; }); } - - private String getAlgorithmName(Class c) { - try { - return c.getDeclaredConstructor().newInstance().getAlgorithmName(); - } catch (InstantiationException | IllegalAccessException | InvocationTargetException - | NoSuchMethodException e) { - return ""; - } - } } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/KeyVaultTrustManagerFactoryProvider.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/KeyVaultTrustManagerFactoryProvider.java index dafba85e114d..d911f502af88 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/KeyVaultTrustManagerFactoryProvider.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/KeyVaultTrustManagerFactoryProvider.java @@ -36,6 +36,7 @@ public final class KeyVaultTrustManagerFactoryProvider extends Provider { /** * Constructor. */ + @SuppressWarnings("deprecation") public KeyVaultTrustManagerFactoryProvider() { super(NAME, VERSION, INFO); initialize(); diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/KeyVaultClient.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/KeyVaultClient.java index ee5315deac1b..aa10798b7d90 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/KeyVaultClient.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/KeyVaultClient.java @@ -38,7 +38,7 @@ import java.security.spec.PKCS8EncodedKeySpec; import java.util.ArrayList; import java.util.Base64; -import java.util.HashMap; +import java.util.Collections; import java.util.List; import java.util.Map; import java.util.Optional; @@ -251,19 +251,19 @@ private AccessToken obtainAccessToken() { managedIdentity = URLEncoder.encode(managedIdentity, "UTF-8"); } - // Priority: 1. Service Principal (Client ID/Secret), 2. Workload Identity, 3. Managed Identity, 4. Provided Access Token + // Priority: 1. Service Principal, 2. Workload Identity, 3. User-assigned Managed Identity, + // 4. Provided Access Token, 5. System-assigned Managed Identity. if (tenantId != null && clientId != null && clientSecret != null) { LOGGER.info("Using client credentials (client ID/secret) for authentication"); String aadAuthenticationUri = getLoginUri(keyVaultUri + "certificates" + API_VERSION_POSTFIX, disableChallengeResourceVerification); - result - = AccessTokenUtil.getAccessToken(resource, aadAuthenticationUri, tenantId, clientId, clientSecret); + result = getAccessToken(resource, aadAuthenticationUri, tenantId, clientId, clientSecret); } else if (AccessTokenUtil.isWorkloadIdentityAvailable(clientId, tenantId)) { LOGGER.info("Using workload identity for authentication"); result = AccessTokenUtil.getAccessTokenWithWorkloadIdentity(keyVaultBaseUri, tenantId, clientId); } else if (managedIdentity != null) { LOGGER.info("Using managed identity for authentication"); - result = AccessTokenUtil.getAccessToken(resource, managedIdentity); + result = getAccessToken(resource, managedIdentity); } else if (providedAccessToken != null && !providedAccessToken.isEmpty()) { LOGGER.info("Using provided access token for authentication"); // Create an AccessToken object from the provided token string @@ -273,7 +273,7 @@ private AccessToken obtainAccessToken() { result = new AccessToken(providedAccessToken, Long.MAX_VALUE / 1000); } else { LOGGER.info("Using managed identity for authentication (default)"); - result = AccessTokenUtil.getAccessToken(resource, null); + result = getAccessToken(resource, null); } } catch (UnsupportedEncodingException e) { LOGGER.log(WARNING, "Could not obtain access token to authenticate with.", e); @@ -292,15 +292,12 @@ private AccessToken obtainAccessToken() { public List getAliases() { LOGGER.entering("KeyVaultClient", "getAliases"); - ArrayList result = new ArrayList<>(); - HashMap headers = new HashMap<>(); - - headers.put("Authorization", "Bearer " + getAccessToken()); - + List result = new ArrayList<>(); + Map headers = Collections.singletonMap("Authorization", "Bearer " + getAccessToken()); String uri = keyVaultUri + "certificates" + API_VERSION_POSTFIX; while (uri != null && !uri.isEmpty()) { - String response = HttpUtil.get(uri, headers); + String response = httpGet(uri, headers); CertificateListResult certificateListResult = null; if (response != null) { @@ -346,12 +343,8 @@ private CertificateBundle getCertificateBundle(String alias) { LOGGER.entering("KeyVaultClient", "getCertificateBundle", alias); CertificateBundle result = null; - HashMap headers = new HashMap<>(); - - headers.put("Authorization", "Bearer " + getAccessToken()); - - String uri = keyVaultUri + "certificates/" + alias + API_VERSION_POSTFIX; - String response = HttpUtil.get(uri, headers); + String response = httpGet(keyVaultUri + "certificates/" + alias + API_VERSION_POSTFIX, + Collections.singletonMap("Authorization", "Bearer " + getAccessToken())); if (response != null) { try { @@ -466,11 +459,8 @@ public Certificate[] getCertificateChainForVersion(CertificateVersion certificat return new Certificate[0]; } - HashMap headers = new HashMap<>(); - - headers.put("Authorization", "Bearer " + getAccessToken()); - - String response = HttpUtil.get(certificateVersion.getSecretId() + API_VERSION_POSTFIX, headers); + String response = httpGet(certificateVersion.getSecretId() + API_VERSION_POSTFIX, + Collections.singletonMap("Authorization", "Bearer " + getAccessToken())); if (response == null) { throw new IllegalStateException("Failed to load certificate chain response for alias: " + alias); @@ -533,7 +523,8 @@ public Key getKeyForVersion(CertificateVersion certificateVersion, char[] passwo if (!exportable) { // Keyless signing uses the versioned key ID instead of exporting private key material. - String keyAlgorithm = keyType.contains("-HSM") ? keyType.substring(0, keyType.indexOf("-HSM")) : keyType; + String keyAlgorithm + = keyType != null && keyType.contains("-HSM") ? keyType.substring(0, keyType.indexOf("-HSM")) : keyType; KeyVaultPrivateKey key = Optional.ofNullable(certificateVersion.getKeyId()) .map(keyId -> new KeyVaultPrivateKey(keyAlgorithm, keyId, this)) @@ -548,11 +539,9 @@ public Key getKeyForVersion(CertificateVersion certificateVersion, char[] passwo if (certificateSecretUri == null) { return null; } - Map headers = new HashMap<>(); - - headers.put("Authorization", "Bearer " + getAccessToken()); - String body = HttpUtil.get(certificateSecretUri + API_VERSION_POSTFIX, headers); + String body = httpGet(certificateSecretUri + API_VERSION_POSTFIX, + Collections.singletonMap("Authorization", "Bearer " + getAccessToken())); if (body == null) { // If the private key is not available the certificate cannot be used for server side certificates or mTLS. @@ -621,13 +610,10 @@ public byte[] getSignedWithPrivateKey(String digestName, String digestValue, Str LOGGER.entering("KeyVaultClient", "getSignedWithPrivateKey", new Object[] { digestName, digestValue, keyId }); SignResult result = null; - String bodyString = String.format("{\"alg\": \"" + digestName + "\", \"value\": \"%s\"}", digestValue); - Map headers = new HashMap<>(); - - headers.put("Authorization", "Bearer " + getAccessToken()); - + String bodyString = "{\"alg\": \"" + digestName + "\", \"value\": \"" + digestValue + "\"}"; + Map headers = Collections.singletonMap("Authorization", "Bearer " + getAccessToken()); String uri = keyId + "/sign" + API_VERSION_POSTFIX; - String response = HttpUtil.post(uri, headers, bodyString, "application/json"); + String response = httpPost(uri, headers, bodyString); if (response != null) { try { @@ -702,4 +688,21 @@ private PrivateKey createPrivateKeyFromPem(String pemString, String keyType) return privateKey; } + + String httpGet(String uri, Map headers) { + return HttpUtil.get(uri, headers); + } + + String httpPost(String uri, Map headers, String body) { + return HttpUtil.post(uri, headers, body, "application/json"); + } + + AccessToken getAccessToken(String resource, String identity) { + return AccessTokenUtil.getAccessToken(resource, identity); + } + + AccessToken getAccessToken(String resource, String aadAuthenticationUri, String tenantId, String clientId, + String clientSecret) { + return AccessTokenUtil.getAccessToken(resource, aadAuthenticationUri, tenantId, clientId, clientSecret); + } } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/certificates/ClasspathCertificates.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/certificates/ClasspathCertificates.java index 1714a3c0334f..b1758198be00 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/certificates/ClasspathCertificates.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/certificates/ClasspathCertificates.java @@ -119,7 +119,7 @@ public void deleteEntry(String alias) { */ public void loadCertificatesFromClasspath() { try { - String[] filenames = getFilenames("/keyvault"); + String[] filenames = getFilenames(); for (String filename : filenames) { try (InputStream inputStream = getClass().getResourceAsStream("/keyvault/" + filename)) { String alias = filename; @@ -147,13 +147,12 @@ public void loadCertificatesFromClasspath() { /** * Get the filenames. * - * @param path the path. * @return the filenames. * @throws IOException when an I/O error occurs. */ - private String[] getFilenames(String path) throws IOException { + private String[] getFilenames() throws IOException { List filenames = new ArrayList<>(); - try (InputStream in = getClass().getResourceAsStream(path)) { + try (InputStream in = getClass().getResourceAsStream("/keyvault")) { if (!Objects.isNull(in)) { try (BufferedReader br = new BufferedReader(new InputStreamReader(in, StandardCharsets.UTF_8))) { String resource; 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 813261322fba..0af06cdabc88 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 @@ -115,7 +115,13 @@ public final class KeyVaultCertificates implements AzureCertificates { * @param parameter The Key Vault load-store configuration. */ public KeyVaultCertificates(KeyVaultLoadStoreParameter parameter) { - updateKeyVaultClient(parameter); + this(parameter, parameter.getUri() == null ? null : new KeyVaultClient(parameter)); + } + + KeyVaultCertificates(KeyVaultLoadStoreParameter parameter, KeyVaultClient keyVaultClient) { + Objects.requireNonNull(parameter, "'parameter' cannot be null."); + updateCertificateConfiguration(parameter); + setKeyVaultClient(keyVaultClient); } private void updateCertificateConfiguration(KeyVaultLoadStoreParameter parameter) { @@ -615,20 +621,18 @@ public String refreshAndGetAliasByCertificate(Certificate certificate) { aliasesSnapshot = new ArrayList<>(aliases); } - aliasesSnapshot.forEach(this::loadCertificateIfNeeded); - - Map certificatesSnapshot; - synchronized (this) { - certificatesSnapshot = new HashMap<>(certificates); + for (String alias : aliasesSnapshot) { + loadCertificateIfNeeded(alias); + Certificate loadedCertificate; + synchronized (this) { + loadedCertificate = certificates.get(alias); + } + if (certificate.equals(loadedCertificate)) { + return alias; + } } - return certificatesSnapshot.entrySet() - .stream() - .filter(entry -> certificate.equals(entry.getValue())) - .findFirst() - .map(Map.Entry::getKey) - .orElse(null); - + return null; } /** diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/certificates/SpecificPathCertificates.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/certificates/SpecificPathCertificates.java index 30d66f230760..a68cbacf7e2f 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/certificates/SpecificPathCertificates.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/certificates/SpecificPathCertificates.java @@ -19,9 +19,7 @@ import java.util.List; import java.util.Map; import java.util.Objects; -import java.util.Optional; import java.util.logging.Logger; -import java.util.stream.Stream; import static com.azure.security.keyvault.jca.implementation.utils.CertificateUtil.loadX509CertificateFromFile; import static com.azure.security.keyvault.jca.implementation.utils.CertificateUtil.loadX509CertificatesFromFile; @@ -203,13 +201,11 @@ private List getFiles() { List files = new ArrayList<>(); File filePackage = new File(certificatePath); File[] array = filePackage.listFiles(); - Optional.ofNullable(array) - .map(Arrays::stream) - .orElseGet(Stream::empty) - .filter(Objects::nonNull) - .filter(File::isFile) - .filter(File::exists) - .filter(File::canRead) + if (array == null) { + return files; + } + Arrays.stream(array) + .filter(file -> Objects.nonNull(file) && file.exists() && file.isFile() && file.canRead()) .forEach(files::add); return files; } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSha256Signature.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSha256Signature.java index a936e2722bdd..564a66586518 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSha256Signature.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSha256Signature.java @@ -6,10 +6,14 @@ * key vault SHA256 */ public final class KeyVaultKeylessEcSha256Signature extends KeyVaultKeylessEcSignature { + /** + * Algorithm name used by this implementation. + */ + public static final String ALGORITHM_NAME = "SHA256withECDSA"; @Override public String getAlgorithmName() { - return "SHA256withECDSA"; + return ALGORITHM_NAME; } /** diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSha384Signature.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSha384Signature.java index 0d267dfcd973..81eea1c06d32 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSha384Signature.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSha384Signature.java @@ -6,10 +6,14 @@ * key vault SHA384 */ public final class KeyVaultKeylessEcSha384Signature extends KeyVaultKeylessEcSignature { + /** + * Algorithm name used by this implementation. + */ + public static final String ALGORITHM_NAME = "SHA384withECDSA"; @Override public String getAlgorithmName() { - return "SHA384withECDSA"; + return ALGORITHM_NAME; } /** diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSha512Signature.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSha512Signature.java index 23b9ff7b443e..4accc59c0b15 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSha512Signature.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSha512Signature.java @@ -6,10 +6,14 @@ * key vault SHA512 */ public final class KeyVaultKeylessEcSha512Signature extends KeyVaultKeylessEcSignature { + /** + * Algorithm name used by this implementation. + */ + public static final String ALGORITHM_NAME = "SHA512withECDSA"; @Override public String getAlgorithmName() { - return "SHA512withECDSA"; + return ALGORITHM_NAME; } /** diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsa256Signature.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsa256Signature.java index 3454bdb4716f..525be4c4f64a 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsa256Signature.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsa256Signature.java @@ -7,6 +7,15 @@ * key vault Rsa signature to support key less */ public class KeyVaultKeylessRsa256Signature extends KeyVaultKeylessRsaSignature { + /** + * Algorithm name used by this implementation. + */ + public static final String ALGORITHM_NAME = "SHA256withRSA"; + + @Override + public String getAlgorithmName() { + return ALGORITHM_NAME; + } /** * Construct a new KeyVaultKeyLessRsaSignature @@ -14,9 +23,4 @@ public class KeyVaultKeylessRsa256Signature extends KeyVaultKeylessRsaSignature public KeyVaultKeylessRsa256Signature() { super("SHA-256", "RS256"); } - - @Override - public String getAlgorithmName() { - return "SHA256withRSA"; - } } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsa512Signature.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsa512Signature.java index 3115057ecbef..d145aa323cb0 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsa512Signature.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsa512Signature.java @@ -7,6 +7,15 @@ * key vault Rsa signature to support key less */ public class KeyVaultKeylessRsa512Signature extends KeyVaultKeylessRsaSignature { + /** + * Algorithm name used by this implementation. + */ + public static final String ALGORITHM_NAME = "SHA512withRSA"; + + @Override + public String getAlgorithmName() { + return ALGORITHM_NAME; + } /** * Construct a new KeyVaultKeyLessRsaSignature @@ -14,9 +23,4 @@ public class KeyVaultKeylessRsa512Signature extends KeyVaultKeylessRsaSignature public KeyVaultKeylessRsa512Signature() { super("SHA-512", "RS512"); } - - @Override - public String getAlgorithmName() { - return "SHA512withRSA"; - } } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsaSignature.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsaSignature.java index cc737a0750f0..83e56830a4ee 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsaSignature.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsaSignature.java @@ -12,7 +12,6 @@ * key vault Rsa signature to support key less */ abstract class KeyVaultKeylessRsaSignature extends AbstractKeyVaultKeylessSignature { - private final String keyVaultDigestName; /** diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsaSsaPssSignature.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsaSsaPssSignature.java index d15df8aec525..0cff98ef0d35 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsaSsaPssSignature.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsaSsaPssSignature.java @@ -13,6 +13,15 @@ * key vault Rsa signature to support key less */ public class KeyVaultKeylessRsaSsaPssSignature extends KeyVaultKeylessRsaSignature { + /** + * Algorithm name used by this implementation. + */ + public static final String ALGORITHM_NAME = "RSASSA-PSS"; + + @Override + public String getAlgorithmName() { + return ALGORITHM_NAME; + } /** * Construct a new KeyVaultKeyLessRsaSignature @@ -42,9 +51,4 @@ protected void engineSetParameter(AlgorithmParameterSpec params) throws InvalidA } } } - - @Override - public String getAlgorithmName() { - return "RSASSA-PSS"; - } } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/AccessTokenUtil.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/AccessTokenUtil.java index d15f15ac0619..74acc9e3e4fb 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/AccessTokenUtil.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/AccessTokenUtil.java @@ -4,7 +4,6 @@ import com.azure.security.keyvault.jca.KeyVaultJcaPropertyNames; import com.azure.security.keyvault.jca.implementation.model.AccessToken; -import org.apache.hc.core5.http.ClassicHttpResponse; import java.io.IOException; import java.io.UnsupportedEncodingException; @@ -13,6 +12,7 @@ import java.net.URLEncoder; import java.util.Collections; import java.util.HashMap; +import java.util.List; import java.util.Locale; import java.util.Map; import java.util.logging.Logger; @@ -97,8 +97,10 @@ public static AccessToken getAccessToken(String resource, String identity) { /* * App Service 2017-09-01: MSI_ENDPOINT, MSI_SECRET - * Azure Container App 2019-08-01: IDENTITY_ENDPOINT, IDENTITY_HEADER, see more from https://learn.microsoft.com/en-us/azure/container-apps/managed-identity?tabs=cli%2Chttp#rest-endpoint-reference - * Azure Virtual Machine 2018-02-01, see more from https://learn.microsoft.com/en-us/entra/identity/managed-identities-azure-resources/how-to-use-vm-token#get-a-token-using-http + * Azure Container App 2019-08-01: IDENTITY_ENDPOINT, IDENTITY_HEADER, see more from + * https://learn.microsoft.com/azure/container-apps/managed-identity?tabs=cli%2Chttp#rest-endpoint-reference + * Azure Virtual Machine 2018-02-01, see more from + * https://learn.microsoft.com/entra/identity/managed-identities-azure-resources/how-to-use-vm-token#get-a-token-using-http */ if (System.getenv("WEBSITE_SITE_NAME") != null && !System.getenv("WEBSITE_SITE_NAME").isEmpty()) { result = getAccessTokenOnAppService(resource, identity); @@ -125,21 +127,21 @@ public static AccessToken getAccessToken(String resource, String identity) { */ public static AccessToken getAccessToken(String resource, String aadAuthenticationUrl, String tenantId, String clientId, String clientSecret) { + return getAccessToken(resource, aadAuthenticationUrl, tenantId, clientId, clientSecret, HttpUtil::post); + } + + // Overloaded method that allows specifying a custom HttpPoster for testing. + static AccessToken getAccessToken(String resource, String aadAuthenticationUrl, String tenantId, String clientId, + String clientSecret, HttpPoster httpPoster) { // The client secret is deliberately left out: entering() renders every parameter in clear text. LOGGER.entering("AccessTokenUtil", "getAccessToken", new Object[] { resource, tenantId, clientId }); LOGGER.info("Getting access token using client ID / client secret"); AccessToken result = null; - StringBuilder oauth2Url = new StringBuilder(); - - if (aadAuthenticationUrl == null) { - oauth2Url.append(OAUTH2_TOKEN_BASE_URL).append(tenantId).append("/"); - } else { - oauth2Url.append(addTrailingSlashIfRequired(aadAuthenticationUrl)); - } - - oauth2Url.append(OAUTH2_TOKEN_POSTFIX); + String oauth2Url = (aadAuthenticationUrl == null) + ? OAUTH2_TOKEN_BASE_URL + tenantId + "/" + OAUTH2_TOKEN_POSTFIX + : addTrailingSlashIfRequired(aadAuthenticationUrl) + OAUTH2_TOKEN_POSTFIX; String encodedClientSecret = ""; @@ -149,17 +151,10 @@ public static AccessToken getAccessToken(String resource, String aadAuthenticati LOGGER.log(WARNING, "Failed to encode client secret for access token request", e); } - StringBuilder requestBody = new StringBuilder(); - - requestBody.append(GRANT_TYPE_FRAGMENT) - .append(CLIENT_ID_FRAGMENT) - .append(clientId) - .append(CLIENT_SECRET_FRAGMENT) - .append(encodedClientSecret) - .append(RESOURCE_FRAGMENT) - .append(resource); + String requestBody = GRANT_TYPE_FRAGMENT + CLIENT_ID_FRAGMENT + clientId + CLIENT_SECRET_FRAGMENT + + encodedClientSecret + RESOURCE_FRAGMENT + resource; - String body = HttpUtil.post(oauth2Url.toString(), requestBody.toString(), "application/x-www-form-urlencoded"); + String body = httpPoster.post(oauth2Url, null, requestBody, "application/x-www-form-urlencoded"); if (body != null) { try { @@ -174,6 +169,26 @@ public static AccessToken getAccessToken(String resource, String aadAuthenticati return result; } + /** + * Functional interface for making HTTP POST requests. + * + *

Introduced to be used for testing purposes, allowing the HTTP POST behavior to be mocked or overridden. + */ + @FunctionalInterface + interface HttpPoster { + String post(String uri, Map headers, String body, String contentType); + } + + /** + * Functional interface for retrieving HTTP response headers. + * + *

Introduced to be used for testing purposes, allowing the HTTP response headers to be mocked or overridden. + */ + @FunctionalInterface + interface HttpHeaderRetriever { + Map> get(String uri); + } + /** * Check if Azure Workload Identity environment is available. * @@ -200,7 +215,6 @@ public static boolean isWorkloadIdentityAvailable(String clientId, String tenant * @param resource The resource scope (will be appended with /.default if not already present). * @param tenantId Tenant ID to use. If blank, fallback to environment variable AZURE_TENANT_ID. * @param clientId Client ID of the managed identity to use. If blank, fallback to environment variable AZURE_CLIENT_ID. - * @param tokenFilePath Path to the federated token file. If blank, fallback to environment variable AZURE_FEDERATED_TOKEN_FILE. * @return An access token, or null if the operation fails. */ public static AccessToken getAccessTokenWithWorkloadIdentity(String resource, String tenantId, String clientId) { @@ -231,7 +245,7 @@ public static AccessToken getAccessTokenWithWorkloadIdentity(String resource, St String requestUrl = buildTokenRequestUrl(authorityHost, effectiveTenantId); String requestBody = buildTokenRequestBody(effectiveClientId, federatedToken, scope); - String response = HttpUtil.post(requestUrl, requestBody, "application/x-www-form-urlencoded"); + String response = HttpUtil.post(requestUrl, null, requestBody, "application/x-www-form-urlencoded"); AccessToken result = parseAccessTokenResponse(response); LOGGER.exiting("AccessTokenUtil", "getAccessTokenWithWorkloadIdentity", result); @@ -324,17 +338,13 @@ private static AccessToken getAccessTokenOnAppService(String resource, String cl LOGGER.info("Getting access token using managed identity based on MSI_SECRET"); AccessToken result = null; - StringBuilder url = new StringBuilder(); - - url.append(System.getenv("MSI_ENDPOINT")) - .append("?api-version=2017-09-01") - .append(RESOURCE_FRAGMENT) - .append(resource); - + String url; if (clientId != null) { - url.append("&clientid=").append(clientId); - + url = System.getenv("MSI_ENDPOINT") + "?api-version=2017-09-01" + RESOURCE_FRAGMENT + resource + + "&clientid=" + clientId; LOGGER.log(INFO, "Using managed identity with client ID: {0}", clientId); + } else { + url = System.getenv("MSI_ENDPOINT") + "?api-version=2017-09-01" + RESOURCE_FRAGMENT + resource; } HashMap headers = new HashMap<>(); @@ -342,7 +352,7 @@ private static AccessToken getAccessTokenOnAppService(String resource, String cl headers.put("Metadata", "true"); headers.put("Secret", System.getenv("MSI_SECRET")); - String body = HttpUtil.get(url.toString(), headers); + String body = HttpUtil.get(url, headers); if (body != null) { try { @@ -381,9 +391,9 @@ private static AccessToken getAccessTokenOnContainerApp(String resource, String LOGGER.log(INFO, "Using managed identity with client ID: {0}", clientId); } - Map headers = new HashMap<>(); + Map headers = Collections.emptyMap(); if (System.getenv(PROPERTY_IDENTITY_HEADER) != null && !System.getenv(PROPERTY_IDENTITY_HEADER).isEmpty()) { - headers.put("X-IDENTITY-HEADER", System.getenv(PROPERTY_IDENTITY_HEADER)); + headers = Collections.singletonMap("X-IDENTITY-HEADER", System.getenv(PROPERTY_IDENTITY_HEADER)); } String body = HttpUtil.get(url.toString(), headers); @@ -426,11 +436,7 @@ private static AccessToken getAccessTokenOnOthers(String resource, String identi url.append("&object_id=").append(identity); } - HashMap headers = new HashMap<>(); - - headers.put("Metadata", "true"); - - String body = HttpUtil.get(url.toString(), headers); + String body = HttpUtil.get(url.toString(), Collections.singletonMap("Metadata", "true")); if (body != null) { try { @@ -446,17 +452,25 @@ private static AccessToken getAccessTokenOnOthers(String resource, String identi } public static String getLoginUri(String resourceUri, boolean disableChallengeResourceVerification) { + return getLoginUri(resourceUri, disableChallengeResourceVerification, HttpUtil::getWithOnlyResponseHeaders); + } + + // Overloaded method that allows specifying a custom HttpHeaderRetriever for testing. + static String getLoginUri(String resourceUri, boolean disableChallengeResourceVerification, + HttpHeaderRetriever httpHeaderRetriever) { LOGGER.entering("AccessTokenUtil", "getLoginUri", resourceUri); LOGGER.log(INFO, "Getting login URI using: {0}", resourceUri); - ClassicHttpResponse response = HttpUtil.getWithResponse(resourceUri, null); + Map> headers = httpHeaderRetriever.get(resourceUri); - if (response == null) { + if (headers == null) { throw new IllegalStateException("Could not obtain login URI to retrieve access token from."); } - Map challengeAttributes - = extractChallengeAttributes(response.getFirstHeader(WWW_AUTHENTICATE).getValue()); + List wwwAuthenticates = headers.get(WWW_AUTHENTICATE); + String wwwAuthenticate + = wwwAuthenticates == null || wwwAuthenticates.isEmpty() ? null : wwwAuthenticates.get(0); + Map challengeAttributes = extractChallengeAttributes(wwwAuthenticate); String scope = challengeAttributes.get("resource"); if (scope != null) { @@ -517,7 +531,7 @@ private static Map extractChallengeAttributes(String authenticat for (String pair : attributes) { String[] keyValue = pair.split("="); - attributeMap.put(keyValue[0].replaceAll("\"", ""), keyValue[1].replaceAll("\"", "")); + attributeMap.put(keyValue[0].replace("\"", ""), keyValue[1].replace("\"", "")); } LOGGER.exiting("AccessTokenUtil", "extractChallengeAttributes", attributeMap); @@ -535,7 +549,7 @@ private static Map extractChallengeAttributes(String authenticat private static boolean isBearerChallenge(String authenticateHeader) { return authenticateHeader != null && !authenticateHeader.isEmpty() - && authenticateHeader.toLowerCase(Locale.ROOT).startsWith(BEARER_TOKEN_PREFIX.toLowerCase(Locale.ROOT)); + && BEARER_TOKEN_PREFIX.regionMatches(true, 0, authenticateHeader, 0, BEARER_TOKEN_PREFIX.length()); } /** diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/AiaCertificateChainUtil.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/AiaCertificateChainUtil.java index f97beed113fb..db0cc0f8c44e 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/AiaCertificateChainUtil.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/AiaCertificateChainUtil.java @@ -58,6 +58,10 @@ final class AiaCertificateChainUtil { private static final AiaResponseCache AIA_CACHE = new AiaResponseCache(AIA_CACHE_MAX_SIZE, System::currentTimeMillis, (message, parameters) -> LOGGER.logp(FINE, AiaResponseCache.class.getName(), "diagnostic", message, parameters)); + // A default HTTP-based response loader for AIA requests. + private static final AiaResponseLoader DEFAULT_RESPONSE_LOADER = HttpUtil::getAiaBytesWithMetadata; + // The currently configured response loader, which can be overridden for tests. + private static volatile AiaResponseLoader responseLoader = DEFAULT_RESPONSE_LOADER; /** * Determines whether a certificate chain should be completed with issuer certificates resolved via AIA. @@ -450,7 +454,7 @@ && isCurrentlyValid(candidate) private static AiaResponseCache.Entry loadAiaResponse(String url) { LOGGER.log(FINE, "Downloading issuer certificate from AIA URL: {0}", url); long now = System.currentTimeMillis(); - HttpUtil.BinaryHttpResponse response = HttpUtil.getBytesWithMetadata(url); + HttpUtil.BinaryHttpResponse response = responseLoader.load(url); byte[] certBytes = response.getBody(); if (certBytes == null) { LOGGER.log(FINE, "Failed to download issuer certificate from AIA URL: {0}", url); @@ -595,6 +599,38 @@ static void clearAiaCache() { AIA_CACHE.clear(); } + /** + * Sets the response loader for AIA requests. + * + *

Introduced to be used for testing purposes, allowing the HTTP response loading behavior to be mocked or + * overridden. + * + * @param loader the response loader to use + */ + static synchronized void setResponseLoader(AiaResponseLoader loader) { + responseLoader = Objects.requireNonNull(loader, "'loader' cannot be null."); + } + + /** + * Resets the response loader for AIA requests to the default HTTP-based loader. + * + *

Introduced to be used for testing purposes, allowing the HTTP response loading behavior to be reset to the + * default. + */ + static synchronized void resetResponseLoader() { + responseLoader = DEFAULT_RESPONSE_LOADER; + } + + /** + * Functional interface for loading AIA responses. + * + *

Introduced to be used for testing purposes, allowing the HTTP POST behavior to be mocked or overridden. + */ + @FunctionalInterface + interface AiaResponseLoader { + HttpUtil.BinaryHttpResponse load(String url); + } + /** * Parses the certificates contained in an AIA response, which may be DER- or PEM-encoded and may hold a bundle * rather than a single certificate. diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/CertificateUtil.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/CertificateUtil.java index 027f81075d67..36068c3e47b5 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/CertificateUtil.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/CertificateUtil.java @@ -32,7 +32,6 @@ import java.util.List; import java.util.Map; import java.util.logging.Logger; -import java.util.stream.Collectors; import static java.util.logging.Level.FINE; @@ -155,11 +154,7 @@ public static Certificate loadX509CertificateFromFile(InputStream inputStream) t public static Certificate[] loadX509CertificatesFromFile(InputStream inputStream) throws CertificateException { CertificateFactory factory = CertificateFactory.getInstance("X.509"); - return factory.generateCertificates(inputStream) - .stream() - .map(o -> (Certificate) o) - .collect(Collectors.toList()) - .toArray(new Certificate[0]); + return factory.generateCertificates(inputStream).toArray(new Certificate[0]); } public static String getCertificateNameFromCertificateItemId(String id) { diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/HttpUtil.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/HttpUtil.java index 51834978bc8a..456f30611cb5 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/HttpUtil.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/HttpUtil.java @@ -3,39 +3,28 @@ package com.azure.security.keyvault.jca.implementation.utils; import com.azure.security.keyvault.jca.implementation.JreKeyStoreFactory; -import org.apache.hc.client5.http.classic.methods.HttpGet; -import org.apache.hc.client5.http.classic.methods.HttpPost; -import org.apache.hc.client5.http.config.RequestConfig; -import org.apache.hc.client5.http.impl.classic.CloseableHttpClient; -import org.apache.hc.client5.http.impl.classic.HttpClients; -import org.apache.hc.client5.http.impl.io.PoolingHttpClientConnectionManager; -import org.apache.hc.client5.http.socket.ConnectionSocketFactory; -import org.apache.hc.client5.http.socket.PlainConnectionSocketFactory; -import org.apache.hc.client5.http.ssl.SSLConnectionSocketFactory; -import org.apache.hc.core5.http.ClassicHttpResponse; -import org.apache.hc.core5.http.ContentType; -import org.apache.hc.core5.http.Header; -import org.apache.hc.core5.http.HttpEntity; -import org.apache.hc.core5.http.config.RegistryBuilder; -import org.apache.hc.core5.http.io.HttpClientResponseHandler; -import org.apache.hc.core5.http.io.entity.EntityUtils; -import org.apache.hc.core5.http.io.entity.StringEntity; -import org.apache.hc.core5.ssl.SSLContexts; -import org.apache.hc.core5.util.Timeout; - -import javax.net.ssl.HostnameVerifier; + +import javax.net.ssl.HttpsURLConnection; import javax.net.ssl.SSLContext; +import javax.net.ssl.TrustManagerFactory; import java.io.BufferedReader; +import java.io.ByteArrayOutputStream; import java.io.IOException; +import java.io.InputStream; import java.io.InputStreamReader; +import java.io.OutputStream; +import java.net.HttpURLConnection; import java.net.URI; import java.net.URISyntaxException; +import java.nio.charset.StandardCharsets; import java.security.KeyManagementException; -import java.security.KeyStore; import java.security.KeyStoreException; import java.security.NoSuchAlgorithmException; +import java.util.List; import java.util.Map; import java.util.Optional; +import java.util.TreeMap; +import java.util.concurrent.TimeUnit; import java.util.logging.Logger; import java.util.stream.Stream; @@ -43,7 +32,7 @@ import static java.util.logging.Level.WARNING; /** - * The RestClient that uses the Apache HttpClient class. + * The REST client that uses the JDK {@link HttpURLConnection} class. */ public final class HttpUtil { public static final String DEFAULT_VERSION = "unknown"; @@ -61,92 +50,284 @@ public final class HttpUtil { private static final Logger LOGGER = Logger.getLogger(HttpUtil.class.getName()); - public static String get(String uri, Map headers) { - String result = null; + static final int HTTP_TIMEOUT_IN_MILLISECONDS = 180_000; + private static final int AIA_HTTP_TIMEOUT_IN_MILLISECONDS = 10_000; + static final int MAX_AIA_RESPONSE_SIZE_IN_BYTES = 10 * 1024 * 1024; + private static final int AIA_HTTP_TOTAL_TIMEOUT_IN_MILLISECONDS = 30_000; + private static final int MAX_AIA_REDIRECTS = 5; - try (CloseableHttpClient client = buildClient()) { - HttpGet httpGet = new HttpGet(uri); + /** + * Functional interface for opening HTTP connections. + * + *

Introduced to be used for testing purposes, allowing the HTTP connection behavior to be mocked or overridden. + */ + @FunctionalInterface + interface ConnectionFactory { + HttpURLConnection open(String url) throws IOException; + } - if (headers != null) { - headers.forEach(httpGet::addHeader); - } + /** + * Performs an HTTP GET request to the specified URI with the given headers. + * + * @param uri the URI to send the GET request to + * @param headers the headers to include in the request + * @return the response body as a string, or {@code null} if the request fails + * @throws RuntimeException if the server returns a non-successful response + */ + public static String get(String uri, Map headers) { + return get(uri, headers, HttpUtil::openConnection); + } - httpGet.addHeader(USER_AGENT_KEY, USER_AGENT_VALUE); + // Overloaded method that allows specifying a custom ConnectionFactory for testing purposes. + static String get(String uri, Map headers, ConnectionFactory connectionFactory) { + HttpURLConnection connection = null; - result = client.execute(httpGet, createResponseHandler()); + try { + connection = connectionFactory.open(uri); + configureConnection(connection, "GET", headers); + + ensureSuccessfulResponse(connection.getResponseCode()); + return readResponseBody(connection); } catch (IOException ioe) { LOGGER.log(WARNING, "Unable to finish the HTTP GET request.", ioe); + + return null; + } finally { + if (connection != null) { + connection.disconnect(); + } } + } - return result; + /** + * Performs an HTTP POST request to the specified URI with the given headers, body, and content type. + * + * @param uri the URI to send the POST request to + * @param headers the headers to include in the request + * @param body the body of the POST request + * @param contentType the content type of the POST request body + * @return the response body as a string, or {@code null} if the request fails + * @throws RuntimeException if the server returns a non-successful response + */ + public static String post(String uri, Map headers, String body, String contentType) { + return post(uri, headers, body, contentType, HttpUtil::openConnection); + } + + // Overloaded method that allows specifying a custom ConnectionFactory for testing purposes. + static String post(String uri, Map headers, String body, String contentType, + ConnectionFactory connectionFactory) { + + HttpURLConnection connection = null; + + try { + connection = connectionFactory.open(uri); + configureConnection(connection, "POST", headers); + connection.setDoOutput(true); + + if (contentType != null) { + connection.setRequestProperty("Content-Type", contentType); + } + + try (OutputStream outputStream = connection.getOutputStream()) { + outputStream.write(body.getBytes(StandardCharsets.UTF_8)); + } + + ensureSuccessfulResponse(connection.getResponseCode()); + return readResponseBody(connection); + } catch (IOException ioe) { + LOGGER.log(WARNING, "Unable to finish the HTTP POST request.", ioe); + + return null; + } finally { + if (connection != null) { + connection.disconnect(); + } + } } /** - * Performs an HTTP GET request and returns the raw response body as a byte array. - * Used primarily for downloading DER-encoded certificates from CA Issuers URLs in - * AIA (Authority Information Access) certificate extensions. + * Performs an HTTP GET request and returns the response body along with HTTP metadata. Used primarily for + * downloading DER-encoded certificates from CA Issuers URLs in AIA (Authority Information Access) certificate + * extensions. * * @param url the URL to fetch * @return the response body bytes, or {@code null} if the request fails or returns non-2xx */ - public static byte[] getBytes(String url) { - return getBytesWithMetadata(url).body; + public static BinaryHttpResponse getAiaBytesWithMetadata(String url) { + return getAiaBytesWithMetadata(url, HttpUtil::openConnection); } - static BinaryHttpResponse getBytesWithMetadata(String url) { - try (CloseableHttpClient client = buildClient()) { - HttpGet httpGet = new HttpGet(url); - httpGet.addHeader(USER_AGENT_KEY, USER_AGENT_VALUE); - // Set reasonable timeouts to prevent indefinite hangs when fetching AIA certificate chain - RequestConfig config = RequestConfig.custom() - .setConnectTimeout(Timeout.ofSeconds(10)) - .setResponseTimeout(Timeout.ofSeconds(10)) - .build(); - httpGet.setConfig(config); - return client.execute(httpGet, response -> toBinaryResponse(response, url)); - } catch (Exception e) { - // Catch all exceptions including IOException, IllegalArgumentException (malformed URL), - // and other runtime exceptions that may occur during HTTP execution. - // Gracefully return null to allow AIA completion to fail silently instead of breaking - // the entire jarsigner/signing operation. - LOGGER.log(WARNING, e, () -> "Unable to finish the HTTP GET (bytes) request for URL: " + url); + // Overloaded method that allows specifying a custom ConnectionFactory for testing purposes. + static BinaryHttpResponse getAiaBytesWithMetadata(String url, ConnectionFactory connectionFactory) { + String currentUrl; + + try { + currentUrl = validateAiaUrl(url); + } catch (IllegalArgumentException e) { + LOGGER.log(WARNING, "Unable to finish the HTTP GET (bytes) request for URL: " + url, e); + return BinaryHttpResponse.empty(); } + + // HttpURLConnection does not follow redirects between different protocols (e.g., HTTP to HTTPS) + // reliably, so we handle redirects ourselves. AIA uses a small number of redirects, if any. + for (int redirectCount = 0; redirectCount <= MAX_AIA_REDIRECTS; redirectCount++) { + HttpURLConnection connection = null; + + try { + connection = connectionFactory.open(currentUrl); + + connection.setInstanceFollowRedirects(false); + connection.setRequestMethod("GET"); + connection.setConnectTimeout(AIA_HTTP_TIMEOUT_IN_MILLISECONDS); + connection.setReadTimeout(AIA_HTTP_TIMEOUT_IN_MILLISECONDS); + connection.setRequestProperty(USER_AGENT_KEY, USER_AGENT_VALUE); + + int status = connection.getResponseCode(); + String cacheControl = getCombinedHeaderValue(connection, "Cache-Control"); + String date = connection.getHeaderField("Date"); + String age = connection.getHeaderField("Age"); + String expires = connection.getHeaderField("Expires"); + + if (isRedirect(status)) { + String location = connection.getHeaderField("Location"); + + if (location == null || redirectCount == MAX_AIA_REDIRECTS) { + LOGGER.log(WARNING, "HTTP GET redirect could not be followed for URL: {0}", currentUrl); + + return new BinaryHttpResponse(null, cacheControl, date, age, expires); + } + + currentUrl = resolveAiaRedirect(currentUrl, location); + + continue; + } + + if (status < 200 || status >= 300) { + LOGGER.log(WARNING, "HTTP GET returned status {0} for URL: {1}", + new Object[] { status, currentUrl }); + + return new BinaryHttpResponse(null, cacheControl, date, age, expires); + } + + long contentLength = connection.getContentLengthLong(); + + if (contentLength > MAX_AIA_RESPONSE_SIZE_IN_BYTES) { + LOGGER.log(WARNING, "AIA response exceeded the maximum size for URL: {0}", currentUrl); + + return new BinaryHttpResponse(null, cacheControl, date, age, expires); + } + + return new BinaryHttpResponse(readResponseBytes(connection.getInputStream(), currentUrl), cacheControl, + date, age, expires); + } catch (IOException | IllegalArgumentException | ClassCastException e) { + // Catch all exceptions including IOException, IllegalArgumentException, and other runtime exceptions + // that may occur during HTTP execution. Gracefully return null to allow AIA completion to fail silently + // the entire jarsigner/signing operation. + LOGGER.log(WARNING, "Unable to finish the HTTP GET (bytes) request for URL: " + currentUrl, e); + + return BinaryHttpResponse.empty(); + } finally { + if (connection != null) { + connection.disconnect(); + } + } + } + + return BinaryHttpResponse.empty(); } - static BinaryHttpResponse toBinaryResponse(ClassicHttpResponse response, String url) throws IOException { - int status = response.getCode(); - if (status < 200 || status >= 300) { - LOGGER.log(WARNING, "HTTP GET returned status {0} for URL: {1}", new Object[] { status, url }); - return new BinaryHttpResponse(null, getCombinedHeaderValue(response, "Cache-Control"), - getHeaderValue(response, "Date"), getHeaderValue(response, "Age"), getHeaderValue(response, "Expires")); + private static boolean isRedirect(int status) { + return status == HttpURLConnection.HTTP_MOVED_PERM + || status == HttpURLConnection.HTTP_MOVED_TEMP + || status == HttpURLConnection.HTTP_SEE_OTHER + || status == 307 + || status == 308; + } + + private static String resolveAiaRedirect(String currentUrl, String location) { + if (location.startsWith("?")) { + int queryIndex = currentUrl.indexOf('?'); + int fragmentIndex = currentUrl.indexOf('#'); + int suffixIndex + = queryIndex < 0 ? fragmentIndex : fragmentIndex < 0 ? queryIndex : Math.min(queryIndex, fragmentIndex); + String currentUrlWithoutSuffix = suffixIndex < 0 ? currentUrl : currentUrl.substring(0, suffixIndex); + + return validateAiaUrl(currentUrlWithoutSuffix + location); + } + + return validateAiaUrl(URI.create(currentUrl).resolve(location).toString()); + } + + private static String validateAiaUrl(String url) { + URI uri = URI.create(url); + String scheme = uri.getScheme(); + + if (!"http".equalsIgnoreCase(scheme) && !"https".equalsIgnoreCase(scheme)) { + throw new IllegalArgumentException("AIA URL must use HTTP or HTTPS."); } - HttpEntity entity = response.getEntity(); - byte[] body = entity != null ? EntityUtils.toByteArray(entity) : null; - return new BinaryHttpResponse(body, getCombinedHeaderValue(response, "Cache-Control"), - getHeaderValue(response, "Date"), getHeaderValue(response, "Age"), getHeaderValue(response, "Expires")); + return uri.toString(); } - private static String getCombinedHeaderValue(ClassicHttpResponse response, String name) { - Header[] headers = response.getHeaders(name); - if (headers.length == 0) { + private static byte[] readResponseBytes(InputStream inputStream, String url) throws IOException { + long deadline = System.nanoTime() + TimeUnit.MILLISECONDS.toNanos(AIA_HTTP_TOTAL_TIMEOUT_IN_MILLISECONDS); + + try (InputStream responseBody = inputStream; ByteArrayOutputStream outputStream = new ByteArrayOutputStream()) { + byte[] buffer = new byte[4096]; + int totalBytesRead = 0; + int read; + + while ((read = responseBody.read(buffer)) != -1) { + if (read > MAX_AIA_RESPONSE_SIZE_IN_BYTES - totalBytesRead) { + LOGGER.log(WARNING, "AIA response exceeded the maximum size for URL: {0}", url); + + return null; + } + + outputStream.write(buffer, 0, read); + + totalBytesRead += read; + + if (System.nanoTime() > deadline) { + LOGGER.log(WARNING, "AIA response exceeded the maximum download time for URL: {0}", url); + + return null; + } + } + + return outputStream.toByteArray(); + } + } + + private static String getCombinedHeaderValue(HttpURLConnection connection, String name) { + Map> headers = connection.getHeaderFields(); + + if (headers == null) { return null; } StringBuilder value = new StringBuilder(); - for (Header header : headers) { - if (value.length() > 0) { - value.append(", "); + + for (Map.Entry> entry : headers.entrySet()) { + if (entry.getKey() == null || !name.equalsIgnoreCase(entry.getKey()) || entry.getValue() == null) { + continue; + } + + for (String headerValue : entry.getValue()) { + if (headerValue == null) { + continue; + } + + if (value.length() > 0) { + value.append(", "); + } + + value.append(headerValue); } - value.append(header.getValue()); } - return value.toString(); - } - private static String getHeaderValue(ClassicHttpResponse response, String name) { - Header header = response.getFirstHeader(name); - return header == null ? null : header.getValue(); + return value.length() == 0 ? null : value.toString(); } static final class BinaryHttpResponse { @@ -189,10 +370,6 @@ String getExpires() { } } - public static String post(String uri, String body, String contentType) { - return post(uri, null, body, contentType); - } - public static String getUserAgentPrefix() { return Optional.of(HttpUtil.class) .map(Class::getClassLoader) @@ -205,99 +382,117 @@ public static String getUserAgentPrefix() { .orElse(DEFAULT_USER_AGENT_VALUE_PREFIX); } - public static String post(String uri, Map headers, String body, String contentType) { - String result = null; + private static String createErrorMessage(int status) { + return "Failed to get a response from Key Vault because the HTTP status code was " + status + ". This may be " + + "caused by missing permissions or roles. For instructions on assigning permissions or roles, see " + + "https://github.com/Azure/azure-sdk-for-java/tree/main/sdk/keyvault/azure-security-keyvault-jca#prerequisites."; + } - try (CloseableHttpClient client = buildClient()) { - HttpPost httpPost = new HttpPost(uri); + private static void ensureSuccessfulResponse(int status) { + if (status < 200 || status >= 300) { + String errorMessage = createErrorMessage(status); + LOGGER.log(SEVERE, errorMessage); + throw new RuntimeException(errorMessage); + } + } + + @SuppressWarnings("StringOperationCanBeSimplified") + private static String readResponseBody(HttpURLConnection connection) throws IOException { + try (InputStream responseBody = connection.getInputStream(); + ByteArrayOutputStream outputStream = new ByteArrayOutputStream()) { - httpPost.addHeader(USER_AGENT_KEY, USER_AGENT_VALUE); + if (responseBody == null) { - if (headers != null) { - headers.forEach(httpPost::addHeader); - httpPost.addHeader("Content-Type", contentType); + return null; } - httpPost.setEntity(new StringEntity(body, ContentType.create(contentType))); + byte[] buffer = new byte[4096]; + int read; - result = client.execute(httpPost, createResponseHandler()); - } catch (IOException ioe) { - LOGGER.log(WARNING, "Unable to finish the HTTP POST request.", ioe); - } + while ((read = responseBody.read(buffer)) != -1) { + outputStream.write(buffer, 0, read); + } - return result; + return new String(outputStream.toByteArray(), StandardCharsets.UTF_8); + } } - public static ClassicHttpResponse getWithResponse(String uri, Map headers) { - ClassicHttpResponse result = null; + /** + * Retrieves only the response headers from an HTTP GET request to the specified URI. + * + * @param uri the URI to send the HTTP GET request to + * @return a map of response headers, or null if the request was not successful + */ + public static Map> getWithOnlyResponseHeaders(String uri) { + return getWithOnlyResponseHeaders(uri, HttpUtil::openConnection); + } - try (CloseableHttpClient client = buildClient()) { - HttpGet httpGet = new HttpGet(uri); + // Overloaded method that allows specifying a custom ConnectionFactory for testing purposes. + static Map> getWithOnlyResponseHeaders(String uri, ConnectionFactory connectionFactory) { + HttpURLConnection connection = null; - if (headers != null) { - headers.forEach(httpGet::addHeader); + try { + connection = connectionFactory.open(uri); + configureConnection(connection, "GET", null); + + if (connection.getResponseCode() == 401) { + Map> headers = new TreeMap<>(String.CASE_INSENSITIVE_ORDER); + Map> responseHeaders = connection.getHeaderFields(); + + if (responseHeaders != null) { + responseHeaders.forEach((name, values) -> { + if (name != null) { + headers.put(name, values); + } + }); + } + + return headers; } - httpGet.addHeader(USER_AGENT_KEY, USER_AGENT_VALUE); - - result = client.execute(httpGet, createResponseHandlerForAuthChallenge()); + return null; } catch (IOException ioe) { LOGGER.log(WARNING, "Unable to finish the HTTP GET request.", ioe); - } - - return result; - } - private static HttpClientResponseHandler createResponseHandler() { - return (ClassicHttpResponse response) -> { - int status = response.getCode(); - String result; - - if (status >= 200 && status < 300) { - HttpEntity entity = response.getEntity(); - result = entity != null ? EntityUtils.toString(entity) : null; - } else { - String errorMessage = "Fail to get response from Key Vault because return http status code is " + status - + ". It " - + "can be caused by missing permissions or roles. To know how to add permissions or roles, see " - + "https://github.com/Azure/azure-sdk-for-java/tree/main/sdk/keyvault/azure-security-keyvault-jca#prerequisites."; - LOGGER.log(SEVERE, errorMessage); - throw new RuntimeException(errorMessage); + return null; + } finally { + if (connection != null) { + connection.disconnect(); } - - return result; - }; + } } - private static HttpClientResponseHandler createResponseHandlerForAuthChallenge() { - return (ClassicHttpResponse response) -> { - int status = response.getCode(); - - return status == 401 ? response : null; - }; + private static void configureConnection(HttpURLConnection connection, String method, Map headers) + throws IOException { + connection.setRequestMethod(method); + connection.setConnectTimeout(HTTP_TIMEOUT_IN_MILLISECONDS); + connection.setReadTimeout(HTTP_TIMEOUT_IN_MILLISECONDS); + if (headers != null) { + headers.forEach(connection::setRequestProperty); + } + connection.setRequestProperty(USER_AGENT_KEY, USER_AGENT_VALUE); } - private static CloseableHttpClient buildClient() { - KeyStore keyStore = JreKeyStoreFactory.getDefaultKeyStore(); + private static HttpURLConnection openConnection(String uri) throws IOException { + HttpURLConnection connection = (HttpURLConnection) URI.create(uri).toURL().openConnection(); - SSLContext sslContext = null; + if (connection instanceof HttpsURLConnection) { + try { + TrustManagerFactory trustManagerFactory + = TrustManagerFactory.getInstance(TrustManagerFactory.getDefaultAlgorithm()); - try { - sslContext = SSLContexts.custom().loadTrustMaterial(keyStore, null).build(); - } catch (NoSuchAlgorithmException | KeyManagementException | KeyStoreException e) { - LOGGER.log(WARNING, "Unable to build the SSL context.", e); - } + trustManagerFactory.init(JreKeyStoreFactory.getDefaultKeyStore()); - SSLConnectionSocketFactory sslConnectionSocketFactory - = new SSLConnectionSocketFactory(sslContext, (HostnameVerifier) null); + SSLContext sslContext = SSLContext.getInstance("TLS"); - PoolingHttpClientConnectionManager manager - = new PoolingHttpClientConnectionManager(RegistryBuilder.create() - .register("http", PlainConnectionSocketFactory.getSocketFactory()) - .register("https", sslConnectionSocketFactory) - .build()); + sslContext.init(null, trustManagerFactory.getTrustManagers(), null); + ((HttpsURLConnection) connection).setSSLSocketFactory(sslContext.getSocketFactory()); + } catch (KeyManagementException | KeyStoreException | NoSuchAlgorithmException e) { + LOGGER.log(WARNING, "Unable to build the SSL context.", e); + } + } - return HttpClients.custom().setConnectionManager(manager).build(); + return connection; } public static String validateUri(String uri, String propertyName) { diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/JsonConverterUtil.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/JsonConverterUtil.java index 66d071587a93..cc1038706ce1 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/JsonConverterUtil.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/java/com/azure/security/keyvault/jca/implementation/utils/JsonConverterUtil.java @@ -6,10 +6,8 @@ import com.azure.json.JsonProviders; import com.azure.json.JsonReader; import com.azure.json.JsonSerializable; -import com.azure.json.JsonWriter; import com.azure.json.ReadValueCallback; -import java.io.ByteArrayOutputStream; import java.io.IOException; import java.util.logging.Logger; @@ -56,27 +54,19 @@ public static > T fromJson(ReadValueCallback jsonSerializable) { LOGGER.entering("JsonConverterUtil", "toJson", jsonSerializable); - - if (jsonSerializable == null) { - return null; - } - - try (ByteArrayOutputStream byteArrayOutputStream = new ByteArrayOutputStream(); - JsonWriter jsonWriter = JsonProviders.createWriter(byteArrayOutputStream)) { - - jsonWriter.writeUntyped(jsonSerializable); - jsonWriter.flush(); - - return byteArrayOutputStream.toString("UTF-8"); - } catch (IOException e) { - LOGGER.log(WARNING, "Unable to convert to JSON", e); + String value = null; + + if (jsonSerializable != null) { + try { + value = jsonSerializable.toJsonString(); + } catch (IOException e) { + LOGGER.log(WARNING, "Unable to convert to JSON", e); + } } LOGGER.exiting("JsonConverterUtil", "toJson"); - - return null; + return value; } } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/main/resources/module-info.java b/sdk/keyvault/azure-security-keyvault-jca/src/main/resources/module-info.java index 74b0860c1d80..47964348d0ad 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/main/resources/module-info.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/main/resources/module-info.java @@ -3,18 +3,9 @@ /** - * Because of maven-shade-plugin, we have some special requirement of module-info: - * 1. When compile current module by java 9. - * module-info.java must have content like "requires org.apache.httpcomponents.client5.httpclient5;". - * 2. When used by other java 9 modules, - * module-info.java must NOT have content like "requires org.apache.httpcomponents.client5.httpclient5;". - * - * To achieve this, we need these steps: - * 1. Avoid compile error in other module by deleting contents like "requires org.apache.httpcomponents.client5.httpclient5;". - * 2. Avoid compile error in current module by moving module-info.java out of src/main/java/. - * 3. Package module-info.class into xxx.jar by configuring moditect-maven-plugin. - * 4. Package module-info.java into xxx-source.jar by putting module-info.java in src/main/resources/. - * 5. Exclude module-info.java into xxx.jar by configuring maven-jar-plugin. + * Dependencies are relocated by the maven-shade-plugin, so this descriptor is kept outside src/main/java and added + * to the shaded JAR by the moditect-maven-plugin. Keeping it in src/main/resources also includes the descriptor in + * the sources JAR, while the maven-jar-plugin excludes the uncompiled source from the binary JAR. */ module com.azure.security.keyvault.jca { requires java.logging; diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/mtls/ClientMTLSSample.java b/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/mtls/ClientMTLSSample.java index c4884dbd2ed0..d3003ed55b8a 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/mtls/ClientMTLSSample.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/mtls/ClientMTLSSample.java @@ -4,23 +4,25 @@ import com.azure.security.keyvault.jca.KeyVaultJcaProvider; import com.azure.security.keyvault.jca.KeyVaultKeyStore; -import org.apache.hc.client5.http.classic.methods.HttpGet; -import org.apache.hc.client5.http.impl.classic.CloseableHttpClient; -import org.apache.hc.client5.http.impl.classic.HttpClients; -import org.apache.hc.client5.http.impl.io.PoolingHttpClientConnectionManager; -import org.apache.hc.client5.http.socket.ConnectionSocketFactory; -import org.apache.hc.client5.http.ssl.SSLConnectionSocketFactory; -import org.apache.hc.client5.http.ssl.TrustSelfSignedStrategy; -import org.apache.hc.core5.http.ClassicHttpResponse; -import org.apache.hc.core5.http.config.RegistryBuilder; -import org.apache.hc.core5.http.io.HttpClientResponseHandler; -import org.apache.hc.core5.http.io.entity.EntityUtils; -import org.apache.hc.core5.ssl.SSLContexts; +import javax.net.ssl.HttpsURLConnection; +import javax.net.ssl.KeyManagerFactory; import javax.net.ssl.SSLContext; +import javax.net.ssl.TrustManager; +import javax.net.ssl.TrustManagerFactory; +import javax.net.ssl.X509TrustManager; import java.io.IOException; +import java.io.InputStreamReader; +import java.io.Reader; +import java.net.URI; +import java.nio.charset.Charset; +import java.nio.charset.StandardCharsets; import java.security.KeyStore; import java.security.Security; +import java.security.cert.CertificateException; +import java.security.cert.X509Certificate; +import java.util.regex.Matcher; +import java.util.regex.Pattern; /** * The ClientMTLS sample. @@ -30,52 +32,114 @@ public class ClientMTLSSample { public static void main(String[] args) throws Exception { // BEGIN: readme-sample-clientMTLS KeyVaultJcaProvider provider = new KeyVaultJcaProvider(); + // Register the provider before requesting its KeyStore implementation. Security.addProvider(provider); System.setProperty("azure.keyvault.uri", ""); System.setProperty("azure.keyvault.tenant-id", ""); System.setProperty("azure.keyvault.client-id", ""); System.setProperty("azure.keyvault.client-secret", ""); + // Load the certificate and private key that identify this client to the server. KeyStore keyStore = KeyVaultKeyStore.getKeyVaultKeyStoreBySystemProperty(); System.setProperty("azure.keyvault.uri", ""); System.setProperty("azure.keyvault.tenant-id", ""); System.setProperty("azure.keyvault.client-id", ""); System.setProperty("azure.keyvault.client-secret", ""); + // Load the server certificates that this client trusts. KeyStore trustStore = KeyVaultKeyStore.getKeyVaultKeyStoreBySystemProperty(); - SSLContext sslContext = SSLContexts - .custom() - .loadTrustMaterial(trustStore, new TrustSelfSignedStrategy()) - .loadKeyMaterial(keyStore, "".toCharArray()) - .build(); + // Create trust managers from the server trust material. + TrustManagerFactory trustManagerFactory + = TrustManagerFactory.getInstance(TrustManagerFactory.getDefaultAlgorithm()); + trustManagerFactory.init(trustStore); + TrustManager[] trustManagers = trustManagerFactory.getTrustManagers(); - SSLConnectionSocketFactory sslConnectionSocketFactory = new SSLConnectionSocketFactory( - sslContext, (hostname, session) -> true); + // The local server may use a self-signed certificate. Accept a one-certificate server chain while delegating + // validation of all other chains to the platform trust manager. Do not use this behavior in production. + for (int i = 0; i < trustManagers.length; i++) { + if (trustManagers[i] instanceof X509TrustManager) { + X509TrustManager delegate = (X509TrustManager) trustManagers[i]; + trustManagers[i] = new X509TrustManager() { + @Override + public void checkClientTrusted(X509Certificate[] chain, String authType) + throws CertificateException { + delegate.checkClientTrusted(chain, authType); + } - PoolingHttpClientConnectionManager manager = new PoolingHttpClientConnectionManager( - RegistryBuilder.create() - .register("https", sslConnectionSocketFactory) - .build()); + @Override + public void checkServerTrusted(X509Certificate[] chain, String authType) + throws CertificateException { + if (chain.length != 1) { + delegate.checkServerTrusted(chain, authType); + } + } + + @Override + public X509Certificate[] getAcceptedIssuers() { + return delegate.getAcceptedIssuers(); + } + }; + } + } + + // Create key managers that select the client certificate and private key during the mTLS handshake. + KeyManagerFactory keyManagerFactory = KeyManagerFactory.getInstance(KeyManagerFactory.getDefaultAlgorithm()); + keyManagerFactory.init(keyStore, "".toCharArray()); + + // Combine the client identity with the server trust configuration. + SSLContext sslContext = SSLContext.getInstance("TLS"); + sslContext.init(keyManagerFactory.getKeyManagers(), trustManagers, null); String result = null; + HttpsURLConnection connection = null; + + try { + // openConnection will return HttpsURLConnection when the protocol is 'https'. + connection = (HttpsURLConnection) URI.create("https://localhost:8765").toURL().openConnection(); + + // Apply the custom identity and trust configuration to this HTTPS connection. + connection.setSSLSocketFactory(sslContext.getSocketFactory()); + // Allow the sample certificate to use a hostname other than localhost. Do not do this in production. + connection.setHostnameVerifier((hostname, session) -> true); - try (CloseableHttpClient client = HttpClients.custom().setConnectionManager(manager).build()) { - HttpGet httpGet = new HttpGet("https://localhost:8765"); - HttpClientResponseHandler responseHandler = (ClassicHttpResponse response) -> { - int status = response.getCode(); - String result1 = "Not success"; - if (status == 200) { - result1 = EntityUtils.toString(response.getEntity()); + connection.setRequestMethod("GET"); + int status = connection.getResponseCode(); + + if (status == 200) { + // Decode the response using its declared charset, or UTF-8 when no charset is present. + Charset responseCharset = StandardCharsets.UTF_8; + String contentType = connection.getContentType(); + if (contentType != null) { + Matcher matcher = Pattern.compile("(?i)\\bcharset\\s*=\\s*\"?([^;\\s\"]+)") + .matcher(contentType); + if (matcher.find()) { + responseCharset = Charset.forName(matcher.group(1)); + } + } + + // Read the complete body without changing its line endings. + try (Reader reader = new InputStreamReader(connection.getInputStream(), responseCharset)) { + StringBuilder responseBody = new StringBuilder(); + char[] buffer = new char[1024]; + int read; + while ((read = reader.read(buffer)) != -1) { + responseBody.append(buffer, 0, read); + } + result = responseBody.toString(); } - return result1; - }; - result = client.execute(httpGet, responseHandler); + } else { + result = "Not success"; + } } catch (IOException ioe) { ioe.printStackTrace(); + } finally { + if (connection != null) { + connection.disconnect(); + } } + System.out.println(result); // END: readme-sample-clientMTLS } - } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/mtls/ServerMTLSSample.java b/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/mtls/ServerMTLSSample.java index 119272fa6edd..574d322522ec 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/mtls/ServerMTLSSample.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/mtls/ServerMTLSSample.java @@ -13,6 +13,7 @@ import javax.net.ssl.TrustManagerFactory; import java.io.BufferedWriter; import java.io.OutputStreamWriter; +import java.nio.charset.StandardCharsets; import java.security.KeyStore; import java.security.Security; @@ -24,14 +25,17 @@ public class ServerMTLSSample { public static void main(String[] args) throws Exception { // BEGIN: readme-sample-serverMTLS KeyVaultJcaProvider provider = new KeyVaultJcaProvider(); + // Register the provider before requesting its KeyStore implementation. Security.addProvider(provider); System.setProperty("azure.keyvault.uri", ""); System.setProperty("azure.keyvault.tenant-id", ""); System.setProperty("azure.keyvault.client-id", ""); System.setProperty("azure.keyvault.client-secret", ""); + // Load the certificate and private key that identify this server to connecting clients. KeyStore keyStore = KeyVaultKeyStore.getKeyVaultKeyStoreBySystemProperty(); + // Key managers select the server certificate and private key during each mTLS handshake. KeyManagerFactory kmf = KeyManagerFactory.getInstance(KeyManagerFactory.getDefaultAlgorithm()); kmf.init(keyStore, "".toCharArray()); @@ -39,26 +43,32 @@ public static void main(String[] args) throws Exception { System.setProperty("azure.keyvault.tenant-id", ""); System.setProperty("azure.keyvault.client-id", ""); System.setProperty("azure.keyvault.client-secret", ""); + // Load the client certificates that this server trusts. KeyStore trustStore = KeyVaultKeyStore.getKeyVaultKeyStoreBySystemProperty(); + // Trust managers validate the certificate presented by each client. TrustManagerFactory tmf = TrustManagerFactory.getInstance(TrustManagerFactory.getDefaultAlgorithm()); tmf.init(trustStore); + // Combine the server identity with the client trust configuration. SSLContext context = SSLContext.getInstance("TLS"); context.init(kmf.getKeyManagers(), tmf.getTrustManagers(), null); SSLServerSocketFactory socketFactory = context.getServerSocketFactory(); SSLServerSocket serverSocket = (SSLServerSocket) socketFactory.createServerSocket(8765); + // Require every client to present a trusted certificate during the TLS handshake. serverSocket.setNeedClientAuth(true); while (true) { + // Accept an mTLS connection and write a minimal HTTP response over it. SSLSocket socket = (SSLSocket) serverSocket.accept(); System.out.println("Client connected: " + socket.getInetAddress()); BufferedWriter out = new BufferedWriter(new OutputStreamWriter(socket.getOutputStream())); String body = "Hello, this is server."; - String response = - "HTTP/1.1 200 OK\r\n" + "Content-Type: text/plain\r\n" + "Content-Length: " + body.getBytes("UTF-8").length + "\r\n" + "Connection: close\r\n" + "\r\n" + body; + // Build a minimal HTTP response and calculate Content-Length from the UTF-8 body bytes. + String response = "HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: " + + body.getBytes(StandardCharsets.UTF_8).length + "\r\nConnection: close\r\n\r\n" + body; out.write(response); out.flush(); diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/tls/ClientSSLSample.java b/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/tls/ClientSSLSample.java index 5e0d4205ce2a..1094d3f1f256 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/tls/ClientSSLSample.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/tls/ClientSSLSample.java @@ -4,23 +4,24 @@ import com.azure.security.keyvault.jca.KeyVaultJcaProvider; import com.azure.security.keyvault.jca.KeyVaultKeyStore; -import org.apache.hc.client5.http.classic.methods.HttpGet; -import org.apache.hc.client5.http.impl.classic.CloseableHttpClient; -import org.apache.hc.client5.http.impl.classic.HttpClients; -import org.apache.hc.client5.http.impl.io.PoolingHttpClientConnectionManager; -import org.apache.hc.client5.http.socket.ConnectionSocketFactory; -import org.apache.hc.client5.http.ssl.SSLConnectionSocketFactory; -import org.apache.hc.client5.http.ssl.TrustSelfSignedStrategy; -import org.apache.hc.core5.http.ClassicHttpResponse; -import org.apache.hc.core5.http.config.RegistryBuilder; -import org.apache.hc.core5.http.io.HttpClientResponseHandler; -import org.apache.hc.core5.http.io.entity.EntityUtils; -import org.apache.hc.core5.ssl.SSLContexts; +import javax.net.ssl.HttpsURLConnection; import javax.net.ssl.SSLContext; +import javax.net.ssl.TrustManager; +import javax.net.ssl.TrustManagerFactory; +import javax.net.ssl.X509TrustManager; import java.io.IOException; +import java.io.InputStreamReader; +import java.io.Reader; +import java.net.URI; +import java.nio.charset.Charset; +import java.nio.charset.StandardCharsets; import java.security.KeyStore; import java.security.Security; +import java.security.cert.CertificateException; +import java.security.cert.X509Certificate; +import java.util.regex.Matcher; +import java.util.regex.Pattern; /** * The ClientSSL sample. @@ -35,38 +36,93 @@ public static void main(String[] args) throws Exception { System.setProperty("azure.keyvault.client-secret", ""); KeyVaultJcaProvider provider = new KeyVaultJcaProvider(); + // Register the provider before requesting its KeyStore implementation. Security.addProvider(provider); KeyStore keyStore = KeyVaultKeyStore.getKeyVaultKeyStoreBySystemProperty(); - SSLContext sslContext = SSLContexts - .custom() - .loadTrustMaterial(keyStore, new TrustSelfSignedStrategy()) - .build(); + // Create trust managers from the certificates in the Key Vault-backed KeyStore. + TrustManagerFactory trustManagerFactory + = TrustManagerFactory.getInstance(TrustManagerFactory.getDefaultAlgorithm()); + trustManagerFactory.init(keyStore); + TrustManager[] trustManagers = trustManagerFactory.getTrustManagers(); - SSLConnectionSocketFactory sslConnectionSocketFactory = new SSLConnectionSocketFactory( - sslContext, (hostname, session) -> true); + // The local server may use a self-signed certificate. Accept a one-certificate server chain while delegating + // validation of all other chains to the platform trust manager. Do not use this behavior in production. + for (int i = 0; i < trustManagers.length; i++) { + if (trustManagers[i] instanceof X509TrustManager) { + X509TrustManager delegate = (X509TrustManager) trustManagers[i]; + trustManagers[i] = new X509TrustManager() { + @Override + public void checkClientTrusted(X509Certificate[] chain, String authType) + throws CertificateException { + delegate.checkClientTrusted(chain, authType); + } - PoolingHttpClientConnectionManager manager = new PoolingHttpClientConnectionManager( - RegistryBuilder.create() - .register("https", sslConnectionSocketFactory) - .build()); + @Override + public void checkServerTrusted(X509Certificate[] chain, String authType) + throws CertificateException { + if (chain.length != 1) { + delegate.checkServerTrusted(chain, authType); + } + } + + @Override + public X509Certificate[] getAcceptedIssuers() { + return delegate.getAcceptedIssuers(); + } + }; + } + } + + // Configure one-way TLS: the client validates the server but doesn't present a client certificate. + SSLContext sslContext = SSLContext.getInstance("TLS"); + sslContext.init(null, trustManagers, null); String result = null; + HttpsURLConnection connection = null; + try { + // openConnection will return HttpsURLConnection when the protocol is 'https'. + connection = (HttpsURLConnection) URI.create("https://localhost:8765").toURL().openConnection(); + + // Apply the custom trust configuration to this HTTPS connection. + connection.setSSLSocketFactory(sslContext.getSocketFactory()); + // Allow the sample certificate to use a hostname other than localhost. Do not do this in production. + connection.setHostnameVerifier((hostname, session) -> true); + + connection.setRequestMethod("GET"); + int status = connection.getResponseCode(); + if (status == 200) { + // Decode the response using its declared charset, or UTF-8 when no charset is present. + Charset responseCharset = StandardCharsets.UTF_8; + String contentType = connection.getContentType(); + if (contentType != null) { + Matcher matcher = Pattern.compile("(?i)\\bcharset\\s*=\\s*\"?([^;\\s\"]+)") + .matcher(contentType); + if (matcher.find()) { + responseCharset = Charset.forName(matcher.group(1)); + } + } - try (CloseableHttpClient client = HttpClients.custom().setConnectionManager(manager).build()) { - HttpGet httpGet = new HttpGet("https://localhost:8765"); - HttpClientResponseHandler responseHandler = (ClassicHttpResponse response) -> { - int status = response.getCode(); - String result1 = "Not success"; - if (status == 200) { - result1 = EntityUtils.toString(response.getEntity()); + // Read the complete body without changing its line endings. + try (Reader reader = new InputStreamReader(connection.getInputStream(), responseCharset)) { + StringBuilder responseBody = new StringBuilder(); + char[] buffer = new char[1024]; + int read; + while ((read = reader.read(buffer)) != -1) { + responseBody.append(buffer, 0, read); + } + result = responseBody.toString(); } - return result1; - }; - result = client.execute(httpGet, responseHandler); + } else { + result = "Not success"; + } } catch (IOException ioe) { ioe.printStackTrace(); + } finally { + if (connection != null) { + connection.disconnect(); + } } System.out.println(result); // END: readme-sample-clientSSL diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/tls/ServerSSLSample.java b/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/tls/ServerSSLSample.java index f925dd04a75c..614bee61948d 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/tls/ServerSSLSample.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/samples/java/com/azure/security/keyvault/jca/tls/ServerSSLSample.java @@ -12,6 +12,7 @@ import javax.net.ssl.SSLSocket; import java.io.BufferedWriter; import java.io.OutputStreamWriter; +import java.nio.charset.StandardCharsets; import java.security.KeyStore; import java.security.Security; @@ -28,13 +29,17 @@ public static void main(String[] args) throws Exception { System.setProperty("azure.keyvault.client-secret", ""); KeyVaultJcaProvider provider = new KeyVaultJcaProvider(); + // Register the provider before requesting its KeyStore implementation. Security.addProvider(provider); + // Load the certificate and private key that identify this server to connecting clients. KeyStore keyStore = KeyVaultKeyStore.getKeyVaultKeyStoreBySystemProperty(); + // Key managers select the server certificate and private key during each TLS handshake. KeyManagerFactory managerFactory = KeyManagerFactory.getInstance(KeyManagerFactory.getDefaultAlgorithm()); managerFactory.init(keyStore, "".toCharArray()); + // Configure one-way TLS: clients aren't required to present a certificate. SSLContext context = SSLContext.getInstance("TLS"); context.init(managerFactory.getKeyManagers(), null, null); @@ -42,13 +47,15 @@ public static void main(String[] args) throws Exception { SSLServerSocket serverSocket = (SSLServerSocket) socketFactory.createServerSocket(8765); while (true) { + // Accept a TLS connection and write a minimal HTTP response over it. SSLSocket socket = (SSLSocket) serverSocket.accept(); System.out.println("Client connected: " + socket.getInetAddress()); BufferedWriter out = new BufferedWriter(new OutputStreamWriter(socket.getOutputStream())); String body = "Hello, this is server."; - String response = - "HTTP/1.1 200 OK\r\n" + "Content-Type: text/plain\r\n" + "Content-Length: " + body.getBytes("UTF-8").length + "\r\n" + "Connection: close\r\n" + "\r\n" + body; + // Build a minimal HTTP response and calculate Content-Length from the UTF-8 body bytes. + String response = "HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: " + + body.getBytes(StandardCharsets.UTF_8).length + "\r\nConnection: close\r\n\r\n" + body; out.write(response); out.flush(); diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/JcaTestUtils.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/JcaTestUtils.java new file mode 100644 index 000000000000..0014f1b83915 --- /dev/null +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/JcaTestUtils.java @@ -0,0 +1,168 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +package com.azure.security.keyvault.jca; + +import javax.net.ssl.KeyManager; +import javax.net.ssl.KeyManagerFactory; +import javax.net.ssl.SSLEngine; +import javax.net.ssl.TrustManager; +import javax.net.ssl.TrustManagerFactory; +import javax.net.ssl.X509ExtendedKeyManager; +import javax.net.ssl.X509TrustManager; +import java.net.Socket; +import java.security.KeyStore; +import java.security.KeyStoreException; +import java.security.NoSuchAlgorithmException; +import java.security.Principal; +import java.security.PrivateKey; +import java.security.UnrecoverableKeyException; +import java.security.cert.CertificateException; +import java.security.cert.X509Certificate; +import java.util.function.BiFunction; +import java.util.function.BiPredicate; + +/** + * Utility methods for testing KeyVault JCA. + */ +public final class JcaTestUtils { + /** + * Loads {@link TrustManager TrustManagers}. + * + * @param keyStore The {@link KeyStore}. + * @param trustStrategy An optional predicate that is used to skip calling + * {@link X509TrustManager#checkServerTrusted(X509Certificate[], String)}. + * @return The {@link TrustManager TrustManagers}. + * @throws NoSuchAlgorithmException If the algorithm used when calling + * {@link TrustManagerFactory#getInstance(String)} doesn't exist. + * @throws KeyStoreException If {@link TrustManagerFactory#init(KeyStore)} fails. + */ + public static TrustManager[] loadTrustMaterial(KeyStore keyStore, + BiPredicate trustStrategy) throws NoSuchAlgorithmException, KeyStoreException { + TrustManagerFactory tmFactory = TrustManagerFactory.getInstance(TrustManagerFactory.getDefaultAlgorithm()); + tmFactory.init(keyStore); + TrustManager[] trustManagers = tmFactory.getTrustManagers(); + + if (trustManagers != null && trustStrategy != null) { + for (int i = 0; i < trustManagers.length; i++) { + TrustManager trustManager = trustManagers[i]; + if (trustManager instanceof X509TrustManager) { + trustManagers[i] = new TrustManagerDelegate((X509TrustManager) trustManager, trustStrategy); + } + } + } + + return trustManagers; + } + + /** + * Loads {@link KeyManager KeyManagers}. + * + * @param keyStore The {@link KeyStore}. + * @param aliasStrategy An optional function to handle aliasing. + * @return The {@link KeyManager KeyManagers}. + * @throws NoSuchAlgorithmException If the algorithm used when calling {@link KeyManagerFactory#getInstance(String)} + * doesn't exist. + * @throws KeyStoreException If {@link KeyManagerFactory#init(KeyStore, char[])} fails. + * @throws UnrecoverableKeyException If the {@link KeyManager} can't be recovered when calling + * {@link KeyManagerFactory#init(KeyStore, char[])}, such as the {@code password is wrong}. + */ + public static KeyManager[] loadKeyMaterial(KeyStore keyStore, char[] password, + BiFunction aliasStrategy) + throws NoSuchAlgorithmException, UnrecoverableKeyException, KeyStoreException { + KeyManagerFactory kmFactory = KeyManagerFactory.getInstance(KeyManagerFactory.getDefaultAlgorithm()); + kmFactory.init(keyStore, password); + KeyManager[] keyManagers = kmFactory.getKeyManagers(); + + if (keyManagers != null && aliasStrategy != null) { + for (int i = 0; i < keyManagers.length; i++) { + KeyManager keyManager = keyManagers[i]; + if (keyManager instanceof X509ExtendedKeyManager) { + keyManagers[i] = new KeyManagerDelegate((X509ExtendedKeyManager) keyManager, aliasStrategy); + } + } + } + + return keyManagers; + } + + private static final class TrustManagerDelegate implements X509TrustManager { + private final X509TrustManager delegate; + private final BiPredicate trustStrategy; + + private TrustManagerDelegate(X509TrustManager delegate, BiPredicate trustStrategy) { + this.delegate = delegate; + this.trustStrategy = trustStrategy; + } + + @Override + public void checkClientTrusted(X509Certificate[] chain, String authType) throws CertificateException { + delegate.checkClientTrusted(chain, authType); + } + + @Override + public void checkServerTrusted(X509Certificate[] chain, String authType) throws CertificateException { + if (!trustStrategy.test(chain, authType)) { + delegate.checkServerTrusted(chain, authType); + } + } + + @Override + public X509Certificate[] getAcceptedIssuers() { + return delegate.getAcceptedIssuers(); + } + } + + private static final class KeyManagerDelegate extends X509ExtendedKeyManager { + private final X509ExtendedKeyManager delegate; + private final BiFunction aliasStrategy; + + private KeyManagerDelegate(X509ExtendedKeyManager delegate, + BiFunction aliasStrategy) { + this.delegate = delegate; + this.aliasStrategy = aliasStrategy; + } + + @Override + public String[] getClientAliases(String keyType, Principal[] issuers) { + return delegate.getClientAliases(keyType, issuers); + } + + @Override + public String chooseClientAlias(String[] keyType, Principal[] issuers, Socket socket) { + return aliasStrategy.apply(keyType, issuers); + } + + @Override + public String[] getServerAliases(String keyType, Principal[] issuers) { + return delegate.getServerAliases(keyType, issuers); + } + + @Override + public String chooseServerAlias(String keyType, Principal[] issuers, Socket socket) { + return aliasStrategy.apply(new String[] { keyType }, issuers); + } + + @Override + public X509Certificate[] getCertificateChain(String alias) { + return delegate.getCertificateChain(alias); + } + + @Override + public PrivateKey getPrivateKey(String alias) { + return delegate.getPrivateKey(alias); + } + + @Override + public String chooseEngineClientAlias(String[] keyType, Principal[] issuers, SSLEngine engine) { + return aliasStrategy.apply(keyType, issuers); + } + + @Override + public String chooseEngineServerAlias(String keyType, Principal[] issuers, SSLEngine engine) { + return aliasStrategy.apply(new String[] { keyType }, issuers); + } + } + + private JcaTestUtils() { + } +} diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/JreKeyStoreTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/JreKeyStoreTest.java index 07b74d695e98..76b056c7f93b 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/JreKeyStoreTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/JreKeyStoreTest.java @@ -4,29 +4,21 @@ package com.azure.security.keyvault.jca; import com.azure.security.keyvault.jca.implementation.certificates.JreCertificates; -import org.apache.hc.client5.http.classic.methods.HttpGet; -import org.apache.hc.client5.http.impl.classic.CloseableHttpClient; -import org.apache.hc.client5.http.impl.classic.HttpClients; -import org.apache.hc.client5.http.impl.io.PoolingHttpClientConnectionManager; -import org.apache.hc.client5.http.socket.ConnectionSocketFactory; -import org.apache.hc.client5.http.ssl.SSLConnectionSocketFactory; -import org.apache.hc.core5.http.ClassicHttpResponse; -import org.apache.hc.core5.http.config.RegistryBuilder; -import org.apache.hc.core5.http.io.HttpClientResponseHandler; -import org.apache.hc.core5.ssl.SSLContexts; import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import javax.net.ssl.HttpsURLConnection; import javax.net.ssl.SSLContext; import java.io.IOException; +import java.net.URI; import java.security.KeyStore; import java.security.cert.Certificate; import java.util.Map; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertNotNull; -import static org.junit.jupiter.api.Assertions.assertTrue; @EnabledIfEnvironmentVariable(named = "AZURE_KEYVAULT_CERTIFICATE_NAME", matches = "myalias") public class JreKeyStoreTest { @@ -48,7 +40,7 @@ public void testJreKsEntries() { assertNotNull(jreCertificates); assertNotNull(jreCertificates.getAliases()); Map certs = jreCertificates.getCertificates(); - assertTrue(certs.size() > 0); + assertFalse(certs.isEmpty()); assertNotNull(jreCertificates.getCertificateKeys()); } @@ -64,33 +56,27 @@ public void testJreKsTrustPeer() throws Exception { * - Create SSL connection factory. * - Set hostname verifier to trust any hostname. */ - - SSLContext sslContext = SSLContexts.custom().loadTrustMaterial(ks, null).build(); - - SSLConnectionSocketFactory sslConnectionSocketFactory - = new SSLConnectionSocketFactory(sslContext, (hostname, session) -> true); - - PoolingHttpClientConnectionManager manager = new PoolingHttpClientConnectionManager( - RegistryBuilder.create().register("https", sslConnectionSocketFactory).build()); + SSLContext sslContext = SSLContext.getInstance("TLS"); + sslContext.init(null, JcaTestUtils.loadTrustMaterial(ks, null), null); /* * And now execute the test. */ String result = null; - - try (CloseableHttpClient client = HttpClients.custom().setConnectionManager(manager).build()) { - HttpGet httpGet = new HttpGet("https://google.com:443"); - HttpClientResponseHandler responseHandler = (ClassicHttpResponse response) -> { - int status = response.getCode(); - String result1 = null; - if (status == 200) { - result1 = "Success"; - } - return result1; - }; - result = client.execute(httpGet, responseHandler); + HttpsURLConnection connection = null; + try { + connection = (HttpsURLConnection) URI.create("https://google.com:443").toURL().openConnection(); + connection.setSSLSocketFactory(sslContext.getSocketFactory()); + connection.setRequestMethod("GET"); + if (connection.getResponseCode() == 200) { + result = "Success"; + } } catch (IOException ioe) { ioe.printStackTrace(); + } finally { + if (connection != null) { + connection.disconnect(); + } } /* diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/KeyVaultEncodeTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/KeyVaultEncodeTest.java index d6e783d68224..2e77690995a8 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/KeyVaultEncodeTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/KeyVaultEncodeTest.java @@ -3,11 +3,11 @@ package com.azure.security.keyvault.jca; -import net.bytebuddy.utility.RandomString; import org.junit.jupiter.api.Assertions; import org.junit.jupiter.api.Test; import java.math.BigInteger; +import java.security.SecureRandom; import java.util.Arrays; import java.util.Base64; import java.util.Random; @@ -21,50 +21,51 @@ public void buildLengthBytesTest() { Random random = new Random(); int a = random.nextInt(1 << 7); byte[] result = KeyVaultEncode.buildLengthBytes(TEST_TAG, a); - Assertions.assertEquals(result.length, 2); - Assertions.assertEquals(result[0], TEST_TAG); - Assertions.assertEquals(result[1], (byte) a); + Assertions.assertEquals(2, result.length); + Assertions.assertEquals(TEST_TAG, result[0]); + Assertions.assertEquals((byte) a, result[1]); a = random.nextInt((1 << 8) - (1 << 7)) + (1 << 7); result = KeyVaultEncode.buildLengthBytes(TEST_TAG, a); - Assertions.assertEquals(result.length, 3); - Assertions.assertEquals(result[0], TEST_TAG); - Assertions.assertEquals(result[1], (byte) 0x081); - Assertions.assertEquals(result[2], (byte) a); + Assertions.assertEquals(3, result.length); + Assertions.assertEquals(TEST_TAG, result[0]); + Assertions.assertEquals((byte) 0x081, result[1]); + Assertions.assertEquals((byte) a, result[2]); a = random.nextInt((1 << 16) - (1 << 8)) + (1 << 8); result = KeyVaultEncode.buildLengthBytes(TEST_TAG, a); - Assertions.assertEquals(result.length, 4); - Assertions.assertEquals(result[0], TEST_TAG); - Assertions.assertEquals(result[1], (byte) 0x082); - Assertions.assertEquals(result[2], (byte) (a >> 8)); - Assertions.assertEquals(result[3], (byte) a); + Assertions.assertEquals(4, result.length); + Assertions.assertEquals(TEST_TAG, result[0]); + Assertions.assertEquals((byte) 0x082, result[1]); + Assertions.assertEquals((byte) (a >> 8), result[2]); + Assertions.assertEquals((byte) a, result[3]); a = random.nextInt((1 << 24) - (1 << 16)) + (1 << 16); result = KeyVaultEncode.buildLengthBytes(TEST_TAG, a); - Assertions.assertEquals(result.length, 5); - Assertions.assertEquals(result[0], TEST_TAG); - Assertions.assertEquals(result[1], (byte) 0x083); - Assertions.assertEquals(result[2], (byte) (a >> 16)); - Assertions.assertEquals(result[3], (byte) (a >> 8)); - Assertions.assertEquals(result[4], (byte) a); + Assertions.assertEquals(5, result.length); + Assertions.assertEquals(TEST_TAG, result[0]); + Assertions.assertEquals((byte) 0x083, result[1]); + Assertions.assertEquals((byte) (a >> 16), result[2]); + Assertions.assertEquals((byte) (a >> 8), result[3]); + Assertions.assertEquals((byte) a, result[4]); a = random.nextInt((1 << 30) - (1 << 24)) + (1 << 24); result = KeyVaultEncode.buildLengthBytes(TEST_TAG, a); - Assertions.assertEquals(result.length, 6); - Assertions.assertEquals(result[0], TEST_TAG); - Assertions.assertEquals(result[1], (byte) 0x084); - Assertions.assertEquals(result[2], (byte) (a >> 24)); - Assertions.assertEquals(result[3], (byte) (a >> 16)); - Assertions.assertEquals(result[4], (byte) (a >> 8)); - Assertions.assertEquals(result[5], (byte) a); + Assertions.assertEquals(6, result.length); + Assertions.assertEquals(TEST_TAG, result[0]); + Assertions.assertEquals((byte) 0x084, result[1]); + Assertions.assertEquals((byte) (a >> 24), result[2]); + Assertions.assertEquals((byte) (a >> 16), result[3]); + Assertions.assertEquals((byte) (a >> 8), result[4]); + Assertions.assertEquals((byte) a, result[5]); } @Test public void concatBytesWithThreeBytes() { - byte[] byte1 = RandomString.make(32).getBytes(); - byte[] byte2 = RandomString.make(32).getBytes(); - byte[] byte3 = RandomString.make(32).getBytes(); + SecureRandom random = new SecureRandom(); + byte[] byte1 = random.generateSeed(32); + byte[] byte2 = random.generateSeed(32); + byte[] byte3 = random.generateSeed(32); byte[] result = KeyVaultEncode.concatBytes(byte1, byte2, byte3); Assertions.assertArrayEquals(byte1, Arrays.copyOfRange(result, 0, byte1.length)); Assertions.assertArrayEquals(byte2, Arrays.copyOfRange(result, byte1.length, byte1.length + byte2.length)); @@ -73,8 +74,9 @@ public void concatBytesWithThreeBytes() { @Test public void concatBytesWithTwoBytes() { - byte[] byte1 = RandomString.make(32).getBytes(); - byte[] byte2 = RandomString.make(32).getBytes(); + SecureRandom random = new SecureRandom(); + byte[] byte1 = random.generateSeed(32); + byte[] byte2 = random.generateSeed(32); byte[] result = KeyVaultEncode.concatBytes(byte1, byte2); Assertions.assertArrayEquals(byte1, Arrays.copyOfRange(result, 0, byte1.length)); Assertions.assertArrayEquals(byte2, Arrays.copyOfRange(result, byte1.length, result.length)); @@ -82,8 +84,8 @@ public void concatBytesWithTwoBytes() { @Test public void toBigIntegerBytesWithLengthPrefixTest() { - byte[] testByte = RandomString.make(32).getBytes(); - Random random = new Random(); + SecureRandom random = new SecureRandom(); + byte[] testByte = random.generateSeed(32); int offset = random.nextInt(testByte.length); int length = random.nextInt(testByte.length - offset); byte[] result = KeyVaultEncode.toBigIntegerBytesWithLengthPrefix(testByte, offset, length); 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 98264d7c197b..36cb6cfc9af9 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 @@ -4,6 +4,7 @@ package com.azure.security.keyvault.jca; import com.azure.security.keyvault.jca.implementation.KeyVaultClient; +import com.azure.security.keyvault.jca.implementation.certificates.KeyVaultCertificates; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; @@ -11,21 +12,16 @@ import org.junit.jupiter.api.parallel.Resources; import java.io.ByteArrayInputStream; +import java.lang.reflect.Field; import java.security.ProviderException; import java.security.cert.CertificateException; import java.security.cert.CertificateFactory; import java.security.cert.X509Certificate; -import java.util.ArrayList; import java.util.Base64; -import java.util.List; - -import org.mockito.MockedConstruction; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertNotNull; -import static org.junit.jupiter.api.Assertions.assertSame; import static org.junit.jupiter.api.Assertions.assertTrue; -import static org.mockito.Mockito.mockConstruction; @ResourceLock(Resources.SYSTEM_PROPERTIES) public class KeyVaultKeyStoreUnitTest { @@ -112,22 +108,27 @@ public void testEngineLoadParameterOverridesSystemProperties() { } @Test - public void testEngineLoadPassesExplicitParameterToKeyVaultClient() { + public void testEngineLoadPassesExplicitParameterToKeyVaultClient() throws ReflectiveOperationException { System.setProperty(KeyVaultJcaPropertyNames.KEYVAULT_URI, "https://ambient.vault.azure.net"); KeyVaultLoadStoreParameter parameter = new KeyVaultLoadStoreParameter("https://explicit.vault.azure.net"); parameter.disableAiaDownload(); - List> constructorArguments = new ArrayList<>(); - try (MockedConstruction mockedConstruction = mockConstruction(KeyVaultClient.class, - (mock, context) -> constructorArguments.add(context.arguments()))) { - KeyVaultKeyStore keyStore = new KeyVaultKeyStore(); - keyStore.engineLoad(parameter); - assertEquals(2, mockedConstruction.constructed().size()); - } + KeyVaultKeyStore keyStore = new KeyVaultKeyStore(); + keyStore.engineLoad(parameter); - assertEquals(2, constructorArguments.size()); - assertSame(parameter, constructorArguments.get(1).get(0)); - assertTrue(((KeyVaultLoadStoreParameter) constructorArguments.get(1).get(0)).isAiaDownloadDisabled()); + Field certificatesField = KeyVaultKeyStore.class.getDeclaredField("keyVaultCertificates"); + certificatesField.setAccessible(true); + KeyVaultCertificates certificates = (KeyVaultCertificates) certificatesField.get(keyStore); + Field clientField = KeyVaultCertificates.class.getDeclaredField("keyVaultClient"); + clientField.setAccessible(true); + KeyVaultClient client = (KeyVaultClient) clientField.get(certificates); + Field uriField = KeyVaultClient.class.getDeclaredField("keyVaultUri"); + uriField.setAccessible(true); + Field disableAiaField = KeyVaultClient.class.getDeclaredField("disableAiaDownload"); + disableAiaField.setAccessible(true); + + assertEquals("https://explicit.vault.azure.net/", uriField.get(client)); + assertTrue((Boolean) disableAiaField.get(client)); } @BeforeEach diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/PropertyConvertorUtils.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/PropertyConvertorUtils.java index b3ca277b5c54..ea0ca7889286 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/PropertyConvertorUtils.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/PropertyConvertorUtils.java @@ -3,7 +3,6 @@ package com.azure.security.keyvault.jca; -import com.azure.core.util.Configuration; import java.io.IOException; import java.security.KeyStore; import java.security.KeyStoreException; @@ -14,9 +13,6 @@ import java.util.List; public class PropertyConvertorUtils { - - private static final Configuration GLOBAL_CONFIGURATION = Configuration.getGlobalConfiguration(); - public static void putEnvironmentPropertyToSystemPropertyForKeyVaultJca() { KEYVAULT_JCA_SYSTEM_PROPERTIES.forEach(environmentPropertyKey -> { String value = getPropertyValue(environmentPropertyKey); @@ -41,11 +37,12 @@ public static void addKeyVaultJcaProvider() { } public static String getPropertyValue(String property) { - return GLOBAL_CONFIGURATION.get(property, System.getenv(property)); - } + String value = System.getProperty(property); + if (value != null) { + return value; + } - public static String getPropertyValue(String property, String defaultValue) { - return GLOBAL_CONFIGURATION.get(property, defaultValue); + return System.getenv(property); } } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/ServerSocketTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/ServerSocketTest.java index bd2d91b7f6ad..6e2c63efcca7 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/ServerSocketTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/ServerSocketTest.java @@ -3,36 +3,22 @@ package com.azure.security.keyvault.jca; -import org.apache.hc.client5.http.classic.methods.HttpGet; -import org.apache.hc.client5.http.impl.classic.CloseableHttpClient; -import org.apache.hc.client5.http.impl.classic.HttpClients; -import org.apache.hc.client5.http.impl.io.PoolingHttpClientConnectionManager; -import org.apache.hc.client5.http.socket.ConnectionSocketFactory; -import org.apache.hc.client5.http.ssl.SSLConnectionSocketFactory; -import org.apache.hc.client5.http.ssl.TrustSelfSignedStrategy; -import org.apache.hc.core5.http.ClassicHttpResponse; -import org.apache.hc.core5.http.config.RegistryBuilder; -import org.apache.hc.core5.http.io.HttpClientResponseHandler; -import org.apache.hc.core5.ssl.PrivateKeyDetails; -import org.apache.hc.core5.ssl.PrivateKeyStrategy; -import org.apache.hc.core5.ssl.SSLContexts; import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import javax.net.ssl.HttpsURLConnection; import javax.net.ssl.KeyManagerFactory; import javax.net.ssl.SSLContext; -import javax.net.ssl.SSLParameters; import javax.net.ssl.SSLServerSocket; import javax.net.ssl.SSLServerSocketFactory; import javax.net.ssl.TrustManagerFactory; import java.io.IOException; import java.io.OutputStream; import java.net.Socket; +import java.net.URI; import java.security.KeyStore; import java.security.Security; -import java.security.cert.X509Certificate; -import java.util.Map; import static org.junit.jupiter.api.Assertions.assertEquals; @@ -57,7 +43,7 @@ public static void beforeEach() throws Exception { KeyVaultJcaProvider provider = new KeyVaultJcaProvider(); Security.addProvider(provider); - /** + /* * - Create an Azure Key Vault specific instance of a KeyStore. * - Set the KeyManagerFactory to use that KeyStore. */ @@ -86,16 +72,16 @@ private void startSocket(SSLServerSocket serverSocket) { @Test public void testHttpsConnectionWithoutClientTrust() throws Exception { - SSLContext sslContext = SSLContexts.custom() - .loadTrustMaterial((final X509Certificate[] chain, final String authType) -> true) - .build(); + SSLContext sslContext = SSLContext.getInstance("TLS"); + sslContext.init(null, JcaTestUtils.loadTrustMaterial(null, (ignoredChain, ignoredAuthType) -> true), null); testHttpsConnection(8765, sslContext); } @Test public void testHttpsConnectionWithSelfSignedClientTrust() throws Exception { - SSLContext sslContext = SSLContexts.custom().loadTrustMaterial(ks, new TrustSelfSignedStrategy()).build(); + SSLContext sslContext = SSLContext.getInstance("TLS"); + sslContext.init(null, JcaTestUtils.loadTrustMaterial(ks, (chain, ignored) -> chain.length == 1), null); testHttpsConnection(8766, sslContext); } @@ -144,7 +130,7 @@ private void testHttpsConnection(Integer port, SSLContext sslContext) throws Exc assertEquals("Success", result); } - private void serverSocketWithTrustManager(Integer port) throws Exception { + private void serverSocketWithTrustManager(int port) throws Exception { /* * Setup server side. * @@ -169,11 +155,10 @@ private void serverSocketWithTrustManager(Integer port) throws Exception { * - Create an SSL context. * - Set SSL context to trust any certificate. */ - - SSLContext sslContext = SSLContexts.custom() - .loadTrustMaterial(ks, new TrustSelfSignedStrategy()) - .loadKeyMaterial(ks, "".toCharArray(), new ClientPrivateKeyStrategy()) - .build(); + SSLContext sslContext = SSLContext.getInstance("TLS"); + sslContext.init( + JcaTestUtils.loadKeyMaterial(ks, "".toCharArray(), (ignoredKeyTypes, ignoredIssuers) -> certificateName), + JcaTestUtils.loadTrustMaterial(ks, (chain, ignored) -> chain.length == 1), null); /* * And now execute the test. @@ -186,41 +171,29 @@ private void serverSocketWithTrustManager(Integer port) throws Exception { assertEquals("Success", result); } - private String sendRequest(SSLContext sslContext, Integer port) { + private String sendRequest(SSLContext sslContext, int port) { - /** + /* * - Create SSL connection factory. * - Set hostname verifier to trust any hostname. */ - SSLConnectionSocketFactory sslConnectionSocketFactory - = new SSLConnectionSocketFactory(sslContext, (hostname, session) -> true); - - PoolingHttpClientConnectionManager manager = new PoolingHttpClientConnectionManager( - RegistryBuilder.create().register("https", sslConnectionSocketFactory).build()); - String result = null; - - try (CloseableHttpClient client = HttpClients.custom().setConnectionManager(manager).build()) { - HttpGet httpGet = new HttpGet("https://localhost:" + port); - HttpClientResponseHandler responseHandler = (ClassicHttpResponse response) -> { - int status = response.getCode(); - String result1 = null; - if (status == 204) { - result1 = "Success"; - } - return result1; - }; - result = client.execute(httpGet, responseHandler); + HttpsURLConnection connection = null; + try { + connection = (HttpsURLConnection) URI.create("https://localhost:" + port).toURL().openConnection(); + connection.setSSLSocketFactory(sslContext.getSocketFactory()); + connection.setHostnameVerifier((hostname, session) -> true); + connection.setRequestMethod("GET"); + if (connection.getResponseCode() == 204) { + result = "Success"; + } } catch (IOException ioe) { ioe.printStackTrace(); + } finally { + if (connection != null) { + connection.disconnect(); + } } return result; } - - private static class ClientPrivateKeyStrategy implements PrivateKeyStrategy { - @Override - public String chooseAlias(Map aliases, SSLParameters sslParameters) { - return certificateName; // It should be your certificate alias used in client-side - } - } } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/SpecificPathCertificatesTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/SpecificPathCertificatesTest.java index 9519a3ce9c6e..694699d6048c 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/SpecificPathCertificatesTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/SpecificPathCertificatesTest.java @@ -13,6 +13,7 @@ import java.io.FileInputStream; import java.io.IOException; import java.io.InputStream; +import java.nio.file.FileSystems; import java.security.KeyStore; import java.security.KeyStoreException; import java.security.NoSuchAlgorithmException; @@ -35,7 +36,7 @@ public static void setEnvironmentProperty() { public static String getFilePath(String packageName) { String filepath = "\\src\\test\\resources\\" + packageName; - return System.getProperty("user.dir") + filepath.replace("\\", System.getProperty("file.separator")); + return System.getProperty("user.dir") + filepath.replace("\\", FileSystems.getDefault().getSeparator()); } @Test diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/JreKeyStoreFactoryTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/JreKeyStoreFactoryTest.java index d059a06792f9..21616730cc9e 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/JreKeyStoreFactoryTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/JreKeyStoreFactoryTest.java @@ -7,12 +7,12 @@ import java.security.KeyStore; -import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotEquals; public class JreKeyStoreFactoryTest { @Test public void test() { KeyStore jreKeyStore = JreKeyStoreFactory.getDefaultKeyStore(); - assertFalse(jreKeyStore.getType().equals(KeyVaultKeyStore.KEY_STORE_TYPE)); + assertNotEquals(KeyVaultKeyStore.KEY_STORE_TYPE, jreKeyStore.getType()); } } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/JreKeyStoreTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/JreKeyStoreTest.java index cfa223e5073d..f8152bd1a1b0 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/JreKeyStoreTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/JreKeyStoreTest.java @@ -5,9 +5,12 @@ import com.azure.security.keyvault.jca.implementation.certificates.JreCertificates; import org.junit.jupiter.api.Test; + import java.security.cert.Certificate; import java.util.Map; -import static org.junit.jupiter.api.Assertions.*; + +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; public class JreKeyStoreTest { @@ -17,7 +20,7 @@ public void testJreKsEntries() { assertNotNull(jreCertificates); assertNotNull(jreCertificates.getAliases()); Map certs = jreCertificates.getCertificates(); - assertTrue(certs.size() > 0); + assertFalse(certs.isEmpty()); assertNotNull(jreCertificates.getCertificateKeys()); } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/KeyVaultClientTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/KeyVaultClientTest.java index 503488e57f8b..7b302732d0df 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/KeyVaultClientTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/KeyVaultClientTest.java @@ -12,26 +12,27 @@ import com.azure.security.keyvault.jca.implementation.model.CertificatePolicy; import com.azure.security.keyvault.jca.implementation.model.KeyProperties; import com.azure.security.keyvault.jca.implementation.model.SecretBundle; -import com.azure.security.keyvault.jca.implementation.utils.AccessTokenUtil; -import com.azure.security.keyvault.jca.implementation.utils.CertificateUtil; import com.azure.security.keyvault.jca.implementation.utils.HttpUtil; import com.azure.security.keyvault.jca.implementation.utils.JsonConverterUtil; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; -import org.mockito.MockedStatic; -import org.mockito.Mockito; +import java.io.ByteArrayOutputStream; import java.io.IOException; import java.nio.charset.StandardCharsets; import java.nio.file.Files; import java.nio.file.Paths; import java.security.Key; +import java.security.KeyStore; import java.security.cert.Certificate; import java.security.cert.CertificateException; import java.util.ArrayList; import java.util.Arrays; import java.util.Base64; +import java.util.HashMap; import java.util.List; +import java.util.Map; +import java.util.concurrent.atomic.AtomicInteger; import java.util.logging.Handler; import java.util.logging.Level; import java.util.logging.LogRecord; @@ -40,20 +41,17 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; import static org.junit.jupiter.api.Assertions.assertSame; import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; import static com.azure.security.keyvault.jca.implementation.utils.HttpUtil.API_VERSION_POSTFIX; -import static org.mockito.ArgumentMatchers.anyMap; -import static org.mockito.ArgumentMatchers.anyString; -import static org.mockito.ArgumentMatchers.eq; -import static org.mockito.ArgumentMatchers.notNull; -import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.times; public class KeyVaultClientTest { private static final String KEY_VAULT_TEST_URI_GLOBAL = "https://fake.vault.azure.net/"; + private static final String TEST_ACCESS_TOKEN = "test-token"; + private static final String CERTIFICATE_ALIAS = "client-cert"; private static final String CERTIFICATE_URI @@ -66,301 +64,309 @@ public class KeyVaultClientTest { @Test public void testGetAliasWithCertificateInfoWith0Page() { - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - utilities.when(() -> HttpUtil.get(anyString(), anyMap())).thenReturn("fakeValue"); - - KeyVaultClient keyVaultClient = mock(KeyVaultClient.class); - List result = keyVaultClient.getAliases(); + KeyVaultClient keyVaultClient = new TestKeyVaultClient(TEST_ACCESS_TOKEN, false, (uri, headers) -> "fakeValue"); - assertEquals(0, result.size()); - } + assertEquals(0, keyVaultClient.getAliases().size()); } @Test public void testGetAliasWithCertificateInfoWith1Page() { - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - utilities.when(() -> HttpUtil.validateUri(anyString(), anyString())).thenCallRealMethod(); - utilities.when(() -> HttpUtil.addTrailingSlashIfRequired(anyString())).thenCallRealMethod(); - - // Create fake certificates. - CertificateItem fakeCertificateItem1 = new CertificateItem(); - fakeCertificateItem1.setId("certificates/fakeCertificateItem1"); + // Create fake certificates. + CertificateItem fakeCertificateItem1 = new CertificateItem(); + fakeCertificateItem1.setId("certificates/fakeCertificateItem1"); - CertificateListResult certificateListResult = new CertificateListResult(); - certificateListResult.setValue(Arrays.asList(fakeCertificateItem1)); + CertificateListResult certificateListResult = new CertificateListResult(); + certificateListResult.setValue(Arrays.asList(fakeCertificateItem1)); - String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); + String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); - utilities.when(() -> HttpUtil.get(notNull(), anyMap())).thenReturn(certificateListResultString); + KeyVaultClient keyVaultClient + = new TestKeyVaultClient(TEST_ACCESS_TOKEN, false, (uri, headers) -> certificateListResultString); - KeyVaultClient keyVaultClient = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null); - List result = keyVaultClient.getAliases(); - - assertEquals(1, result.size()); - assertTrue(result.contains("fakeCertificateItem1")); - } + List result = keyVaultClient.getAliases(); + assertEquals(1, result.size()); + assertTrue(result.contains("fakeCertificateItem1")); } @Test public void testGetAliasWithCertificateInfoWith2Pages() { - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - utilities.when(() -> HttpUtil.validateUri(anyString(), anyString())).thenCallRealMethod(); - utilities.when(() -> HttpUtil.addTrailingSlashIfRequired(anyString())).thenCallRealMethod(); - - // create fake certificates - CertificateItem fakeCertificateItem1 = new CertificateItem(); - fakeCertificateItem1.setId("certificates/fakeCertificateItem1"); + // create fake certificates + CertificateItem fakeCertificateItem1 = new CertificateItem(); + fakeCertificateItem1.setId("certificates/fakeCertificateItem1"); - CertificateItem fakeCertificateItem2 = new CertificateItem(); - fakeCertificateItem2.setId("certificates/fakeCertificateItem2"); + CertificateItem fakeCertificateItem2 = new CertificateItem(); + fakeCertificateItem2.setId("certificates/fakeCertificateItem2"); - CertificateItem fakeCertificateItem3 = new CertificateItem(); - fakeCertificateItem3.setId("certificates/fakeCertificateItem3"); + CertificateItem fakeCertificateItem3 = new CertificateItem(); + fakeCertificateItem3.setId("certificates/fakeCertificateItem3"); - // Create first page certificate result. - CertificateListResult certificateListResult = new CertificateListResult(); - certificateListResult.setNextLink("fakeNextLint"); - certificateListResult.setValue(Arrays.asList(fakeCertificateItem1)); + // Create first page certificate result. + CertificateListResult certificateListResult = new CertificateListResult(); + certificateListResult.setNextLink("fakeNextLink"); + certificateListResult.setValue(Arrays.asList(fakeCertificateItem1)); - // Create next page certificate result. - CertificateListResult certificateListResultNext = new CertificateListResult(); - certificateListResultNext.setValue(Arrays.asList(fakeCertificateItem2, fakeCertificateItem3)); + // Create next page certificate result. + CertificateListResult certificateListResultNext = new CertificateListResult(); + certificateListResultNext.setValue(Arrays.asList(fakeCertificateItem2, fakeCertificateItem3)); - String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); - String certificateListResultStringNext = JsonConverterUtil.toJson(certificateListResultNext); + String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); + String certificateListResultStringNext = JsonConverterUtil.toJson(certificateListResultNext); - utilities.when(() -> HttpUtil.get(notNull(), anyMap())).thenReturn(certificateListResultString); - utilities.when(() -> HttpUtil.get(eq("fakeNextLint"), anyMap())) - .thenReturn(certificateListResultStringNext); + KeyVaultClient keyVaultClient = new TestKeyVaultClient(TEST_ACCESS_TOKEN, false, (uri, + headers) -> "fakeNextLink".equals(uri) ? certificateListResultStringNext : certificateListResultString); - KeyVaultClient keyVaultClient = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null); - List result = keyVaultClient.getAliases(); - - assertEquals(3, result.size()); - assertTrue(result - .containsAll(Arrays.asList("fakeCertificateItem1", "fakeCertificateItem2", "fakeCertificateItem3"))); - } + List result = keyVaultClient.getAliases(); + assertEquals(3, result.size()); + assertTrue( + result.containsAll(Arrays.asList("fakeCertificateItem1", "fakeCertificateItem2", "fakeCertificateItem3"))); } @Test public void testGetAliasFiltersOutDisabledCertificate() { - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - utilities.when(() -> HttpUtil.validateUri(anyString(), anyString())).thenCallRealMethod(); - utilities.when(() -> HttpUtil.addTrailingSlashIfRequired(anyString())).thenCallRealMethod(); - - // Enabled certificate. - CertificateItemAttributes enabledAttributes = new CertificateItemAttributes(); - enabledAttributes.setEnabled(true); - CertificateItem enabledCertificate = new CertificateItem(); - enabledCertificate.setId("certificates/client-cert-active"); - enabledCertificate.setAttributes(enabledAttributes); - - // Disabled certificate. This one previously caused an HTTP 403 while initializing the keystore. - CertificateItemAttributes disabledAttributes = new CertificateItemAttributes(); - disabledAttributes.setEnabled(false); - CertificateItem disabledCertificate = new CertificateItem(); - disabledCertificate.setId("certificates/client-cert-unused"); - disabledCertificate.setAttributes(disabledAttributes); - - CertificateListResult certificateListResult = new CertificateListResult(); - certificateListResult.setValue(Arrays.asList(enabledCertificate, disabledCertificate)); - - String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); - utilities.when(() -> HttpUtil.get(notNull(), anyMap())).thenReturn(certificateListResultString); - - KeyVaultClient keyVaultClient = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null); - List result = keyVaultClient.getAliases(); - - assertEquals(1, result.size()); - assertTrue(result.contains("client-cert-active")); - assertFalse(result.contains("client-cert-unused")); - } + // Enabled certificate. + CertificateItemAttributes enabledAttributes = new CertificateItemAttributes(); + enabledAttributes.setEnabled(true); + CertificateItem enabledCertificate = new CertificateItem(); + enabledCertificate.setId("certificates/client-cert-active"); + enabledCertificate.setAttributes(enabledAttributes); + + // Disabled certificate. This one previously caused an HTTP 403 while initializing the keystore. + CertificateItemAttributes disabledAttributes = new CertificateItemAttributes(); + disabledAttributes.setEnabled(false); + CertificateItem disabledCertificate = new CertificateItem(); + disabledCertificate.setId("certificates/client-cert-unused"); + disabledCertificate.setAttributes(disabledAttributes); + + CertificateListResult certificateListResult = new CertificateListResult(); + certificateListResult.setValue(Arrays.asList(enabledCertificate, disabledCertificate)); + + String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); + + KeyVaultClient keyVaultClient + = new TestKeyVaultClient(TEST_ACCESS_TOKEN, false, (uri, headers) -> certificateListResultString); + List result = keyVaultClient.getAliases(); + + assertEquals(1, result.size()); + assertTrue(result.contains("client-cert-active")); + assertFalse(result.contains("client-cert-unused")); } @Test public void testGetAliasKeepsEnabledAndAttributelessCertificates() { - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - utilities.when(() -> HttpUtil.validateUri(anyString(), anyString())).thenCallRealMethod(); - utilities.when(() -> HttpUtil.addTrailingSlashIfRequired(anyString())).thenCallRealMethod(); + // Certificate explicitly enabled. + CertificateItemAttributes enabledAttributes = new CertificateItemAttributes(); + enabledAttributes.setEnabled(true); + CertificateItem enabledCertificate = new CertificateItem(); + enabledCertificate.setId("certificates/enabledCertificate"); + enabledCertificate.setAttributes(enabledAttributes); - // Certificate explicitly enabled. - CertificateItemAttributes enabledAttributes = new CertificateItemAttributes(); - enabledAttributes.setEnabled(true); - CertificateItem enabledCertificate = new CertificateItem(); - enabledCertificate.setId("certificates/enabledCertificate"); - enabledCertificate.setAttributes(enabledAttributes); + // Certificate without attributes, which must be treated as enabled for backward compatibility. + CertificateItem attributelessCertificate = new CertificateItem(); + attributelessCertificate.setId("certificates/attributelessCertificate"); - // Certificate without attributes, which must be treated as enabled for backward compatibility. - CertificateItem attributelessCertificate = new CertificateItem(); - attributelessCertificate.setId("certificates/attributelessCertificate"); + CertificateListResult certificateListResult = new CertificateListResult(); + certificateListResult.setValue(Arrays.asList(enabledCertificate, attributelessCertificate)); - CertificateListResult certificateListResult = new CertificateListResult(); - certificateListResult.setValue(Arrays.asList(enabledCertificate, attributelessCertificate)); + String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); - String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); - utilities.when(() -> HttpUtil.get(notNull(), anyMap())).thenReturn(certificateListResultString); + KeyVaultClient keyVaultClient + = new TestKeyVaultClient(TEST_ACCESS_TOKEN, false, (uri, headers) -> certificateListResultString); + List result = keyVaultClient.getAliases(); - KeyVaultClient keyVaultClient = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null); - List result = keyVaultClient.getAliases(); - - assertEquals(2, result.size()); - assertTrue(result.containsAll(Arrays.asList("enabledCertificate", "attributelessCertificate"))); - } + assertEquals(2, result.size()); + assertTrue(result.containsAll(Arrays.asList("enabledCertificate", "attributelessCertificate"))); } @Test public void testGetAliasFiltersDisabledCertificateFromRawResponse() { - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - utilities.when(() -> HttpUtil.validateUri(anyString(), anyString())).thenCallRealMethod(); - utilities.when(() -> HttpUtil.addTrailingSlashIfRequired(anyString())).thenCallRealMethod(); - - // A response that mirrors the shape returned by the Azure Key Vault "list certificates" REST API, with one - // enabled and one disabled certificate. - String rawResponse = "{\"value\":[" - + "{\"id\":\"https://fake.vault.azure.net/certificates/client-cert-active\"," + // A response that mirrors the shape returned by the Azure Key Vault "list certificates" REST API, with one + // enabled and one disabled certificate. + String rawResponse + = "{\"value\":[" + "{\"id\":\"https://fake.vault.azure.net/certificates/client-cert-active\"," + "\"attributes\":{\"enabled\":true,\"nbf\":1783324860,\"exp\":1814861460}}," + "{\"id\":\"https://fake.vault.azure.net/certificates/client-cert-unused\"," + "\"attributes\":{\"enabled\":false,\"nbf\":1783324860,\"exp\":1814861460}}]," + "\"nextLink\":null}"; - utilities.when(() -> HttpUtil.get(notNull(), anyMap())).thenReturn(rawResponse); + KeyVaultClient keyVaultClient = new TestKeyVaultClient(TEST_ACCESS_TOKEN, false, (uri, headers) -> rawResponse); + List result = keyVaultClient.getAliases(); - KeyVaultClient keyVaultClient = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null); - List result = keyVaultClient.getAliases(); - - assertEquals(1, result.size()); - assertTrue(result.contains("client-cert-active")); - assertFalse(result.contains("client-cert-unused")); - } + assertEquals(1, result.size()); + assertTrue(result.contains("client-cert-active")); + assertFalse(result.contains("client-cert-unused")); } @Test public void testCacheToken() { - try (MockedStatic tokenUtilMockedStatic = Mockito.mockStatic(AccessTokenUtil.class); - MockedStatic httpUtilMockedStatic = Mockito.mockStatic(HttpUtil.class)) { - - httpUtilMockedStatic.when(() -> HttpUtil.validateUri(anyString(), anyString())).thenCallRealMethod(); - httpUtilMockedStatic.when(() -> HttpUtil.addTrailingSlashIfRequired(anyString())).thenCallRealMethod(); + AccessToken cacheToken = new AccessToken(); + cacheToken.setExpiresIn(300); // 300 seconds. - AccessToken cacheToken = new AccessToken(); - cacheToken.setExpiresIn(300); // 300 seconds. + CertificateItem fakeCertificateItem = new CertificateItem(); + fakeCertificateItem.setId("certificates/fakeCertificateItem"); - tokenUtilMockedStatic.when(() -> AccessTokenUtil.getAccessToken(anyString(), anyString())) - .thenReturn(cacheToken); + CertificateListResult certificateListResult = new CertificateListResult(); + certificateListResult.setValue(Arrays.asList(fakeCertificateItem)); - CertificateItem fakeCertificateItem = new CertificateItem(); - fakeCertificateItem.setId("certificates/fakeCertificateItem"); + String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); - CertificateListResult certificateListResult = new CertificateListResult(); - certificateListResult.setValue(Arrays.asList(fakeCertificateItem)); - - String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); - httpUtilMockedStatic.when(() -> HttpUtil.get(anyString(), anyMap())) - .thenReturn(certificateListResultString); + AtomicInteger getAccessTokenCount = new AtomicInteger(); + KeyVaultClient keyVaultClient = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, "") { + @Override + String httpGet(String uri, Map headers) { + return certificateListResultString; + } - KeyVaultClient keyVaultClient = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, ""); - keyVaultClient.getAliases(); - keyVaultClient.getAliases(); // Get aliases the second time. + @Override + AccessToken getAccessToken(String resource, String identity) { + getAccessTokenCount.incrementAndGet(); + return cacheToken; + } + }; + keyVaultClient.getAliases(); + keyVaultClient.getAliases(); // Get aliases the second time. - tokenUtilMockedStatic.verify(() -> AccessTokenUtil.getAccessToken(anyString(), anyString()), times(1)); - } + assertEquals(1, getAccessTokenCount.get()); } @Test public void testCacheTokenExpired() { - try (MockedStatic tokenUtilMockedStatic = Mockito.mockStatic(AccessTokenUtil.class); - MockedStatic httpUtilMockedStatic = Mockito.mockStatic(HttpUtil.class)) { - - httpUtilMockedStatic.when(() -> HttpUtil.validateUri(anyString(), anyString())).thenCallRealMethod(); - httpUtilMockedStatic.when(() -> HttpUtil.addTrailingSlashIfRequired(anyString())).thenCallRealMethod(); + AccessToken cacheToken = new AccessToken(); + cacheToken.setExpiresIn(50); // 50 seconds. - AccessToken cacheToken = new AccessToken(); - cacheToken.setExpiresIn(50); // 50 seconds. + CertificateItem fakeCertificateItem = new CertificateItem(); + fakeCertificateItem.setId("certificates/fakeCertificateItem"); - tokenUtilMockedStatic.when(() -> AccessTokenUtil.getAccessToken(anyString(), anyString())) - .thenReturn(cacheToken); + CertificateListResult certificateListResult = new CertificateListResult(); + certificateListResult.setValue(Arrays.asList(fakeCertificateItem)); - CertificateItem fakeCertificateItem = new CertificateItem(); - fakeCertificateItem.setId("certificates/fakeCertificateItem"); + String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); - CertificateListResult certificateListResult = new CertificateListResult(); - certificateListResult.setValue(Arrays.asList(fakeCertificateItem)); + AtomicInteger getAccessTokenCount = new AtomicInteger(); + KeyVaultClient keyVaultClient = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, "") { + @Override + String httpGet(String uri, Map headers) { + return certificateListResultString; + } - String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); - httpUtilMockedStatic.when(() -> HttpUtil.get(anyString(), anyMap())) - .thenReturn(certificateListResultString); + @Override + AccessToken getAccessToken(String resource, String identity) { + getAccessTokenCount.incrementAndGet(); + return cacheToken; + } + }; - KeyVaultClient keyVaultClient = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, ""); - keyVaultClient.getAliases(); - keyVaultClient.getAliases(); // Get aliases the second time. + keyVaultClient.getAliases(); + keyVaultClient.getAliases(); // Get aliases the second time. - tokenUtilMockedStatic.verify(() -> AccessTokenUtil.getAccessToken(anyString(), anyString()), times(2)); - } + assertEquals(2, getAccessTokenCount.get()); } @Test public void testAccessTokenAuthentication() { - try (MockedStatic httpUtilMockedStatic = Mockito.mockStatic(HttpUtil.class)) { - httpUtilMockedStatic.when(() -> HttpUtil.validateUri(anyString(), anyString())).thenCallRealMethod(); - httpUtilMockedStatic.when(() -> HttpUtil.addTrailingSlashIfRequired(anyString())).thenCallRealMethod(); - - CertificateItem fakeCertificateItem = new CertificateItem(); - fakeCertificateItem.setId("certificates/fakeCertificateItem"); + CertificateItem fakeCertificateItem = new CertificateItem(); + fakeCertificateItem.setId("certificates/fakeCertificateItem"); - CertificateListResult certificateListResult = new CertificateListResult(); - certificateListResult.setValue(Arrays.asList(fakeCertificateItem)); + CertificateListResult certificateListResult = new CertificateListResult(); + certificateListResult.setValue(Arrays.asList(fakeCertificateItem)); - String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); - httpUtilMockedStatic.when(() -> HttpUtil.get(anyString(), anyMap())) - .thenReturn(certificateListResultString); + String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); - // Create client with access token - String testAccessToken = "test-bearer-token-12345"; - KeyVaultClient keyVaultClient - = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null, null, null, null, testAccessToken, false); + // Create client with access token + String testAccessToken = "test-bearer-token-12345"; + KeyVaultClient keyVaultClient + = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null, null, null, null, testAccessToken, false) { + @Override + String httpGet(String uri, Map headers) { + return certificateListResultString; + } + }; - List result = keyVaultClient.getAliases(); + List result = keyVaultClient.getAliases(); - // Verify that the access token was used - assertEquals(1, result.size()); - assertTrue(result.contains("fakeCertificateItem")); - } + // Verify that the access token was used + assertEquals(1, result.size()); + assertTrue(result.contains("fakeCertificateItem")); } @Test public void testAuthenticationPriority() { - try (MockedStatic httpUtilMockedStatic = Mockito.mockStatic(HttpUtil.class); - MockedStatic tokenUtilMockedStatic = Mockito.mockStatic(AccessTokenUtil.class)) { - - httpUtilMockedStatic.when(() -> HttpUtil.validateUri(anyString(), anyString())).thenCallRealMethod(); - httpUtilMockedStatic.when(() -> HttpUtil.addTrailingSlashIfRequired(anyString())).thenCallRealMethod(); - - AccessToken accessToken = new AccessToken("fake-token", 3600); - tokenUtilMockedStatic.when(() -> AccessTokenUtil.getAccessToken(anyString(), anyString())) - .thenReturn(accessToken); - - CertificateItem fakeCertificateItem = new CertificateItem(); - fakeCertificateItem.setId("certificates/fakeCertificateItem"); - - CertificateListResult certificateListResult = new CertificateListResult(); - certificateListResult.setValue(Arrays.asList(fakeCertificateItem)); - - String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); - httpUtilMockedStatic.when(() -> HttpUtil.get(anyString(), anyMap())) - .thenReturn(certificateListResultString); - - // Test 1: Managed Identity should take priority over access token - KeyVaultClient client1 - = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null, null, null, "managed-id", "bearer-token", false); - client1.getAliases(); - tokenUtilMockedStatic.verify(() -> AccessTokenUtil.getAccessToken(anyString(), eq("managed-id")), times(1)); - - // Test 2: Access token should be used when managed identity is not set - KeyVaultClient client2 - = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null, null, null, null, "bearer-token", false); - List result = client2.getAliases(); - assertEquals(1, result.size()); - assertTrue(result.contains("fakeCertificateItem")); - } + AtomicInteger getAccessTokenCount = new AtomicInteger(); + AccessToken accessToken = new AccessToken("fake-token", 3600); + + CertificateItem fakeCertificateItem = new CertificateItem(); + fakeCertificateItem.setId("certificates/fakeCertificateItem"); + + CertificateListResult certificateListResult = new CertificateListResult(); + certificateListResult.setValue(Arrays.asList(fakeCertificateItem)); + + String certificateListResultString = JsonConverterUtil.toJson(certificateListResult); + + // Test 1: Managed Identity should take priority over access token + KeyVaultClient client1 + = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null, null, null, "managed-id", "bearer-token", false) { + @Override + String httpGet(String uri, Map headers) { + return certificateListResultString; + } + + @Override + AccessToken getAccessToken(String resource, String identity) { + if ("managed-id".equals(identity)) { + getAccessTokenCount.incrementAndGet(); + } + return accessToken; + } + }; + client1.getAliases(); + + // Test 2: Access token should be used when managed identity is not set + KeyVaultClient client2 + = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null, null, null, null, "bearer-token", false) { + @Override + String httpGet(String uri, Map headers) { + return certificateListResultString; + } + + @Override + AccessToken getAccessToken(String resource, String identity) { + if ("managed-id".equals(identity)) { + getAccessTokenCount.incrementAndGet(); + } + return accessToken; + } + }; + + List result = client2.getAliases(); + assertEquals(1, result.size()); + assertTrue(result.contains("fakeCertificateItem")); + + assertEquals(1, getAccessTokenCount.get()); + } + + @Test + public void testSystemAssignedManagedIdentityFallback() { + CertificateItem certificateItem = new CertificateItem(); + certificateItem.setId("certificates/fakeCertificateItem"); + CertificateListResult certificateListResult = new CertificateListResult(); + certificateListResult.setValue(Arrays.asList(certificateItem)); + String response = JsonConverterUtil.toJson(certificateListResult); + AtomicInteger getAccessTokenCount = new AtomicInteger(); + + KeyVaultClient client = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null) { + @Override + String httpGet(String uri, Map headers) { + return response; + } + + @Override + AccessToken getAccessToken(String resource, String identity) { + assertNull(identity); + getAccessTokenCount.incrementAndGet(); + return new AccessToken("fake-token", 3600); + } + }; + + assertEquals(1, client.getAliases().size()); + assertEquals(1, getAccessTokenCount.get()); } @Test @@ -372,102 +378,79 @@ public void testCertificateChainUsesVersionedSecretId() throws Exception { Paths.get("src/test/resources/certificate-util/SecretBundle.value/3-certificates-in-chain.pem")), StandardCharsets.UTF_8)); - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(CERTIFICATE_URI), anyMap())) - .thenReturn(JsonConverterUtil.toJson(certificateBundle)); - utilities.when(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap())) - .thenReturn(JsonConverterUtil.toJson(secretBundle)); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(CERTIFICATE_URI, JsonConverterUtil.toJson(certificateBundle)); + keyVaultClient.addHttpResponses(VERSIONED_SECRET_ID + API_VERSION_POSTFIX, + JsonConverterUtil.toJson(secretBundle)); - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - CertificateVersion certificateVersion = keyVaultClient.resolveCertificateVersion(CERTIFICATE_ALIAS); - Certificate[] chain = keyVaultClient.getCertificateChainForVersion(certificateVersion); + CertificateVersion certificateVersion = keyVaultClient.resolveCertificateVersion(CERTIFICATE_ALIAS); + Certificate[] chain = keyVaultClient.getCertificateChainForVersion(certificateVersion); - assertEquals(3, chain.length); - utilities.verify(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap()), times(1)); - } + assertEquals(3, chain.length); + assertEquals(1, keyVaultClient.getHttpCallCount(VERSIONED_SECRET_ID + API_VERSION_POSTFIX)); } @Test public void testCertificateChainJsonParsingFailureIsPropagated() { - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap())) - .thenReturn("{invalid-json"); - - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - IllegalStateException exception = assertThrows(IllegalStateException.class, - () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); - - assertEquals("Failed to parse certificate chain response for alias: " + CERTIFICATE_ALIAS, - exception.getMessage()); - assertTrue(exception.getCause() instanceof IOException); - } + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(VERSIONED_SECRET_ID + API_VERSION_POSTFIX, "{invalid-json"); + + IllegalStateException exception = assertThrows(IllegalStateException.class, + () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); + + assertEquals("Failed to parse certificate chain response for alias: " + CERTIFICATE_ALIAS, + exception.getMessage()); + assertTrue(exception.getCause() instanceof IOException); } @Test public void testCertificateChainMissingHttpResponseIsPropagated() { - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap())) - .thenReturn(null); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(VERSIONED_SECRET_ID + API_VERSION_POSTFIX, (Object) null); - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - IllegalStateException exception = assertThrows(IllegalStateException.class, - () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); + IllegalStateException exception = assertThrows(IllegalStateException.class, + () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); - assertEquals("Failed to load certificate chain response for alias: " + CERTIFICATE_ALIAS, - exception.getMessage()); - } + assertEquals("Failed to load certificate chain response for alias: " + CERTIFICATE_ALIAS, + exception.getMessage()); } @Test public void testCertificateChainHttpFailureIsPropagatedWithoutWrapping() { RuntimeException httpFailure = new RuntimeException("Key Vault returned HTTP 429"); - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap())) - .thenThrow(httpFailure); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(VERSIONED_SECRET_ID + API_VERSION_POSTFIX, httpFailure); - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - RuntimeException exception = assertThrows(RuntimeException.class, - () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); + RuntimeException exception = assertThrows(RuntimeException.class, + () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); - assertSame(httpFailure, exception); - } + assertSame(httpFailure, exception); } @Test public void testCertificateChainMissingSecretValueIsPropagated() { SecretBundle secretBundle = new SecretBundle(); - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap())) - .thenReturn(JsonConverterUtil.toJson(secretBundle)); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(VERSIONED_SECRET_ID + API_VERSION_POSTFIX, + JsonConverterUtil.toJson(secretBundle)); - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - IllegalStateException exception = assertThrows(IllegalStateException.class, - () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); + IllegalStateException exception = assertThrows(IllegalStateException.class, + () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); - assertEquals("Certificate chain response has no secret value for alias: " + CERTIFICATE_ALIAS, - exception.getMessage()); - } + assertEquals("Certificate chain response has no secret value for alias: " + CERTIFICATE_ALIAS, + exception.getMessage()); } @Test public void testCertificateChainMissingSecretBundleIsPropagated() { - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap())) - .thenReturn("null"); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(VERSIONED_SECRET_ID + API_VERSION_POSTFIX, "null"); - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - IllegalStateException exception = assertThrows(IllegalStateException.class, - () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); + IllegalStateException exception = assertThrows(IllegalStateException.class, + () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); - assertEquals("Certificate chain response has no secret value for alias: " + CERTIFICATE_ALIAS, - exception.getMessage()); - } + assertEquals("Certificate chain response has no secret value for alias: " + CERTIFICATE_ALIAS, + exception.getMessage()); } @Test @@ -475,18 +458,15 @@ public void testCertificateChainPemDecodingFailureIsPropagated() { SecretBundle secretBundle = new SecretBundle(); secretBundle.setValue("-----BEGIN CERTIFICATE-----\ninvalid\n-----END CERTIFICATE-----"); - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap())) - .thenReturn(JsonConverterUtil.toJson(secretBundle)); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(VERSIONED_SECRET_ID + API_VERSION_POSTFIX, + JsonConverterUtil.toJson(secretBundle)); - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - IllegalStateException exception = assertThrows(IllegalStateException.class, - () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); + IllegalStateException exception = assertThrows(IllegalStateException.class, + () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); - assertEquals("Failed to decode certificate chain for alias: " + CERTIFICATE_ALIAS, exception.getMessage()); - assertTrue(exception.getCause() instanceof CertificateException); - } + assertEquals("Failed to decode certificate chain for alias: " + CERTIFICATE_ALIAS, exception.getMessage()); + assertTrue(exception.getCause() instanceof CertificateException); } @Test @@ -494,19 +474,16 @@ public void testCertificateChainUnterminatedPemFailureIsPropagated() throws Exce SecretBundle secretBundle = new SecretBundle(); secretBundle.setValue(readUnterminatedCertificatePem()); - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap())) - .thenReturn(JsonConverterUtil.toJson(secretBundle)); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(VERSIONED_SECRET_ID + API_VERSION_POSTFIX, + JsonConverterUtil.toJson(secretBundle)); - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - IllegalStateException exception = assertThrows(IllegalStateException.class, - () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); + IllegalStateException exception = assertThrows(IllegalStateException.class, + () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); - assertEquals("Failed to decode certificate chain for alias: " + CERTIFICATE_ALIAS, exception.getMessage()); - assertTrue(exception.getCause() instanceof CertificateException); - assertEquals("Certificate PEM block is not terminated.", exception.getCause().getMessage()); - } + assertEquals("Failed to decode certificate chain for alias: " + CERTIFICATE_ALIAS, exception.getMessage()); + assertTrue(exception.getCause() instanceof CertificateException); + assertEquals("Certificate PEM block is not terminated.", exception.getCause().getMessage()); } @Test @@ -514,18 +491,15 @@ public void testCertificateChainPkcs12DecodingFailureIsPropagated() { SecretBundle secretBundle = new SecretBundle(); secretBundle.setValue(Base64.getEncoder().encodeToString(new byte[] { 1, 2, 3 })); - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap())) - .thenReturn(JsonConverterUtil.toJson(secretBundle)); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(VERSIONED_SECRET_ID + API_VERSION_POSTFIX, + JsonConverterUtil.toJson(secretBundle)); - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - IllegalStateException exception = assertThrows(IllegalStateException.class, - () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); + IllegalStateException exception = assertThrows(IllegalStateException.class, + () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); - assertEquals("Failed to decode certificate chain for alias: " + CERTIFICATE_ALIAS, exception.getMessage()); - assertNotNull(exception.getCause()); - } + assertEquals("Failed to decode certificate chain for alias: " + CERTIFICATE_ALIAS, exception.getMessage()); + assertNotNull(exception.getCause()); } @Test @@ -533,18 +507,15 @@ public void testCertificateChainInvalidBase64IsPropagated() { SecretBundle secretBundle = new SecretBundle(); secretBundle.setValue("not-valid-base64!"); - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap())) - .thenReturn(JsonConverterUtil.toJson(secretBundle)); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(VERSIONED_SECRET_ID + API_VERSION_POSTFIX, + JsonConverterUtil.toJson(secretBundle)); - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - IllegalStateException exception = assertThrows(IllegalStateException.class, - () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); + IllegalStateException exception = assertThrows(IllegalStateException.class, + () -> keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID))); - assertEquals("Failed to decode certificate chain for alias: " + CERTIFICATE_ALIAS, exception.getMessage()); - assertTrue(exception.getCause() instanceof IllegalArgumentException); - } + assertEquals("Failed to decode certificate chain for alias: " + CERTIFICATE_ALIAS, exception.getMessage()); + assertTrue(exception.getCause() instanceof IllegalArgumentException); } @Test @@ -557,20 +528,17 @@ public void testCertificateChainDecodingFailureIsRetried() throws Exception { Paths.get("src/test/resources/certificate-util/SecretBundle.value/3-certificates-in-chain.pem")), StandardCharsets.UTF_8)); - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap())) - .thenReturn(JsonConverterUtil.toJson(invalidBundle), JsonConverterUtil.toJson(validBundle)); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(VERSIONED_SECRET_ID + API_VERSION_POSTFIX, + JsonConverterUtil.toJson(invalidBundle), JsonConverterUtil.toJson(validBundle)); + CertificateVersion certificateVersion = createCertificateVersion(VERSIONED_SECRET_ID); - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - CertificateVersion certificateVersion = createCertificateVersion(VERSIONED_SECRET_ID); - assertThrows(IllegalStateException.class, - () -> keyVaultClient.getCertificateChainForVersion(certificateVersion)); - Certificate[] chain = keyVaultClient.getCertificateChainForVersion(certificateVersion); + assertThrows(IllegalStateException.class, + () -> keyVaultClient.getCertificateChainForVersion(certificateVersion)); + Certificate[] chain = keyVaultClient.getCertificateChainForVersion(certificateVersion); - assertEquals(3, chain.length); - utilities.verify(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap()), times(2)); - } + assertEquals(3, chain.length); + assertEquals(2, keyVaultClient.getHttpCallCount(VERSIONED_SECRET_ID + API_VERSION_POSTFIX)); } private static String readUnterminatedCertificatePem() throws IOException { @@ -583,48 +551,34 @@ private static String readUnterminatedCertificatePem() throws IOException { @Test public void testCertificateChainWithoutSecretIdReturnsEmpty() { - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - Certificate[] chain = keyVaultClient.getCertificateChainForVersion(createCertificateVersion(null)); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + Certificate[] chain = keyVaultClient.getCertificateChainForVersion(createCertificateVersion(null)); - assertEquals(0, chain.length); - utilities.verifyNoInteractions(); - } + assertEquals(0, chain.length); + assertEquals(0, keyVaultClient.getTotalHttpCallCount()); } @Test public void testCertificateChainWithoutResolvedVersionReturnsEmpty() { - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - Certificate[] chain = keyVaultClient.getCertificateChainForVersion(null); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + Certificate[] chain = keyVaultClient.getCertificateChainForVersion(null); - assertEquals(0, chain.length); - utilities.verifyNoInteractions(); - } + assertEquals(0, chain.length); + assertEquals(0, keyVaultClient.getTotalHttpCallCount()); } @Test public void testCertificateChainDecodedWithoutCertificatesReturnsEmpty() throws Exception { SecretBundle secretBundle = new SecretBundle(); - secretBundle.setValue("valid-empty-chain"); - - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class); - MockedStatic certificateUtilities = Mockito.mockStatic(CertificateUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap())) - .thenReturn(JsonConverterUtil.toJson(secretBundle)); - certificateUtilities - .when(() -> CertificateUtil.loadCertificatesFromSecretBundleValue("valid-empty-chain", false)) - .thenReturn(new Certificate[0]); - - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - Certificate[] chain - = keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID)); - - assertEquals(0, chain.length); - } + secretBundle.setValue(createEmptyPkcs12()); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(VERSIONED_SECRET_ID + API_VERSION_POSTFIX, + JsonConverterUtil.toJson(secretBundle)); + + Certificate[] chain + = keyVaultClient.getCertificateChainForVersion(createCertificateVersion(VERSIONED_SECRET_ID)); + + assertEquals(0, chain.length); } @Test @@ -637,38 +591,30 @@ public void testExportableKeyUsesVersionedSecretId() throws Exception { Paths.get("src/test/resources/certificate-util/SecretBundle.value/pkcs12-exportable-key.pfx")), StandardCharsets.UTF_8)); - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(CERTIFICATE_URI), anyMap())) - .thenReturn(JsonConverterUtil.toJson(certificateBundle)); - utilities.when(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap())) - .thenReturn(JsonConverterUtil.toJson(secretBundle)); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(CERTIFICATE_URI, JsonConverterUtil.toJson(certificateBundle)); + keyVaultClient.addHttpResponses(VERSIONED_SECRET_ID + API_VERSION_POSTFIX, + JsonConverterUtil.toJson(secretBundle)); - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - CertificateVersion certificateVersion = keyVaultClient.resolveCertificateVersion(CERTIFICATE_ALIAS); - Key key = keyVaultClient.getKeyForVersion(certificateVersion, null); + CertificateVersion certificateVersion = keyVaultClient.resolveCertificateVersion(CERTIFICATE_ALIAS); + Key key = keyVaultClient.getKeyForVersion(certificateVersion, null); - assertNotNull(key); - utilities.verify(() -> HttpUtil.get(eq(VERSIONED_SECRET_ID + API_VERSION_POSTFIX), anyMap()), times(1)); - } + assertNotNull(key); + assertEquals(1, keyVaultClient.getHttpCallCount(VERSIONED_SECRET_ID + API_VERSION_POSTFIX)); } @Test public void testKeylessKeyUsesVersionedKeyId() { CertificateBundle certificateBundle = createCertificateBundle(false); - try (MockedStatic utilities = Mockito.mockStatic(HttpUtil.class)) { - configureHttpUtilityMethods(utilities); - utilities.when(() -> HttpUtil.get(eq(CERTIFICATE_URI), anyMap())) - .thenReturn(JsonConverterUtil.toJson(certificateBundle)); + ScriptedKeyVaultClient keyVaultClient = createClientWithAccessToken(); + keyVaultClient.addHttpResponses(CERTIFICATE_URI, JsonConverterUtil.toJson(certificateBundle)); - KeyVaultClient keyVaultClient = createClientWithAccessToken(); - CertificateVersion certificateVersion = keyVaultClient.resolveCertificateVersion(CERTIFICATE_ALIAS); - Key key = keyVaultClient.getKeyForVersion(certificateVersion, null); + CertificateVersion certificateVersion = keyVaultClient.resolveCertificateVersion(CERTIFICATE_ALIAS); + Key key = keyVaultClient.getKeyForVersion(certificateVersion, null); - assertTrue(key instanceof KeyVaultPrivateKey); - assertEquals(VERSIONED_KEY_ID, ((KeyVaultPrivateKey) key).getKid()); - } + assertTrue(key instanceof KeyVaultPrivateKey); + assertEquals(VERSIONED_KEY_ID, ((KeyVaultPrivateKey) key).getKid()); } @Test @@ -721,19 +667,13 @@ public void close() { logger.setLevel(Level.ALL); logger.setUseParentHandlers(false); - try (MockedStatic httpUtilMockedStatic = Mockito.mockStatic(HttpUtil.class)) { - httpUtilMockedStatic.when(() -> HttpUtil.validateUri(anyString(), anyString())).thenCallRealMethod(); - httpUtilMockedStatic.when(() -> HttpUtil.addTrailingSlashIfRequired(anyString())).thenCallRealMethod(); - httpUtilMockedStatic - .when(() -> HttpUtil.get( - eq(KEY_VAULT_TEST_URI_GLOBAL + "certificates/" + alias + HttpUtil.API_VERSION_POSTFIX), anyMap())) - .thenReturn(JsonConverterUtil.toJson(certificateBundle)); - httpUtilMockedStatic - .when(() -> HttpUtil.get(eq(certificateSecretUri + HttpUtil.API_VERSION_POSTFIX), anyMap())) - .thenReturn(JsonConverterUtil.toJson(secretBundle)); - - KeyVaultClient keyVaultClient - = new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null, null, null, null, "bearer-token", false); + try { + ScriptedKeyVaultClient keyVaultClient = new ScriptedKeyVaultClient("bearer-token"); + keyVaultClient.addHttpResponses( + KEY_VAULT_TEST_URI_GLOBAL + "certificates/" + alias + HttpUtil.API_VERSION_POSTFIX, + JsonConverterUtil.toJson(certificateBundle)); + keyVaultClient.addHttpResponses(certificateSecretUri + HttpUtil.API_VERSION_POSTFIX, + JsonConverterUtil.toJson(secretBundle)); Key key = keyVaultClient.getKey(alias, null); assertNotNull(key); @@ -767,13 +707,56 @@ private static CertificateVersion createCertificateVersion(String secretId) { return new CertificateVersion(CERTIFICATE_ALIAS, null, VERSIONED_KEY_ID, secretId, true, "RSA"); } - private static KeyVaultClient createClientWithAccessToken() { - return new KeyVaultClient(KEY_VAULT_TEST_URI_GLOBAL, null, null, null, null, "test-token", false); + private static ScriptedKeyVaultClient createClientWithAccessToken() { + return new ScriptedKeyVaultClient("test-token"); + } + + private static String createEmptyPkcs12() throws Exception { + KeyStore keyStore = KeyStore.getInstance("PKCS12"); + keyStore.load(null, new char[0]); + ByteArrayOutputStream output = new ByteArrayOutputStream(); + keyStore.store(output, new char[0]); + return Base64.getEncoder().encodeToString(output.toByteArray()); } - private static void configureHttpUtilityMethods(MockedStatic utilities) { - utilities.when(() -> HttpUtil.validateUri(anyString(), anyString())).thenCallRealMethod(); - utilities.when(() -> HttpUtil.addTrailingSlashIfRequired(anyString())).thenCallRealMethod(); + private static final class ScriptedKeyVaultClient extends KeyVaultClient { + private final Map> httpResponses = new HashMap<>(); + private final Map httpCallCounts = new HashMap<>(); + + private ScriptedKeyVaultClient(String accessToken) { + super(KEY_VAULT_TEST_URI_GLOBAL, null, null, null, null, accessToken, false); + } + + private void addHttpResponses(String uri, Object... responses) { + httpResponses.put(uri, new ArrayList<>(Arrays.asList(responses))); + } + + private int getHttpCallCount(String uri) { + AtomicInteger count = httpCallCounts.get(uri); + return count == null ? 0 : count.get(); + } + + private int getTotalHttpCallCount() { + return httpCallCounts.values().stream().mapToInt(AtomicInteger::get).sum(); + } + + @Override + String httpGet(String uri, Map headers) { + int invocation = httpCallCounts.computeIfAbsent(uri, ignored -> new AtomicInteger()).getAndIncrement(); + List responses = httpResponses.get(uri); + if (responses == null || responses.isEmpty()) { + throw new AssertionError("Unexpected HTTP GET: " + uri); + } + + Object response = responses.get(Math.min(invocation, responses.size() - 1)); + if (response instanceof RuntimeException) { + throw (RuntimeException) response; + } + if (response instanceof Error) { + throw (Error) response; + } + return (String) response; + } } @EnabledIfEnvironmentVariable(named = "AZURE_KEYVAULT_CERTIFICATE_NAME", matches = "myalias") diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/MockKeyVaultClient.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/MockKeyVaultClient.java new file mode 100644 index 000000000000..5d3d656297e9 --- /dev/null +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/MockKeyVaultClient.java @@ -0,0 +1,9 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +package com.azure.security.keyvault.jca.implementation; + +public class MockKeyVaultClient extends KeyVaultClient { + public MockKeyVaultClient() { + super("https://accountname.vault.azure.net", "tenant-id", "client-id", "client-secret"); + } +} diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/TestCertificateVersions.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/TestCertificateVersions.java new file mode 100644 index 000000000000..6ca825e54cfa --- /dev/null +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/TestCertificateVersions.java @@ -0,0 +1,43 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +package com.azure.security.keyvault.jca.implementation; + +/** + * Shared test-only helper used to construct {@link CertificateVersion} instances from tests that live outside the + * {@code com.azure.security.keyvault.jca.implementation} package, since {@link CertificateVersion} is a final class + * with a package-private constructor. + */ +public final class TestCertificateVersions { + + private TestCertificateVersions() { + } + + /** + * Creates a {@link CertificateVersion} with the given alias and metadata. + * + * @param alias The certificate alias. + * @param certificateData The Base64-encoded DER certificate data. + * @param keyId The versioned Key Vault key ID. + * @param secretId The versioned Key Vault secret ID. + * @param exportable Whether the private key is exportable through the certificate's secret. + * @param keyType The Key Vault key type. + * @return A new {@link CertificateVersion} instance. + */ + public static CertificateVersion create(String alias, String certificateData, String keyId, String secretId, + boolean exportable, String keyType) { + return new CertificateVersion(alias, certificateData, keyId, secretId, exportable, keyType); + } + + /** + * Creates a {@link CertificateVersion} for the given alias with no additional metadata populated. Each + * invocation returns a distinct instance, which is useful for tests that need to tell "versions" apart by + * identity. + * + * @param alias The certificate alias. + * @return A new {@link CertificateVersion} instance. + */ + public static CertificateVersion create(String alias) { + return create(alias, null, null, null, false, null); + } +} diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/TestKeyVaultClient.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/TestKeyVaultClient.java new file mode 100644 index 000000000000..5dc879d3d136 --- /dev/null +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/TestKeyVaultClient.java @@ -0,0 +1,32 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +package com.azure.security.keyvault.jca.implementation; + +import java.util.Map; +import java.util.function.BiFunction; + +/** + * Test client that replaces Key Vault HTTP calls with a caller-provided response function. + */ +public final class TestKeyVaultClient extends KeyVaultClient { + private final BiFunction, String> httpGet; + + /** + * Creates a test client. + * + * @param accessToken the fixed access token + * @param disableAiaDownload whether AIA downloads are disabled + * @param httpGet the HTTP GET response function + */ + public TestKeyVaultClient(String accessToken, boolean disableAiaDownload, + BiFunction, String> httpGet) { + super("https://fake.vault.azure.net/", null, null, null, null, accessToken, false, disableAiaDownload); + this.httpGet = httpGet; + } + + @Override + String httpGet(String uri, Map headers) { + return httpGet.apply(uri, headers); + } +} diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/certificates/ClasspathCertificatesTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/certificates/ClasspathCertificatesTest.java index e2dab5e2f7b1..fed6265ab347 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/certificates/ClasspathCertificatesTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/certificates/ClasspathCertificatesTest.java @@ -3,15 +3,14 @@ package com.azure.security.keyvault.jca.implementation.certificates; -import static org.mockito.Mockito.mock; - -import java.security.cert.Certificate; +import com.azure.security.keyvault.jca.implementation.mocking.MockCertificate; import org.junit.jupiter.api.Assertions; import org.junit.jupiter.api.Test; -public class ClasspathCertificatesTest { +import java.security.cert.Certificate; - private final Certificate certificate = mock(Certificate.class); +public class ClasspathCertificatesTest { + private final Certificate certificate = new MockCertificate(); @Test public void testSetCertificateEntry() { 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 9b3ab4c25090..30a725aea0ef 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 @@ -3,21 +3,26 @@ 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.KeyVaultLoadStoreParameter; import com.azure.security.keyvault.jca.implementation.CertificateVersion; +import com.azure.security.keyvault.jca.implementation.TestCertificateVersions; import com.azure.security.keyvault.jca.implementation.KeyVaultClient; +import com.azure.security.keyvault.jca.implementation.mocking.MockCertificate; +import com.azure.security.keyvault.jca.implementation.mocking.MockKey; + +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + import java.lang.reflect.Field; import java.security.Key; import java.security.cert.Certificate; +import java.util.ArrayDeque; import java.util.ArrayList; import java.util.Arrays; import java.util.Collections; +import java.util.Deque; +import java.util.HashMap; import java.util.HashSet; import java.util.List; import java.util.Map; @@ -27,11 +32,6 @@ import java.util.concurrent.atomic.AtomicInteger; import java.util.function.Consumer; import java.util.function.Supplier; -import org.junit.jupiter.api.Assertions; -import org.junit.jupiter.api.BeforeEach; -import org.junit.jupiter.api.Test; -import org.mockito.invocation.InvocationOnMock; -import org.mockito.stubbing.Answer; public class KeyVaultCertificatesTest { @@ -39,13 +39,13 @@ public class KeyVaultCertificatesTest { private static final long TIMEOUT_MILLIS = 10_000; - private final KeyVaultClient keyVaultClient = mock(KeyVaultClient.class); + private final FakeKeyVaultClient keyVaultClient = new FakeKeyVaultClient(); - private final CertificateVersion certificateVersion = mock(CertificateVersion.class); + private final CertificateVersion certificateVersion = TestCertificateVersions.create("myalias"); - private final Key key = mock(Key.class); + private final Key key = new MockKey(); - private final Certificate certificate = mock(Certificate.class); + private final Certificate certificate = new MockCertificate(); private final Certificate[] certificateChain = new Certificate[] { certificate }; @@ -59,11 +59,14 @@ private KeyVaultCertificates createKeyVaultCertificates(KeyVaultClient client, S KeyVaultLoadStoreParameter parameter = new KeyVaultLoadStoreParameter(null).setCertificatesRefreshIntervalInMs(60_000) .setCertificateAliasFilterPatterns(filterPatterns); - KeyVaultCertificates certificates = new KeyVaultCertificates(parameter); - setKeyVaultClient(certificates, client); - return certificates; + return new KeyVaultCertificates(parameter, client); } + /** + * Reinjects a {@link KeyVaultClient} after {@link KeyVaultCertificates#updateKeyVaultClient} has nulled it out + * (which happens whenever the supplied {@link KeyVaultLoadStoreParameter} has no URI). There is no public API + * for this, so reflection is used purely to keep testing against the fake client afterward. + */ private void setKeyVaultClient(KeyVaultCertificates certificates, KeyVaultClient client) { try { Field keyVaultClientField = KeyVaultCertificates.class.getDeclaredField("keyVaultClient"); @@ -78,11 +81,11 @@ private void setKeyVaultClient(KeyVaultCertificates certificates, KeyVaultClient public void beforeEach() { List aliases = new ArrayList<>(); aliases.add("myalias"); - when(keyVaultClient.getAliases()).thenReturn(aliases); - when(keyVaultClient.resolveCertificateVersion("myalias")).thenReturn(certificateVersion); - when(keyVaultClient.getKeyForVersion(certificateVersion, null)).thenReturn(key); - when(keyVaultClient.getCertificateForVersion(certificateVersion)).thenReturn(certificate); - when(keyVaultClient.getCertificateChainForVersion(certificateVersion)).thenReturn(certificateChain); + keyVaultClient.stubAliases(aliases); + keyVaultClient.stubResolveCertificateVersion("myalias", certificateVersion); + keyVaultClient.stubKeyForVersion(certificateVersion, key); + keyVaultClient.stubCertificateForVersion(certificateVersion, certificate); + keyVaultClient.stubCertificateChainForVersion(certificateVersion, certificateChain); keyVaultCertificates = createKeyVaultCertificates(keyVaultClient); } @@ -168,11 +171,40 @@ public void testGetCertificateChainsReturnsSnapshot() { public void testRefreshAndGetAliasByCertificate() { Assertions.assertEquals(keyVaultCertificates.refreshAndGetAliasByCertificate(certificate), "myalias"); Assertions.assertEquals(keyVaultCertificates.getCertificates().get("myalias"), certificate); - when(keyVaultClient.getAliases()).thenReturn(null); + keyVaultClient.stubAliases(null); Assertions.assertNotEquals(keyVaultCertificates.refreshAndGetAliasByCertificate(certificate), "myalias"); Assertions.assertNull(keyVaultCertificates.getCertificates().get("myalias")); } + @Test + public void testRefreshAndGetAliasByCertificateReturnsMatchBeforeLaterRefreshInvalidatesIt() { + String matchingAlias = "first-alias"; + String otherAlias = "second-alias"; + CertificateVersion matchingVersion = TestCertificateVersions.create(matchingAlias); + CertificateVersion otherVersion = TestCertificateVersions.create(otherAlias); + Certificate otherCertificate = new MockCertificate() { + @Override + public byte[] getEncoded() { + return new byte[] { 1 }; + } + }; + + keyVaultClient.stubAliases(Arrays.asList(matchingAlias, otherAlias)); + keyVaultClient.stubResolveCertificateVersion(matchingAlias, matchingVersion); + keyVaultClient.stubResolveCertificateVersion(otherAlias, otherVersion); + keyVaultClient.stubCertificateForVersionAnswer(matchingVersion, () -> { + sleepUnchecked(50); + return certificate; + }); + keyVaultClient.stubCertificateForVersion(otherVersion, otherCertificate); + KeyVaultLoadStoreParameter parameter + = new KeyVaultLoadStoreParameter(null).setCertificatesRefreshIntervalInMs(1); + keyVaultCertificates = new KeyVaultCertificates(parameter, keyVaultClient); + + Assertions.assertEquals(matchingAlias, keyVaultCertificates.refreshAndGetAliasByCertificate(certificate)); + Assertions.assertEquals(0, keyVaultClient.resolveCertificateVersionCallCount(otherAlias)); + } + @Test public void testRefreshAndGetAliasByCertificateWithNullCertificate() { Assertions.assertNull(keyVaultCertificates.refreshAndGetAliasByCertificate(null)); @@ -189,10 +221,10 @@ public void testDeleteAlias() { public void testGetAliasesDoesNotLoadCertificateDetailsEagerly() { keyVaultCertificates.getAliases(); - verify(keyVaultClient, never()).resolveCertificateVersion("myalias"); - verify(keyVaultClient, never()).getKeyForVersion(certificateVersion, null); - verify(keyVaultClient, never()).getCertificateForVersion(certificateVersion); - verify(keyVaultClient, never()).getCertificateChainForVersion(certificateVersion); + Assertions.assertEquals(0, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(0, keyVaultClient.keyForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); } @Test @@ -201,73 +233,73 @@ public void testLoadCertificateDetailsForRequestedAliasOnly() { aliases.add("myalias"); aliases.add("otheralias"); - when(keyVaultClient.getAliases()).thenReturn(aliases); + keyVaultClient.stubAliases(aliases); keyVaultCertificates.getCertificate("myalias"); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getCertificateForVersion(certificateVersion); - verify(keyVaultClient, never()).getKeyForVersion(certificateVersion, null); - verify(keyVaultClient, never()).getCertificateChainForVersion(certificateVersion); - verify(keyVaultClient, never()).resolveCertificateVersion("otheralias"); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.certificateForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.keyForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.resolveCertificateVersionCallCount("otheralias")); } @Test public void testGetKeyLoadsOnlyKeyForRequestedAlias() { keyVaultCertificates.getCertificateKey("myalias"); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getKeyForVersion(certificateVersion, null); - verify(keyVaultClient, never()).getCertificateForVersion(certificateVersion); - verify(keyVaultClient, never()).getCertificateChainForVersion(certificateVersion); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.keyForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); } @Test public void testGetCertificateChainLoadsOnlyChainForRequestedAlias() { keyVaultCertificates.getCertificateChain("myalias"); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getCertificateChainForVersion(certificateVersion); - verify(keyVaultClient, never()).getKeyForVersion(certificateVersion, null); - verify(keyVaultClient, never()).getCertificateForVersion(certificateVersion); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.keyForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateForVersionCallCount(certificateVersion)); } @Test public void testCertificateChainAndKeyUseSameResolvedVersionUntilRefresh() { - CertificateVersion version1 = mock(CertificateVersion.class); - CertificateVersion version2 = mock(CertificateVersion.class); - Certificate version1Certificate = mock(Certificate.class); + CertificateVersion version1 = TestCertificateVersions.create("myalias"); + CertificateVersion version2 = TestCertificateVersions.create("myalias"); + Certificate version1Certificate = new MockCertificate(); Certificate[] version1Chain = new Certificate[] { version1Certificate }; - Key version1Key = mock(Key.class); - Key version2Key = mock(Key.class); + Key version1Key = new MockKey(); + Key version2Key = new MockKey(); // These stubs reproduce the pre-fix behavior when material was fetched independently by alias. - when(keyVaultClient.getCertificateChain("myalias")).thenReturn(version1Chain); - when(keyVaultClient.getKey("myalias", null)).thenReturn(version2Key); + keyVaultClient.stubLegacyCertificateChain("myalias", version1Chain); + keyVaultClient.stubLegacyKey("myalias", version2Key); - when(keyVaultClient.resolveCertificateVersion("myalias")).thenReturn(version1, version2); - when(keyVaultClient.getCertificateForVersion(version1)).thenReturn(version1Certificate); - when(keyVaultClient.getCertificateChainForVersion(version1)).thenReturn(version1Chain); - when(keyVaultClient.getKeyForVersion(version1, null)).thenReturn(version1Key); - when(keyVaultClient.getKeyForVersion(version2, null)).thenReturn(version2Key); + keyVaultClient.stubResolveCertificateVersion("myalias", version1, version2); + keyVaultClient.stubCertificateForVersion(version1, version1Certificate); + keyVaultClient.stubCertificateChainForVersion(version1, version1Chain); + keyVaultClient.stubKeyForVersion(version1, version1Key); + keyVaultClient.stubKeyForVersion(version2, version2Key); Assertions.assertArrayEquals(version1Chain, keyVaultCertificates.getCertificateChain("myalias")); Assertions.assertSame(version1Certificate, keyVaultCertificates.getCertificate("myalias")); Assertions.assertSame(version1Key, keyVaultCertificates.getCertificateKey("myalias")); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, never()).getKeyForVersion(version2, null); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(0, keyVaultClient.keyForVersionCallCount(version2)); keyVaultCertificates.refreshCertificates(); Assertions.assertSame(version2Key, keyVaultCertificates.getCertificateKey("myalias")); - verify(keyVaultClient, times(2)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getKeyForVersion(version2, null); + Assertions.assertEquals(2, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.keyForVersionCallCount(version2)); } @Test public void testConcurrentMaterialLoadsShareCertificateVersionResolution() throws Exception { BlockingAnswer blockingAnswer = new BlockingAnswer<>(certificateVersion); - when(keyVaultClient.resolveCertificateVersion("myalias")).thenAnswer(blockingAnswer); + keyVaultClient.stubResolveCertificateVersionAnswer("myalias", blockingAnswer); CountDownLatch readersReady = new CountDownLatch(3); CountDownLatch readersMayStart = new CountDownLatch(1); List readers = Arrays.asList( @@ -283,10 +315,10 @@ public void testConcurrentMaterialLoadsShareCertificateVersionResolution() throw blockingAnswer.release(); joinThreads(readers); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getCertificateForVersion(certificateVersion); - verify(keyVaultClient, times(1)).getCertificateChainForVersion(certificateVersion); - verify(keyVaultClient, times(1)).getKeyForVersion(certificateVersion, null); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.certificateForVersionCallCount(certificateVersion)); + Assertions.assertEquals(1, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); + Assertions.assertEquals(1, keyVaultClient.keyForVersionCallCount(certificateVersion)); } @Test @@ -294,7 +326,7 @@ public void testConfiguredAliasesFilter() { List aliases = new ArrayList<>(); aliases.add("myalias"); aliases.add("otheralias"); - when(keyVaultClient.getAliases()).thenReturn(aliases); + keyVaultClient.stubAliases(aliases); keyVaultCertificates = createKeyVaultCertificates(keyVaultClient, Collections.singleton("myalias")); @@ -309,12 +341,12 @@ public void testFilterPatternsIncludeRegex() { List aliases = new ArrayList<>(); aliases.add("prod-cert"); aliases.add("dev-cert"); - when(keyVaultClient.getAliases()).thenReturn(aliases); + keyVaultClient.stubAliases(aliases); keyVaultCertificates = createKeyVaultCertificates(keyVaultClient, Collections.singleton("^prod-.*")); Assertions.assertEquals(Collections.singletonList("prod-cert"), keyVaultCertificates.getAliases()); - verify(keyVaultClient, times(1)).getAliases(); + Assertions.assertEquals(1, keyVaultClient.aliasesCallCount()); } @Test @@ -322,13 +354,13 @@ public void testFilterPatternsExcludeRegex() { List aliases = new ArrayList<>(); aliases.add("prod-active"); aliases.add("prod-deprecated"); - when(keyVaultClient.getAliases()).thenReturn(aliases); + keyVaultClient.stubAliases(aliases); Set filterPatterns = new HashSet<>(Arrays.asList("^prod-.*", "!^prod-deprecated$")); keyVaultCertificates = createKeyVaultCertificates(keyVaultClient, filterPatterns); Assertions.assertEquals(Collections.singletonList("prod-active"), keyVaultCertificates.getAliases()); - verify(keyVaultClient, times(1)).getAliases(); + Assertions.assertEquals(1, keyVaultClient.aliasesCallCount()); } @Test @@ -344,17 +376,17 @@ public void testConfiguredAliasesFilterAfterRefresh() { Assertions.assertTrue(refreshedAliases.contains("myalias")); Assertions.assertFalse(refreshedAliases.contains("otheralias")); Assertions.assertFalse(refreshedAliases.contains("new")); - verify(keyVaultClient, times(2)).getAliases(); + Assertions.assertEquals(2, keyVaultClient.aliasesCallCount()); } @Test public void testConfiguredAliasesFilterUsesListApi() { - when(keyVaultClient.getAliases()).thenReturn(Arrays.asList("configured-alias", "other-alias")); + keyVaultClient.stubAliases(Arrays.asList("configured-alias", "other-alias")); keyVaultCertificates = createKeyVaultCertificates(keyVaultClient, Collections.singleton("configured-alias")); Assertions.assertEquals(Collections.singletonList("configured-alias"), keyVaultCertificates.getAliases()); - verify(keyVaultClient, times(1)).getAliases(); + Assertions.assertEquals(1, keyVaultClient.aliasesCallCount()); } @Test @@ -363,7 +395,7 @@ public void testConfiguredAliasesIgnoreNullEntries() { keyVaultCertificates = createKeyVaultCertificates(keyVaultClient, configuredAliases); Assertions.assertEquals(Collections.singletonList("myalias"), keyVaultCertificates.getAliases()); - verify(keyVaultClient, times(1)).getAliases(); + Assertions.assertEquals(1, keyVaultClient.aliasesCallCount()); } @Test @@ -376,7 +408,7 @@ public void testInvalidFilterPatternThrows() { @Test public void testFilterPatternWithBoundedQuantifier() { - when(keyVaultClient.getAliases()).thenReturn(Arrays.asList("cert-42", "cert-1234567", "cert-abc")); + keyVaultClient.stubAliases(Arrays.asList("cert-42", "cert-1234567", "cert-abc")); keyVaultCertificates = createKeyVaultCertificates(keyVaultClient, Collections.singleton("^cert-\\d{1,5}$")); @@ -389,147 +421,142 @@ public void testGetCertificateWithUnconfiguredAliasDoesNotFetchDetails() { Assertions.assertNull(keyVaultCertificates.getCertificate("otheralias")); - verify(keyVaultClient, never()).resolveCertificateVersion("otheralias"); + Assertions.assertEquals(0, keyVaultClient.resolveCertificateVersionCallCount("otheralias")); } @Test public void testAliasCertificateLoadFailureIsRetriedOnNextAccess() { RuntimeException loadFailure = new RuntimeException("transient certificate error"); - when(keyVaultClient.getCertificateForVersion(certificateVersion)).thenThrow(loadFailure) - .thenReturn(certificate); + keyVaultClient.stubCertificateForVersionThrowThenReturn(certificateVersion, loadFailure, certificate); RuntimeException thrown = Assertions.assertThrows(RuntimeException.class, () -> keyVaultCertificates.getCertificate("myalias")); Assertions.assertSame(loadFailure, thrown); Assertions.assertEquals(certificate, keyVaultCertificates.getCertificate("myalias")); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(2)).getCertificateForVersion(certificateVersion); - verify(keyVaultClient, never()).getKeyForVersion(certificateVersion, null); - verify(keyVaultClient, never()).getCertificateChainForVersion(certificateVersion); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(2, keyVaultClient.certificateForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.keyForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); } @Test public void testAliasCertificateNullLoadIsRetriedOnNextAccess() { - when(keyVaultClient.getCertificateForVersion(certificateVersion)).thenReturn(null).thenReturn(certificate); + keyVaultClient.stubCertificateForVersion(certificateVersion, null, certificate); Assertions.assertNull(keyVaultCertificates.getCertificate("myalias")); Assertions.assertEquals(certificate, keyVaultCertificates.getCertificate("myalias")); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(2)).getCertificateForVersion(certificateVersion); - verify(keyVaultClient, never()).getKeyForVersion(certificateVersion, null); - verify(keyVaultClient, never()).getCertificateChainForVersion(certificateVersion); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(2, keyVaultClient.certificateForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.keyForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); } @Test public void testAliasKeyLoadFailureIsRetriedOnNextAccess() { RuntimeException loadFailure = new RuntimeException("transient key error"); - when(keyVaultClient.getKeyForVersion(certificateVersion, null)).thenThrow(loadFailure).thenReturn(key); + keyVaultClient.stubKeyForVersionThrowThenReturn(certificateVersion, loadFailure, key); RuntimeException thrown = Assertions.assertThrows(RuntimeException.class, () -> keyVaultCertificates.getCertificateKey("myalias")); Assertions.assertSame(loadFailure, thrown); Assertions.assertEquals(key, keyVaultCertificates.getCertificateKey("myalias")); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(2)).getKeyForVersion(certificateVersion, null); - verify(keyVaultClient, never()).getCertificateForVersion(certificateVersion); - verify(keyVaultClient, never()).getCertificateChainForVersion(certificateVersion); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(2, keyVaultClient.keyForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); } @Test public void testAliasKeyNullLoadIsRetriedOnNextAccess() { - when(keyVaultClient.getKeyForVersion(certificateVersion, null)).thenReturn(null).thenReturn(key); + keyVaultClient.stubKeyForVersion(certificateVersion, null, key); Assertions.assertNull(keyVaultCertificates.getCertificateKey("myalias")); Assertions.assertEquals(key, keyVaultCertificates.getCertificateKey("myalias")); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(2)).getKeyForVersion(certificateVersion, null); - verify(keyVaultClient, never()).getCertificateForVersion(certificateVersion); - verify(keyVaultClient, never()).getCertificateChainForVersion(certificateVersion); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(2, keyVaultClient.keyForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); } @Test public void testAliasChainLoadFailureIsRetriedOnNextAccess() { RuntimeException loadFailure = new RuntimeException("Key Vault returned HTTP 403"); - when(keyVaultClient.getCertificateChainForVersion(certificateVersion)).thenThrow(loadFailure) - .thenReturn(certificateChain); + keyVaultClient.stubCertificateChainForVersionThrowThenReturn(certificateVersion, loadFailure, certificateChain); RuntimeException thrown = Assertions.assertThrows(RuntimeException.class, () -> keyVaultCertificates.getCertificateChain("myalias")); Assertions.assertSame(loadFailure, thrown); Assertions.assertArrayEquals(certificateChain, keyVaultCertificates.getCertificateChain("myalias")); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(2)).getCertificateChainForVersion(certificateVersion); - verify(keyVaultClient, never()).getCertificateForVersion(certificateVersion); - verify(keyVaultClient, never()).getKeyForVersion(certificateVersion, null); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(2, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.keyForVersionCallCount(certificateVersion)); } @Test public void testAliasChainEmptyLoadIsCached() { - when(keyVaultClient.getCertificateChainForVersion(certificateVersion)).thenReturn(new Certificate[0]) - .thenReturn(certificateChain); + keyVaultClient.stubCertificateChainForVersion(certificateVersion, new Certificate[0], certificateChain); Assertions.assertArrayEquals(new Certificate[0], keyVaultCertificates.getCertificateChain("myalias")); Assertions.assertArrayEquals(new Certificate[0], keyVaultCertificates.getCertificateChain("myalias")); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getCertificateChainForVersion(certificateVersion); - verify(keyVaultClient, never()).getCertificateForVersion(certificateVersion); - verify(keyVaultClient, never()).getKeyForVersion(certificateVersion, null); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.certificateForVersionCallCount(certificateVersion)); + Assertions.assertEquals(0, keyVaultClient.keyForVersionCallCount(certificateVersion)); } @Test public void testCertificateVersionResolutionFailureIsRetriedOnNextAccess() { RuntimeException resolutionFailure = new RuntimeException("Key Vault returned HTTP 403"); - when(keyVaultClient.resolveCertificateVersion("myalias")).thenThrow(resolutionFailure) - .thenReturn(certificateVersion); + keyVaultClient.stubResolveCertificateVersionThrowThenReturn("myalias", resolutionFailure, certificateVersion); RuntimeException thrown = Assertions.assertThrows(RuntimeException.class, () -> keyVaultCertificates.getCertificate("myalias")); Assertions.assertSame(resolutionFailure, thrown); Assertions.assertSame(certificate, keyVaultCertificates.getCertificate("myalias")); - verify(keyVaultClient, times(2)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getCertificateForVersion(certificateVersion); + Assertions.assertEquals(2, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.certificateForVersionCallCount(certificateVersion)); } @Test public void testConcurrentChainLoadFailureIsSharedAndRetried() throws Exception { RuntimeException loadFailure = new RuntimeException("Key Vault returned HTTP 429"); BlockingFailureAnswer blockingFailure = new BlockingFailureAnswer<>(loadFailure); - when(keyVaultClient.getCertificateChainForVersion(certificateVersion)).thenAnswer(blockingFailure) - .thenReturn(certificateChain); + keyVaultClient.stubCertificateChainForVersionAnswerThenReturn(certificateVersion, blockingFailure, + certificateChain); assertConcurrentFailureIsShared(() -> keyVaultCertificates.getCertificateChain("myalias"), blockingFailure, loadFailure, "loadMaterialIfNeeded"); - verify(keyVaultClient, times(1)).getCertificateChainForVersion(certificateVersion); + Assertions.assertEquals(1, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); Assertions.assertArrayEquals(certificateChain, keyVaultCertificates.getCertificateChain("myalias")); - verify(keyVaultClient, times(2)).getCertificateChainForVersion(certificateVersion); + Assertions.assertEquals(2, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); } @Test public void testConcurrentCertificateVersionResolutionFailureIsSharedAndRetried() throws Exception { RuntimeException resolutionFailure = new RuntimeException("Key Vault returned HTTP 403"); BlockingFailureAnswer blockingFailure = new BlockingFailureAnswer<>(resolutionFailure); - when(keyVaultClient.resolveCertificateVersion("myalias")).thenAnswer(blockingFailure) - .thenReturn(certificateVersion); + keyVaultClient.stubResolveCertificateVersionAnswerThenReturn("myalias", blockingFailure, certificateVersion); assertConcurrentFailureIsShared(() -> keyVaultCertificates.getCertificate("myalias"), blockingFailure, resolutionFailure, "resolveCertificateVersionIfNeeded"); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); Assertions.assertSame(certificate, keyVaultCertificates.getCertificate("myalias")); - verify(keyVaultClient, times(2)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getCertificateForVersion(certificateVersion); + Assertions.assertEquals(2, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.certificateForVersionCallCount(certificateVersion)); } @Test public void testRefreshDiscardsInFlightFailureFromPreviousGeneration() throws Exception { RuntimeException staleFailure = new RuntimeException("stale generation failure"); - CertificateVersion freshVersion = mock(CertificateVersion.class); - Certificate[] freshChain = new Certificate[] { mock(Certificate.class) }; + CertificateVersion freshVersion = TestCertificateVersions.create("myalias"); + Certificate[] freshChain = new Certificate[] { new MockCertificate() }; BlockingFailureAnswer blockingFailure = new BlockingFailureAnswer<>(staleFailure); - when(keyVaultClient.resolveCertificateVersion("myalias")).thenReturn(certificateVersion, freshVersion); - when(keyVaultClient.getCertificateChainForVersion(certificateVersion)).thenAnswer(blockingFailure); - when(keyVaultClient.getCertificateChainForVersion(freshVersion)).thenReturn(freshChain); + keyVaultClient.stubResolveCertificateVersion("myalias", certificateVersion, freshVersion); + keyVaultClient.stubCertificateChainForVersionAnswer(certificateVersion, blockingFailure); + keyVaultClient.stubCertificateChainForVersion(freshVersion, freshChain); List loadedChains = Collections.synchronizedList(new ArrayList<>()); List failures = Collections.synchronizedList(new ArrayList<>()); Thread reader = new Thread(() -> { @@ -549,8 +576,8 @@ public void testRefreshDiscardsInFlightFailureFromPreviousGeneration() throws Ex Assertions.assertTrue(failures.isEmpty()); Assertions.assertEquals(1, loadedChains.size()); Assertions.assertArrayEquals(freshChain, loadedChains.get(0)); - verify(keyVaultClient, times(1)).getCertificateChainForVersion(certificateVersion); - verify(keyVaultClient, times(1)).getCertificateChainForVersion(freshVersion); + Assertions.assertEquals(1, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); + Assertions.assertEquals(1, keyVaultClient.certificateChainForVersionCallCount(freshVersion)); } @Test @@ -569,7 +596,7 @@ public void testUpdateKeyVaultClientClearsCachedState() { @Test public void testUpdateKeyVaultClientAppliesAliasFilterPatterns() { - when(keyVaultClient.getAliases()).thenReturn(Arrays.asList("prod-cert", "dev-cert")); + keyVaultClient.stubAliases(Arrays.asList("prod-cert", "dev-cert")); KeyVaultLoadStoreParameter parameter = new KeyVaultLoadStoreParameter(null).setCertificateAliasFilterPatterns(Collections.singleton("^prod-.*")); @@ -581,7 +608,7 @@ public void testUpdateKeyVaultClientAppliesAliasFilterPatterns() { @Test public void testInvalidClientUpdatePreservesExistingAliasFilters() { - when(keyVaultClient.getAliases()).thenReturn(Arrays.asList("myalias", "otheralias")); + keyVaultClient.stubAliases(Arrays.asList("myalias", "otheralias")); keyVaultCertificates = createKeyVaultCertificates(keyVaultClient, Collections.singleton("myalias")); Assertions.assertEquals(Collections.singletonList("myalias"), keyVaultCertificates.getAliases()); @@ -599,94 +626,94 @@ public void testInvalidClientUpdatePreservesExistingAliasFilters() { @Test public void testConcurrentCertificateLoadsShareSingleRequest() throws Exception { BlockingAnswer blockingAnswer = new BlockingAnswer<>(certificate); - when(keyVaultClient.getCertificateForVersion(certificateVersion)).thenAnswer(blockingAnswer); + keyVaultClient.stubCertificateForVersionAnswer(certificateVersion, blockingAnswer); assertConcurrentLoadsShareSingleRequest(() -> keyVaultCertificates.getCertificate("myalias"), blockingAnswer, loadedCertificate -> Assertions.assertSame(certificate, loadedCertificate)); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getCertificateForVersion(certificateVersion); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.certificateForVersionCallCount(certificateVersion)); } @Test public void testConcurrentCertificateChainLoadsShareSingleRequest() throws Exception { BlockingAnswer blockingAnswer = new BlockingAnswer<>(certificateChain); - when(keyVaultClient.getCertificateChainForVersion(certificateVersion)).thenAnswer(blockingAnswer); + keyVaultClient.stubCertificateChainForVersionAnswer(certificateVersion, blockingAnswer); assertConcurrentLoadsShareSingleRequest(() -> keyVaultCertificates.getCertificateChain("myalias"), blockingAnswer, loadedChain -> Assertions.assertArrayEquals(certificateChain, loadedChain)); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getCertificateChainForVersion(certificateVersion); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); } @Test public void testConcurrentKeyLoadsShareSingleRequest() throws Exception { BlockingAnswer blockingAnswer = new BlockingAnswer<>(key); - when(keyVaultClient.getKeyForVersion(certificateVersion, null)).thenAnswer(blockingAnswer); + keyVaultClient.stubKeyForVersionAnswer(certificateVersion, blockingAnswer); assertConcurrentLoadsShareSingleRequest(() -> keyVaultCertificates.getCertificateKey("myalias"), blockingAnswer, loadedKey -> Assertions.assertSame(key, loadedKey)); - verify(keyVaultClient, times(1)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getKeyForVersion(certificateVersion, null); + Assertions.assertEquals(1, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.keyForVersionCallCount(certificateVersion)); } @Test public void testRefreshDiscardsInFlightCertificateFromPreviousGeneration() throws Exception { - Certificate freshCertificate = mock(Certificate.class); - CertificateVersion freshVersion = mock(CertificateVersion.class); + Certificate freshCertificate = new MockCertificate(); + CertificateVersion freshVersion = TestCertificateVersions.create("myalias"); BlockingAnswer staleAnswer = new BlockingAnswer<>(certificate); - when(keyVaultClient.resolveCertificateVersion("myalias")).thenReturn(certificateVersion, freshVersion); - when(keyVaultClient.getCertificateForVersion(certificateVersion)).thenAnswer(staleAnswer); - when(keyVaultClient.getCertificateForVersion(freshVersion)).thenReturn(freshCertificate); + keyVaultClient.stubResolveCertificateVersion("myalias", certificateVersion, freshVersion); + keyVaultClient.stubCertificateForVersionAnswer(certificateVersion, staleAnswer); + keyVaultClient.stubCertificateForVersion(freshVersion, freshCertificate); assertRefreshDiscardsStaleLoad(() -> keyVaultCertificates.getCertificate("myalias"), staleAnswer, loadedCertificate -> Assertions.assertSame(freshCertificate, loadedCertificate)); - verify(keyVaultClient, times(2)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getCertificateForVersion(certificateVersion); - verify(keyVaultClient, times(1)).getCertificateForVersion(freshVersion); + Assertions.assertEquals(2, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.certificateForVersionCallCount(certificateVersion)); + Assertions.assertEquals(1, keyVaultClient.certificateForVersionCallCount(freshVersion)); } @Test public void testRefreshDiscardsInFlightCertificateChainFromPreviousGeneration() throws Exception { - Certificate[] freshChain = new Certificate[] { mock(Certificate.class) }; - CertificateVersion freshVersion = mock(CertificateVersion.class); + Certificate[] freshChain = new Certificate[] { new MockCertificate() }; + CertificateVersion freshVersion = TestCertificateVersions.create("myalias"); BlockingAnswer staleAnswer = new BlockingAnswer<>(certificateChain); - when(keyVaultClient.resolveCertificateVersion("myalias")).thenReturn(certificateVersion, freshVersion); - when(keyVaultClient.getCertificateChainForVersion(certificateVersion)).thenAnswer(staleAnswer); - when(keyVaultClient.getCertificateChainForVersion(freshVersion)).thenReturn(freshChain); + keyVaultClient.stubResolveCertificateVersion("myalias", certificateVersion, freshVersion); + keyVaultClient.stubCertificateChainForVersionAnswer(certificateVersion, staleAnswer); + keyVaultClient.stubCertificateChainForVersion(freshVersion, freshChain); assertRefreshDiscardsStaleLoad(() -> keyVaultCertificates.getCertificateChain("myalias"), staleAnswer, loadedChain -> Assertions.assertArrayEquals(freshChain, loadedChain)); - verify(keyVaultClient, times(2)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getCertificateChainForVersion(certificateVersion); - verify(keyVaultClient, times(1)).getCertificateChainForVersion(freshVersion); + Assertions.assertEquals(2, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.certificateChainForVersionCallCount(certificateVersion)); + Assertions.assertEquals(1, keyVaultClient.certificateChainForVersionCallCount(freshVersion)); } @Test public void testRefreshDiscardsInFlightKeyFromPreviousGeneration() throws Exception { - Key freshKey = mock(Key.class); - CertificateVersion freshVersion = mock(CertificateVersion.class); + Key freshKey = new MockKey(); + CertificateVersion freshVersion = TestCertificateVersions.create("myalias"); BlockingAnswer staleAnswer = new BlockingAnswer<>(key); - when(keyVaultClient.resolveCertificateVersion("myalias")).thenReturn(certificateVersion, freshVersion); - when(keyVaultClient.getKeyForVersion(certificateVersion, null)).thenAnswer(staleAnswer); - when(keyVaultClient.getKeyForVersion(freshVersion, null)).thenReturn(freshKey); + keyVaultClient.stubResolveCertificateVersion("myalias", certificateVersion, freshVersion); + keyVaultClient.stubKeyForVersionAnswer(certificateVersion, staleAnswer); + keyVaultClient.stubKeyForVersion(freshVersion, freshKey); assertRefreshDiscardsStaleLoad(() -> keyVaultCertificates.getCertificateKey("myalias"), staleAnswer, loadedKey -> Assertions.assertSame(freshKey, loadedKey)); - verify(keyVaultClient, times(2)).resolveCertificateVersion("myalias"); - verify(keyVaultClient, times(1)).getKeyForVersion(certificateVersion, null); - verify(keyVaultClient, times(1)).getKeyForVersion(freshVersion, null); + Assertions.assertEquals(2, keyVaultClient.resolveCertificateVersionCallCount("myalias")); + Assertions.assertEquals(1, keyVaultClient.keyForVersionCallCount(certificateVersion)); + Assertions.assertEquals(1, keyVaultClient.keyForVersionCallCount(freshVersion)); } @Test public void testClientReplacementDiscardsInFlightCertificate() throws Exception { BlockingAnswer staleAnswer = new BlockingAnswer<>(certificate); - when(keyVaultClient.getCertificateForVersion(certificateVersion)).thenAnswer(staleAnswer); + keyVaultClient.stubCertificateForVersionAnswer(certificateVersion, staleAnswer); List loadedValues = Collections.synchronizedList(new ArrayList<>()); Thread reader = new Thread(() -> loadedValues.add(keyVaultCertificates.getCertificate("myalias"))); reader.start(); @@ -699,7 +726,7 @@ public void testClientReplacementDiscardsInFlightCertificate() throws Exception Assertions.assertEquals(1, loadedValues.size()); Assertions.assertNull(loadedValues.get(0)); Assertions.assertTrue(keyVaultCertificates.getCertificates().isEmpty()); - verify(keyVaultClient, times(1)).getCertificateForVersion(certificateVersion); + Assertions.assertEquals(1, keyVaultClient.certificateForVersionCallCount(certificateVersion)); } @Test @@ -708,10 +735,10 @@ public void testConcurrentForceRefreshAppliesLatestAliases() throws Exception { CountDownLatch firstListCallMayFinish = new CountDownLatch(1); AtomicInteger listCallCount = new AtomicInteger(); - when(keyVaultClient.getAliases()).thenAnswer(invocation -> { + keyVaultClient.stubAliasesAnswer(() -> { if (listCallCount.getAndIncrement() == 0) { firstListCallStarted.countDown(); - awaitLatch(firstListCallMayFinish); + awaitLatchUnchecked(firstListCallMayFinish); return Collections.singletonList("stale-alias"); } return Collections.singletonList("fresh-alias"); @@ -739,9 +766,9 @@ public void testConcurrentRefreshIssuesSingleAliasListCall() throws Exception { CountDownLatch listCallStarted = new CountDownLatch(1); CountDownLatch listCallMayFinish = new CountDownLatch(1); - when(keyVaultClient.getAliases()).thenAnswer(invocation -> { + keyVaultClient.stubAliasesAnswer(() -> { listCallStarted.countDown(); - awaitLatch(listCallMayFinish); + awaitLatchUnchecked(listCallMayFinish); return Collections.singletonList("myalias"); }); @@ -762,7 +789,7 @@ public void testConcurrentRefreshIssuesSingleAliasListCall() throws Exception { reader.join(TIMEOUT_MILLIS); } - verify(keyVaultClient, times(1)).getAliases(); + Assertions.assertEquals(1, keyVaultClient.aliasesCallCount()); } private void assertConcurrentLoadsShareSingleRequest(Supplier load, BlockingAnswer blockingAnswer, @@ -875,6 +902,28 @@ private static void awaitLatch(CountDownLatch latch) throws InterruptedException } } + /** + * Same as {@link #awaitLatch(CountDownLatch)}, but usable from {@link Supplier#get()} lambdas, which cannot + * declare a checked {@link InterruptedException}. + */ + private static void awaitLatchUnchecked(CountDownLatch latch) { + try { + awaitLatch(latch); + } catch (InterruptedException exception) { + Thread.currentThread().interrupt(); + throw new IllegalStateException("Interrupted while waiting for the test latch.", exception); + } + } + + private static void sleepUnchecked(long milliseconds) { + try { + Thread.sleep(milliseconds); + } catch (InterruptedException exception) { + Thread.currentThread().interrupt(); + throw new IllegalStateException("Interrupted while waiting for the refresh interval to expire.", exception); + } + } + /** * Waits until none of the threads can still reach Key Vault, so the pending call cannot be released too early. */ @@ -919,9 +968,8 @@ private static void awaitSingleFlightWaiters(List readers, String waiter long deadline = System.nanoTime() + TimeUnit.MILLISECONDS.toNanos(TIMEOUT_MILLIS); while (System.nanoTime() < deadline) { - long ownerCount = readers.stream() - .filter(reader -> hasStackFrame(reader, BlockingFailureAnswer.class, "answer")) - .count(); + long ownerCount + = readers.stream().filter(reader -> hasStackFrame(reader, BlockingFailureAnswer.class, "get")).count(); long waiterCount = readers.stream() .filter(reader -> hasStackFrameCalledBy(reader, KeyVaultCertificates.class, "awaitInFlightOperation", waiterCallerMethod)) @@ -959,7 +1007,11 @@ private static boolean hasStackFrame(Thread thread, Class declaringClass, Str && frame.getMethodName().equals(methodName)); } - private static final class BlockingAnswer implements Answer { + /** + * Scripted response that blocks until released, so tests can deterministically observe an in-flight Key Vault + * call before letting it complete. + */ + private static final class BlockingAnswer implements Supplier { private final CountDownLatch started = new CountDownLatch(1); @@ -972,9 +1024,9 @@ private BlockingAnswer(T result) { } @Override - public T answer(InvocationOnMock invocation) throws InterruptedException { + public T get() { started.countDown(); - awaitLatch(mayFinish); + awaitLatchUnchecked(mayFinish); return result; } @@ -987,7 +1039,10 @@ private void release() { } } - private static final class BlockingFailureAnswer implements Answer { + /** + * Same as {@link BlockingAnswer}, but throws a scripted failure instead of returning a value once released. + */ + private static final class BlockingFailureAnswer implements Supplier { private final CountDownLatch started = new CountDownLatch(1); @@ -1000,9 +1055,9 @@ private BlockingFailureAnswer(RuntimeException failure) { } @Override - public T answer(InvocationOnMock invocation) throws InterruptedException { + public T get() { started.countDown(); - awaitLatch(mayFail); + awaitLatchUnchecked(mayFail); throw failure; } @@ -1015,4 +1070,244 @@ private void release() { } } + /** + * Handwritten {@link KeyVaultClient} test double. Every remote call is scripted through a + * {@link ScriptedResponses} instance keyed by whatever uniquely identifies the call (nothing for + * {@link #getAliases()}, the alias for {@link #resolveCertificateVersion(String)}, and the resolved + * {@link CertificateVersion} identity for the per-version material getters), which mirrors how the production + * {@code KeyVaultCertificates} class keys its own caches. + */ + private static final class FakeKeyVaultClient extends KeyVaultClient { + + private final ScriptedResponses> aliasesScript = new ScriptedResponses<>(); + + private final ScriptedResponses resolveCertificateVersionScript + = new ScriptedResponses<>(); + + private final ScriptedResponses certificateForVersionScript + = new ScriptedResponses<>(); + + private final ScriptedResponses certificateChainForVersionScript + = new ScriptedResponses<>(); + + private final ScriptedResponses keyForVersionScript = new ScriptedResponses<>(); + + private final Map legacyCertificateChains = new HashMap<>(); + + private final Map legacyKeys = new HashMap<>(); + + private FakeKeyVaultClient() { + super("https://accountname.vault.azure.net", "tenant-id", "client-id", "client-secret"); + } + + private void stubAliases(List aliases) { + aliasesScript.returnValues(null, Collections.singletonList(aliases)); + } + + private void stubAliasesAnswer(Supplier> supplier) { + aliasesScript.answer(null, supplier); + } + + private int aliasesCallCount() { + return aliasesScript.callCount(null); + } + + @Override + public List getAliases() { + return aliasesScript.invoke(null); + } + + private void stubResolveCertificateVersion(String alias, CertificateVersion... versions) { + resolveCertificateVersionScript.returnValues(alias, Arrays.asList(versions)); + } + + private void stubResolveCertificateVersionThrowThenReturn(String alias, RuntimeException exception, + CertificateVersion version) { + resolveCertificateVersionScript.throwThenReturn(alias, exception, version); + } + + private void stubResolveCertificateVersionAnswer(String alias, Supplier supplier) { + resolveCertificateVersionScript.answer(alias, supplier); + } + + private void stubResolveCertificateVersionAnswerThenReturn(String alias, Supplier supplier, + CertificateVersion version) { + resolveCertificateVersionScript.answerThenReturn(alias, supplier, version); + } + + private int resolveCertificateVersionCallCount(String alias) { + return resolveCertificateVersionScript.callCount(alias); + } + + @Override + public CertificateVersion resolveCertificateVersion(String alias) { + return resolveCertificateVersionScript.invoke(alias); + } + + private void stubCertificateForVersion(CertificateVersion version, Certificate... certificates) { + certificateForVersionScript.returnValues(version, Arrays.asList(certificates)); + } + + private void stubCertificateForVersionThrowThenReturn(CertificateVersion version, RuntimeException exception, + Certificate value) { + certificateForVersionScript.throwThenReturn(version, exception, value); + } + + private void stubCertificateForVersionAnswer(CertificateVersion version, Supplier supplier) { + certificateForVersionScript.answer(version, supplier); + } + + private int certificateForVersionCallCount(CertificateVersion version) { + return certificateForVersionScript.callCount(version); + } + + @Override + public Certificate getCertificateForVersion(CertificateVersion version) { + return certificateForVersionScript.invoke(version); + } + + private void stubCertificateChainForVersion(CertificateVersion version, Certificate[]... chains) { + certificateChainForVersionScript.returnValues(version, Arrays.asList(chains)); + } + + private void stubCertificateChainForVersionThrowThenReturn(CertificateVersion version, + RuntimeException exception, Certificate[] value) { + certificateChainForVersionScript.throwThenReturn(version, exception, value); + } + + private void stubCertificateChainForVersionAnswer(CertificateVersion version, + Supplier supplier) { + certificateChainForVersionScript.answer(version, supplier); + } + + private void stubCertificateChainForVersionAnswerThenReturn(CertificateVersion version, + Supplier supplier, Certificate[] value) { + certificateChainForVersionScript.answerThenReturn(version, supplier, value); + } + + private int certificateChainForVersionCallCount(CertificateVersion version) { + return certificateChainForVersionScript.callCount(version); + } + + @Override + public Certificate[] getCertificateChainForVersion(CertificateVersion version) { + return certificateChainForVersionScript.invoke(version); + } + + private void stubKeyForVersion(CertificateVersion version, Key... keys) { + keyForVersionScript.returnValues(version, Arrays.asList(keys)); + } + + private void stubKeyForVersionThrowThenReturn(CertificateVersion version, RuntimeException exception, + Key value) { + keyForVersionScript.throwThenReturn(version, exception, value); + } + + private void stubKeyForVersionAnswer(CertificateVersion version, Supplier supplier) { + keyForVersionScript.answer(version, supplier); + } + + private int keyForVersionCallCount(CertificateVersion version) { + return keyForVersionScript.callCount(version); + } + + @Override + public Key getKeyForVersion(CertificateVersion version, char[] password) { + return keyForVersionScript.invoke(version); + } + + // Legacy, alias-keyed accessors. Only used to reproduce pre-refactor stubs in one test; never verified. + private void stubLegacyCertificateChain(String alias, Certificate[] chain) { + legacyCertificateChains.put(alias, chain); + } + + private void stubLegacyKey(String alias, Key key) { + legacyKeys.put(alias, key); + } + + @Override + public Certificate[] getCertificateChain(String alias) { + return legacyCertificateChains.get(alias); + } + + @Override + public Key getKey(String alias, char[] password) { + return legacyKeys.get(alias); + } + } + + /** + * Scripted return values, failures, responses, and call counts keyed by whatever value identifies a given + * response (for example a {@code null} constant, an alias, or a {@link CertificateVersion}). + *

+ * Each key is scripted with a queue of {@link Supplier}s. Calling {@link #invoke(Object)} consumes queued + * suppliers one at a time until only one remains, at which point that last supplier is replayed for every + * subsequent call. + * + * @param The type of the key used to identify a stub (for example a {@link String} alias). + * @param The type of value returned by a stub. + */ + private static final class ScriptedResponses { + + private final Map>> queuedSuppliers = new HashMap<>(); + + private final Map callCounts = new HashMap<>(); + + private synchronized void sequence(K key, List> suppliers) { + queuedSuppliers.put(key, new ArrayDeque<>(suppliers)); + callCounts.put(key, new AtomicInteger()); + } + + private void returnValues(K key, List values) { + List> suppliers = new ArrayList<>(); + for (V value : values) { + suppliers.add(() -> value); + } + sequence(key, suppliers); + } + + private void throwThenReturn(K key, RuntimeException exception, V value) { + List> suppliers = new ArrayList<>(); + suppliers.add(() -> { + throw exception; + }); + suppliers.add(() -> value); + sequence(key, suppliers); + } + + private void answer(K key, Supplier supplier) { + sequence(key, Collections.singletonList(supplier)); + } + + private void answerThenReturn(K key, Supplier supplier, V value) { + List> suppliers = new ArrayList<>(); + suppliers.add(supplier); + suppliers.add(() -> value); + sequence(key, suppliers); + } + + private int callCount(K key) { + AtomicInteger counter; + synchronized (this) { + counter = callCounts.get(key); + } + return counter == null ? 0 : counter.get(); + } + + private V invoke(K key) { + Supplier supplier; + synchronized (this) { + AtomicInteger counter = callCounts.computeIfAbsent(key, unused -> new AtomicInteger()); + counter.incrementAndGet(); + Deque> suppliers = queuedSuppliers.get(key); + if (suppliers == null || suppliers.isEmpty()) { + throw new AssertionError("Unexpected call for key: " + key); + } + supplier = suppliers.size() > 1 ? suppliers.poll() : suppliers.peek(); + } + // The actual supplier invocation happens outside the synchronized block so blocking suppliers used by + // concurrency tests don't serialize unrelated calls against this same script. + return supplier.get(); + } + } } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/certificates/SpecificPathCertificatesTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/certificates/SpecificPathCertificatesTest.java index 3a038616a521..687ab4700736 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/certificates/SpecificPathCertificatesTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/certificates/SpecificPathCertificatesTest.java @@ -6,13 +6,15 @@ import org.junit.jupiter.api.Assertions; import org.junit.jupiter.api.Test; +import java.nio.file.FileSystems; + public class SpecificPathCertificatesTest { SpecificPathCertificates specificPathCertificates; public static String getFilePath(String packageName) { String filepath = "\\src\\test\\resources\\" + packageName; - return System.getProperty("user.dir") + filepath.replace("\\", System.getProperty("file.separator")); + return System.getProperty("user.dir") + filepath.replace("\\", FileSystems.getDefault().getSeparator()); } @Test diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/mocking/MockCertificate.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/mocking/MockCertificate.java new file mode 100644 index 000000000000..b4afd69c7ea9 --- /dev/null +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/mocking/MockCertificate.java @@ -0,0 +1,48 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +package com.azure.security.keyvault.jca.implementation.mocking; + +import java.security.InvalidKeyException; +import java.security.NoSuchAlgorithmException; +import java.security.NoSuchProviderException; +import java.security.PublicKey; +import java.security.SignatureException; +import java.security.cert.Certificate; +import java.security.cert.CertificateEncodingException; +import java.security.cert.CertificateException; + +/** + * Mock of {@link Certificate}. + */ +public class MockCertificate extends Certificate { + public MockCertificate() { + super("mock"); + } + + @Override + public byte[] getEncoded() throws CertificateEncodingException { + return new byte[0]; + } + + @Override + public void verify(PublicKey key) throws CertificateException, NoSuchAlgorithmException, InvalidKeyException, + NoSuchProviderException, SignatureException { + + } + + @Override + public void verify(PublicKey key, String sigProvider) throws CertificateException, NoSuchAlgorithmException, + InvalidKeyException, NoSuchProviderException, SignatureException { + + } + + @Override + public String toString() { + return ""; + } + + @Override + public PublicKey getPublicKey() { + return null; + } +} diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/mocking/MockKey.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/mocking/MockKey.java new file mode 100644 index 000000000000..c6021f17dc0f --- /dev/null +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/mocking/MockKey.java @@ -0,0 +1,22 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +package com.azure.security.keyvault.jca.implementation.mocking; + +import java.security.Key; + +public class MockKey implements Key { + @Override + public String getAlgorithm() { + return ""; + } + + @Override + public String getFormat() { + return ""; + } + + @Override + public byte[] getEncoded() { + return new byte[0]; + } +} diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/mocking/MockPrivateKey.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/mocking/MockPrivateKey.java new file mode 100644 index 000000000000..b0fbd78f676b --- /dev/null +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/mocking/MockPrivateKey.java @@ -0,0 +1,22 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +package com.azure.security.keyvault.jca.implementation.mocking; + +import java.security.PrivateKey; + +public class MockPrivateKey implements PrivateKey { + @Override + public String getAlgorithm() { + return null; + } + + @Override + public String getFormat() { + return null; + } + + @Override + public byte[] getEncoded() { + return new byte[0]; + } +} diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/mocking/MockPublicKey.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/mocking/MockPublicKey.java new file mode 100644 index 000000000000..4ee99bb9511e --- /dev/null +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/mocking/MockPublicKey.java @@ -0,0 +1,22 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +package com.azure.security.keyvault.jca.implementation.mocking; + +import java.security.PublicKey; + +public class MockPublicKey implements PublicKey { + @Override + public String getAlgorithm() { + return null; + } + + @Override + public String getFormat() { + return null; + } + + @Override + public byte[] getEncoded() { + return new byte[0]; + } +} diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSignatureTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSignatureTest.java index 684e9f72891c..698030af7bb2 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSignatureTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessEcSignatureTest.java @@ -5,73 +5,59 @@ import com.azure.security.keyvault.jca.KeyVaultEncode; import com.azure.security.keyvault.jca.KeyVaultJcaPropertyNames; -import com.azure.security.keyvault.jca.implementation.KeyVaultPrivateKey; import com.azure.security.keyvault.jca.implementation.KeyVaultClient; +import com.azure.security.keyvault.jca.implementation.KeyVaultPrivateKey; +import com.azure.security.keyvault.jca.implementation.MockKeyVaultClient; +import com.azure.security.keyvault.jca.implementation.mocking.MockPrivateKey; +import com.azure.security.keyvault.jca.implementation.mocking.MockPublicKey; import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; -import org.mockito.ArgumentMatchers; +import org.junit.jupiter.api.parallel.ResourceLock; +import org.junit.jupiter.api.parallel.Resources; import java.security.PrivateKey; import java.security.PublicKey; -import static org.junit.jupiter.api.Assertions.*; -import static org.mockito.ArgumentMatchers.anyString; -import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.when; +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +@ResourceLock(Resources.SYSTEM_PROPERTIES) public class KeyVaultKeylessEcSignatureTest { KeyVaultKeylessEcSignature keyVaultKeylessEcSignature; - private final KeyVaultClient keyVaultClient = mock(KeyVaultClient.class); - - private final KeyVaultPrivateKey keyVaultPrivateKey = mock(KeyVaultPrivateKey.class); + private KeyVaultClient keyVaultClient; private final byte[] signedWithES256 = "fake256Value".getBytes(); private final byte[] signedWithES384 = "fake384Value".getBytes(); + private final PublicKey publicKey = new MockPublicKey(); + private final PrivateKey privateKey = new MockPrivateKey(); static final String KEY_VAULT_TEST_URI_GLOBAL = "https://fake.vault.azure.net/"; + private String previousKeyVaultUri; + @BeforeEach public void before() { + previousKeyVaultUri = System.getProperty(KeyVaultJcaPropertyNames.KEYVAULT_URI); System.setProperty(KeyVaultJcaPropertyNames.KEYVAULT_URI, KEY_VAULT_TEST_URI_GLOBAL); keyVaultKeylessEcSignature = new KeyVaultKeylessEcSha256Signature(); } - private final PublicKey publicKey = new PublicKey() { - @Override - public String getAlgorithm() { - return null; - } - - @Override - public String getFormat() { - return null; - } - - @Override - public byte[] getEncoded() { - return new byte[0]; - } - }; - - private final PrivateKey privateKey = new PrivateKey() { - @Override - public String getAlgorithm() { - return null; - } - - @Override - public String getFormat() { - return null; - } + @AfterEach + public void after() { + restoreKeyVaultUri(previousKeyVaultUri); + } - @Override - public byte[] getEncoded() { - return new byte[0]; + private static void restoreKeyVaultUri(String value) { + if (value == null) { + System.clearProperty(KeyVaultJcaPropertyNames.KEYVAULT_URI); + } else { + System.setProperty(KeyVaultJcaPropertyNames.KEYVAULT_URI, value); } - }; + } @Test public void engineInitVerifyTest() { @@ -102,18 +88,31 @@ public void engineSetParameterTest() { @Test public void setDigestNameAndEngineSignTest() { + keyVaultClient = new MockKeyVaultClient() { + @Override + public byte[] getSignedWithPrivateKey(String digestName, String digestValue, String keyId) { + return "ES256".equals(digestName) ? signedWithES256 : null; + } + }; + KeyVaultPrivateKey keyVaultPrivateKey = new KeyVaultPrivateKey("algorithm", "kid") { + @Override + public KeyVaultClient getKeyVaultClient() { + return keyVaultClient; + } + }; keyVaultKeylessEcSignature = new KeyVaultKeylessEcSha256Signature(); - when(keyVaultClient.getSignedWithPrivateKey(ArgumentMatchers.eq("ES256"), anyString(), - ArgumentMatchers.eq(null))).thenReturn(signedWithES256); - when(keyVaultPrivateKey.getKeyVaultClient()).thenReturn(keyVaultClient); keyVaultKeylessEcSignature.engineInitSign(keyVaultPrivateKey, null); Assertions.assertArrayEquals(KeyVaultEncode.encodeByte(signedWithES256), keyVaultKeylessEcSignature.engineSign()); + keyVaultClient = new MockKeyVaultClient() { + @Override + public byte[] getSignedWithPrivateKey(String digestName, String digestValue, String keyId) { + return "ES384".equals(digestName) ? signedWithES384 : null; + } + }; keyVaultKeylessEcSignature = new KeyVaultKeylessEcSha384Signature(); keyVaultKeylessEcSignature.engineInitSign(keyVaultPrivateKey, null); - when(keyVaultClient.getSignedWithPrivateKey(ArgumentMatchers.eq("ES384"), anyString(), - ArgumentMatchers.eq(null))).thenReturn(signedWithES384); assertArrayEquals(KeyVaultEncode.encodeByte(signedWithES384), keyVaultKeylessEcSignature.engineSign()); } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsaSsaPssSignatureTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsaSsaPssSignatureTest.java index 7502bfd49590..2e549a8e5596 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsaSsaPssSignatureTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/signature/KeyVaultKeylessRsaSsaPssSignatureTest.java @@ -4,12 +4,18 @@ package com.azure.security.keyvault.jca.implementation.signature; import com.azure.security.keyvault.jca.KeyVaultJcaPropertyNames; -import com.azure.security.keyvault.jca.implementation.KeyVaultPrivateKey; import com.azure.security.keyvault.jca.implementation.KeyVaultClient; +import com.azure.security.keyvault.jca.implementation.KeyVaultPrivateKey; +import com.azure.security.keyvault.jca.implementation.MockKeyVaultClient; +import com.azure.security.keyvault.jca.implementation.mocking.MockPrivateKey; +import com.azure.security.keyvault.jca.implementation.mocking.MockPublicKey; +import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; -import org.mockito.ArgumentMatchers; +import org.junit.jupiter.api.parallel.ResourceLock; +import org.junit.jupiter.api.parallel.Resources; +import java.nio.charset.StandardCharsets; import java.security.InvalidAlgorithmParameterException; import java.security.PrivateKey; import java.security.PublicKey; @@ -19,59 +25,38 @@ import static org.junit.jupiter.api.Assertions.assertArrayEquals; import static org.junit.jupiter.api.Assertions.assertThrows; -import static org.mockito.ArgumentMatchers.anyString; -import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.when; +@ResourceLock(Resources.SYSTEM_PROPERTIES) public class KeyVaultKeylessRsaSsaPssSignatureTest { KeyVaultKeylessRsaSsaPssSignature keyVaultKeylessRsaSsaPssSignature; - static final String KEY_VAULT_TEST_URI_GLOBAL = "https://fake.vault.azure.net/"; + private final PublicKey publicKey = new MockPublicKey(); + private final PrivateKey privateKey = new MockPrivateKey(); - private final KeyVaultClient keyVaultClient = mock(KeyVaultClient.class); + static final String KEY_VAULT_TEST_URI_GLOBAL = "https://fake.vault.azure.net/"; - private final KeyVaultPrivateKey keyVaultPrivateKey = mock(KeyVaultPrivateKey.class); + private String previousKeyVaultUri; @BeforeEach public void before() { + previousKeyVaultUri = System.getProperty(KeyVaultJcaPropertyNames.KEYVAULT_URI); System.setProperty(KeyVaultJcaPropertyNames.KEYVAULT_URI, KEY_VAULT_TEST_URI_GLOBAL); keyVaultKeylessRsaSsaPssSignature = new KeyVaultKeylessRsaSsaPssSignature(); } - private final PublicKey publicKey = new PublicKey() { - @Override - public String getAlgorithm() { - return null; - } - - @Override - public String getFormat() { - return null; - } - - @Override - public byte[] getEncoded() { - return new byte[0]; - } - }; - - private final PrivateKey privateKey = new PrivateKey() { - @Override - public String getAlgorithm() { - return null; - } - - @Override - public String getFormat() { - return null; - } + @AfterEach + public void after() { + restoreKeyVaultUri(previousKeyVaultUri); + } - @Override - public byte[] getEncoded() { - return new byte[0]; + private static void restoreKeyVaultUri(String value) { + if (value == null) { + System.clearProperty(KeyVaultJcaPropertyNames.KEYVAULT_URI); + } else { + System.setProperty(KeyVaultJcaPropertyNames.KEYVAULT_URI, value); } - }; + } @Test public void engineInitVerifyTest() { @@ -104,13 +89,17 @@ public void engineSetParameterTest() { @Test public void setDigestNameAndEngineSignTest() throws InvalidAlgorithmParameterException { + KeyVaultClient keyVaultClient = new MockKeyVaultClient() { + @Override + public byte[] getSignedWithPrivateKey(String digestName, String digestValue, String keyId) { + return "PS256".equals(digestName) ? "fakeValue".getBytes(StandardCharsets.UTF_8) : null; + } + }; + KeyVaultPrivateKey keyVaultPrivateKey = new KeyVaultPrivateKey("algorithm", "kid", keyVaultClient); keyVaultKeylessRsaSsaPssSignature = new KeyVaultKeylessRsaSsaPssSignature(); - when(keyVaultPrivateKey.getKeyVaultClient()).thenReturn(keyVaultClient); keyVaultKeylessRsaSsaPssSignature.engineInitSign(keyVaultPrivateKey, null); keyVaultKeylessRsaSsaPssSignature .engineSetParameter(new PSSParameterSpec("SHA-1", "MGF1", MGF1ParameterSpec.SHA1, 20, 1)); - when(keyVaultClient.getSignedWithPrivateKey(ArgumentMatchers.eq("PS256"), anyString(), - ArgumentMatchers.eq(null))).thenReturn("fakeValue".getBytes()); assertArrayEquals("fakeValue".getBytes(), keyVaultKeylessRsaSsaPssSignature.engineSign()); } @@ -124,7 +113,8 @@ public void engineSetParameterWithNullParameterTest() { @Test public void engineSetParameterWithNotPSSParameterSpecTest() { keyVaultKeylessRsaSsaPssSignature = new KeyVaultKeylessRsaSsaPssSignature(); - AlgorithmParameterSpec algorithmParameterSpec = mock(AlgorithmParameterSpec.class); + AlgorithmParameterSpec algorithmParameterSpec = new AlgorithmParameterSpec() { + }; assertThrows(InvalidAlgorithmParameterException.class, () -> keyVaultKeylessRsaSsaPssSignature.engineSetParameter(algorithmParameterSpec)); } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/AccessTokenUtilTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/AccessTokenUtilTest.java index 71e01892e2ec..baace3c5f8f7 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/AccessTokenUtilTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/AccessTokenUtilTest.java @@ -5,14 +5,14 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.io.TempDir; -import org.mockito.MockedStatic; -import org.mockito.Mockito; import java.nio.charset.StandardCharsets; import java.nio.file.Files; import java.nio.file.Path; import java.util.ArrayList; +import java.util.Collections; import java.util.List; +import java.util.Map; import java.util.logging.Handler; import java.util.logging.Level; import java.util.logging.LogRecord; @@ -73,11 +73,9 @@ public void close() { logger.setLevel(Level.ALL); logger.setUseParentHandlers(false); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - httpMock.when(() -> HttpUtil.post(Mockito.anyString(), Mockito.anyString(), Mockito.anyString())) - .thenReturn(null); - - AccessTokenUtil.getAccessToken("https://vault.azure.net", null, "tenant-id", "client-id", clientSecret); + try { + AccessTokenUtil.getAccessToken("https://vault.azure.net", null, "tenant-id", "client-id", clientSecret, + (uri, headers, body, contentType) -> null); } finally { logger.removeHandler(collector); logger.setLevel(originalLevel); @@ -87,4 +85,13 @@ public void close() { assertFalse(loggedValues.contains(clientSecret), "The client secret must never be logged"); assertTrue(loggedValues.contains("client-id"), "Non-secret parameters stay available for diagnostics"); } + + @Test + void getLoginUriReturnsNullForEmptyAuthenticateHeader() { + Map> headers = Collections.singletonMap("WWW-Authenticate", Collections.emptyList()); + + String result = AccessTokenUtil.getLoginUri("https://vault.azure.net", false, ignored -> headers); + + assertNull(result); + } } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/AiaCertificateChainTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/AiaCertificateChainTest.java index 486d91d33d46..00c7ee3ce031 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/AiaCertificateChainTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/AiaCertificateChainTest.java @@ -6,6 +6,8 @@ import com.azure.security.keyvault.jca.KeyVaultJcaPropertyNames; import com.azure.security.keyvault.jca.implementation.CertificateVersion; import com.azure.security.keyvault.jca.implementation.KeyVaultClient; +import com.azure.security.keyvault.jca.implementation.TestCertificateVersions; +import com.azure.security.keyvault.jca.implementation.TestKeyVaultClient; import com.azure.security.keyvault.jca.implementation.model.SecretBundle; import org.bouncycastle.asn1.x500.X500Name; import org.bouncycastle.asn1.x509.AccessDescription; @@ -26,8 +28,7 @@ import org.junit.jupiter.api.Timeout; import org.junit.jupiter.api.parallel.Execution; import org.junit.jupiter.api.parallel.ExecutionMode; -import org.mockito.MockedStatic; -import org.mockito.Mockito; +import org.junit.jupiter.api.parallel.Isolated; import java.math.BigInteger; import java.nio.charset.StandardCharsets; @@ -52,10 +53,14 @@ import java.util.Base64; import java.util.Collections; import java.util.Date; +import java.util.HashMap; import java.util.List; +import java.util.Map; import java.util.Set; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.atomic.AtomicLong; +import java.util.function.Function; import java.util.logging.Handler; import java.util.logging.Level; import java.util.logging.LogRecord; @@ -68,8 +73,6 @@ import static org.junit.jupiter.api.Assertions.assertNull; import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; -import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.when; /** * Tests for AIA-based certificate chain completion in {@link CertificateUtil}. @@ -78,9 +81,10 @@ * only its leaf certificate in the secret bundle. The missing intermediate CA certificates * must be downloaded via the CA Issuers URL in the AIA extension of each certificate. * - *

Tests must run sequentially because they share JVM-global state (system properties, the AIA response cache - * and Mockito static mocks). Parallel execution would cause property-pollution flakiness. + *

Tests must run in isolation because they share JVM-global state (system properties, the AIA response cache, + * and the AIA response loader). Running another test class concurrently could make it use this class's response loader. */ +@Isolated("Mutates the global AIA response loader and cache") @Execution(ExecutionMode.SAME_THREAD) public class AiaCertificateChainTest { @@ -93,6 +97,7 @@ public class AiaCertificateChainTest { private static X509Certificate rootCert; private static X509Certificate intermediateCert; private static X509Certificate leafCert; + private TestAiaResponseLoader responseLoader; @BeforeAll static void generateTestChain() throws Exception { @@ -120,6 +125,8 @@ void setupClean() { // Ensure each test starts with a clean state - clear the disable property System.clearProperty(KeyVaultJcaPropertyNames.KEYVAULT_JCA_DISABLE_AIA_DOWNLOAD); AiaCertificateChainUtil.clearAiaCache(); + responseLoader = new TestAiaResponseLoader(); + AiaCertificateChainUtil.setResponseLoader(responseLoader); } @AfterEach @@ -127,6 +134,7 @@ void cleanup() { // Clear the property after each test to prevent interference with subsequent tests System.clearProperty(KeyVaultJcaPropertyNames.KEYVAULT_JCA_DISABLE_AIA_DOWNLOAD); AiaCertificateChainUtil.clearAiaCache(); + AiaCertificateChainUtil.resetResponseLoader(); } // ----------------------------------------------------------------------- @@ -166,17 +174,15 @@ void completeChainViaAiaLeafOnlyDownloadsIntermediateAndRoot() throws Exception // Simulate AKV returning only the leaf cert (non-exportable, leaf-only secret) Certificate[] leafOnly = new Certificate[] { leafCert }; - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); - mockAiaResponse(httpMock, AIA_ROOT_URL, rootCert.getEncoded()); + addAiaResponse(AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); + addAiaResponse(AIA_ROOT_URL, rootCert.getEncoded()); - Certificate[] completed = AiaCertificateChainUtil.completeChainViaAia(leafOnly, false); + Certificate[] completed = AiaCertificateChainUtil.completeChainViaAia(leafOnly, false); - assertEquals(3, completed.length, "Chain should contain leaf + intermediate + root"); - assertEquals(leafCert, completed[0], "First cert should be the leaf"); - assertEquals(intermediateCert, completed[1], "Second cert should be the intermediate CA"); - assertEquals(rootCert, completed[2], "Third cert should be the root CA"); - } + assertEquals(3, completed.length, "Chain should contain leaf + intermediate + root"); + assertEquals(leafCert, completed[0], "First cert should be the leaf"); + assertEquals(intermediateCert, completed[1], "Second cert should be the intermediate CA"); + assertEquals(rootCert, completed[2], "Third cert should be the root CA"); } @Test @@ -184,14 +190,12 @@ void completeChainViaAiaLeafAndIntermediateDownloadsRootOnly() throws Exception // Chain already has leaf + intermediate; only root is missing Certificate[] partial = new Certificate[] { leafCert, intermediateCert }; - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_ROOT_URL, rootCert.getEncoded()); + addAiaResponse(AIA_ROOT_URL, rootCert.getEncoded()); - Certificate[] completed = AiaCertificateChainUtil.completeChainViaAia(partial, false); + Certificate[] completed = AiaCertificateChainUtil.completeChainViaAia(partial, false); - assertEquals(3, completed.length, "Chain should contain leaf + intermediate + root"); - assertEquals(rootCert, completed[2]); - } + assertEquals(3, completed.length, "Chain should contain leaf + intermediate + root"); + assertEquals(rootCert, completed[2]); } @Test @@ -199,25 +203,21 @@ void completeChainViaAiaFullChainNoDownloadNeeded() throws Exception { // Already complete: root is self-signed, no AIA download should happen Certificate[] full = new Certificate[] { leafCert, intermediateCert, rootCert }; - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - Certificate[] result = AiaCertificateChainUtil.completeChainViaAia(full, false); + Certificate[] result = AiaCertificateChainUtil.completeChainViaAia(full, false); - assertEquals(3, result.length); - httpMock.verifyNoInteractions(); - } + assertEquals(3, result.length); + assertEquals(0, responseLoader.getTotalCallCount()); } @Test void completeChainViaAiaDownloadFailsReturnsOriginal() throws Exception { Certificate[] leafOnly = new Certificate[] { leafCert }; - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, null); + addAiaResponse(AIA_INTERMEDIATE_URL, null); - Certificate[] result = AiaCertificateChainUtil.completeChainViaAia(leafOnly, false); + Certificate[] result = AiaCertificateChainUtil.completeChainViaAia(leafOnly, false); - assertEquals(1, result.length, "Should return original chain when download fails"); - } + assertEquals(1, result.length, "Should return original chain when download fails"); } @Test @@ -237,14 +237,12 @@ void completeChainViaAiaEmptyInputReturnsEmpty() { @Test void downloadIssuerCertificateFromAiaReturnsDerEncodedCert() throws Exception { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); + addAiaResponse(AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); - X509Certificate result = AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert); + X509Certificate result = AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert); - assertNotNull(result); - assertEquals(intermediateCert, result); - } + assertNotNull(result); + assertEquals(intermediateCert, result); } @Test @@ -258,15 +256,13 @@ void downloadIssuerCertificateFromAiaNoCertWithoutAiaReturnsNull() throws Except void downloadIssuerCertificateFromAiaPemBundleSelectsMatchingIssuer() throws Exception { String pemBundle = toPem(rootCert) + toPem(intermediateCert); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, pemBundle.getBytes(StandardCharsets.UTF_8)); + addAiaResponse(AIA_INTERMEDIATE_URL, pemBundle.getBytes(StandardCharsets.UTF_8)); - X509Certificate result = AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert); + X509Certificate result = AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert); - assertNotNull(result); - assertEquals(intermediateCert, result, - "Should select the matching issuer from PEM bundle, not the first certificate"); - } + assertNotNull(result); + assertEquals(intermediateCert, result, + "Should select the matching issuer from PEM bundle, not the first certificate"); } @Test @@ -282,16 +278,14 @@ void completeChainViaAiaRejectsIssuerWithoutKeyCertSign() throws Exception { X509Certificate leafWithBadIssuerAia = buildCertificate(leafKeyPair.getPublic(), "CN=Leaf With Bad Issuer", "CN=Bad Issuer", badIssuerKeyPair.getPrivate(), false, AIA_BAD_ISSUER_URL); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_BAD_ISSUER_URL, badIssuerCert.getEncoded()); + addAiaResponse(AIA_BAD_ISSUER_URL, badIssuerCert.getEncoded()); - Certificate[] result - = AiaCertificateChainUtil.completeChainViaAia(new Certificate[] { leafWithBadIssuerAia }, false); + Certificate[] result + = AiaCertificateChainUtil.completeChainViaAia(new Certificate[] { leafWithBadIssuerAia }, false); - assertEquals(1, result.length, - "Issuer without keyCertSign should be rejected even if basicConstraints indicates CA"); - assertEquals(leafWithBadIssuerAia, result[0]); - } + assertEquals(1, result.length, + "Issuer without keyCertSign should be rejected even if basicConstraints indicates CA"); + assertEquals(leafWithBadIssuerAia, result[0]); } @Test @@ -312,16 +306,14 @@ void completeChainViaAiaRejectsExpiredIssuer() throws Exception { X509Certificate leafWithExpiredAia = buildCertificate(leafKeyPair.getPublic(), "CN=Leaf", "CN=Expired Issuer", expiredIssuerKeyPair.getPrivate(), false, AIA_BAD_ISSUER_URL); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_BAD_ISSUER_URL, expiredIssuerCert.getEncoded()); + addAiaResponse(AIA_BAD_ISSUER_URL, expiredIssuerCert.getEncoded()); - Certificate[] result - = AiaCertificateChainUtil.completeChainViaAia(new Certificate[] { leafWithExpiredAia }, false); + Certificate[] result + = AiaCertificateChainUtil.completeChainViaAia(new Certificate[] { leafWithExpiredAia }, false); - assertEquals(1, result.length, - "An expired issuer certificate must be rejected and not inserted into the chain"); - assertEquals(leafWithExpiredAia, result[0]); - } + assertEquals(1, result.length, + "An expired issuer certificate must be rejected and not inserted into the chain"); + assertEquals(leafWithExpiredAia, result[0]); } // ----------------------------------------------------------------------- @@ -389,11 +381,9 @@ void pkixPathBuildingWithFixSucceeds() throws Exception { Certificate[] leafOnly = new Certificate[] { leafCert }; Certificate[] completedChain; - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); - mockAiaResponse(httpMock, AIA_ROOT_URL, rootCert.getEncoded()); - completedChain = AiaCertificateChainUtil.completeChainViaAia(leafOnly, false); - } + addAiaResponse(AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); + addAiaResponse(AIA_ROOT_URL, rootCert.getEncoded()); + completedChain = AiaCertificateChainUtil.completeChainViaAia(leafOnly, false); assertEquals(3, completedChain.length, "Chain should be leaf + intermediate + root after fix"); @@ -420,13 +410,11 @@ void pkixPathBuildingWithFixSucceeds() throws Exception { @Test void aiaDownloadCanBeDisabled() throws Exception { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - Certificate[] result = CertificateUtil.loadCertificatesFromSecretBundleValue(toPem(leafCert), true); + Certificate[] result = CertificateUtil.loadCertificatesFromSecretBundleValue(toPem(leafCert), true); - assertEquals(1, result.length, "Chain should remain unchanged when AIA download is disabled"); - assertEquals(leafCert, result[0], "The returned certificate should be the leaf certificate"); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(Mockito.anyString()), Mockito.never()); - } + assertEquals(1, result.length, "Chain should remain unchanged when AIA download is disabled"); + assertEquals(leafCert, result[0], "The returned certificate should be the leaf certificate"); + assertEquals(0, responseLoader.getTotalCallCount()); } @Test @@ -434,31 +422,25 @@ void keyVaultClientKeepsAiaDownloadSettingFromConstruction() throws Exception { String secretId = "https://fake.vault.azure.net/secrets/aia-test/version"; SecretBundle secretBundle = new SecretBundle(); secretBundle.setValue(toPem(leafCert)); - CertificateVersion certificateVersion = mock(CertificateVersion.class); - when(certificateVersion.getAlias()).thenReturn("aia-test"); - when(certificateVersion.getSecretId()).thenReturn(secretId); - - KeyVaultClient keyVaultClient - = new KeyVaultClient("https://fake.vault.azure.net/", null, null, null, null, "test-token", false, true); + CertificateVersion certificateVersion + = TestCertificateVersions.create("aia-test", null, null, secretId, false, null); + KeyVaultClient keyVaultClient = new TestKeyVaultClient("test-token", true, (uri, headers) -> { + assertEquals(secretId + HttpUtil.API_VERSION_POSTFIX, uri); + return JsonConverterUtil.toJson(secretBundle); + }); // Simulate another SSL bundle replacing the JVM-global value before this client lazily loads its chain. System.setProperty(KeyVaultJcaPropertyNames.KEYVAULT_JCA_DISABLE_AIA_DOWNLOAD, "false"); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - httpMock - .when(() -> HttpUtil.get(secretId + HttpUtil.API_VERSION_POSTFIX, - Collections.singletonMap("Authorization", "Bearer test-token"))) - .thenReturn(JsonConverterUtil.toJson(secretBundle)); - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); - mockAiaResponse(httpMock, AIA_ROOT_URL, rootCert.getEncoded()); + addAiaResponse(AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); + addAiaResponse(AIA_ROOT_URL, rootCert.getEncoded()); - Certificate[] result = keyVaultClient.getCertificateChainForVersion(certificateVersion); + Certificate[] result = keyVaultClient.getCertificateChainForVersion(certificateVersion); - assertArrayEquals(new Certificate[] { leafCert }, result, - "The client must keep the AIA setting captured when it was constructed"); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.never()); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_ROOT_URL), Mockito.never()); - } + assertArrayEquals(new Certificate[] { leafCert }, result, + "The client must keep the AIA setting captured when it was constructed"); + assertEquals(0, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); + assertEquals(0, responseLoader.getCallCount(AIA_ROOT_URL)); } @Test @@ -466,31 +448,25 @@ void keyVaultClientKeepsAiaDownloadEnabledFromConstruction() throws Exception { String secretId = "https://fake.vault.azure.net/secrets/aia-enabled/version"; SecretBundle secretBundle = new SecretBundle(); secretBundle.setValue(toPem(leafCert)); - CertificateVersion certificateVersion = mock(CertificateVersion.class); - when(certificateVersion.getAlias()).thenReturn("aia-enabled"); - when(certificateVersion.getSecretId()).thenReturn(secretId); - - KeyVaultClient keyVaultClient - = new KeyVaultClient("https://fake.vault.azure.net/", null, null, null, null, "test-token", false, false); + CertificateVersion certificateVersion + = TestCertificateVersions.create("aia-enabled", null, null, secretId, false, null); + KeyVaultClient keyVaultClient = new TestKeyVaultClient("test-token", false, (uri, headers) -> { + assertEquals(secretId + HttpUtil.API_VERSION_POSTFIX, uri); + return JsonConverterUtil.toJson(secretBundle); + }); // Simulate another SSL bundle replacing the JVM-global value before this client lazily loads its chain. System.setProperty(KeyVaultJcaPropertyNames.KEYVAULT_JCA_DISABLE_AIA_DOWNLOAD, "true"); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - httpMock - .when(() -> HttpUtil.get(secretId + HttpUtil.API_VERSION_POSTFIX, - Collections.singletonMap("Authorization", "Bearer test-token"))) - .thenReturn(JsonConverterUtil.toJson(secretBundle)); - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); - mockAiaResponse(httpMock, AIA_ROOT_URL, rootCert.getEncoded()); + addAiaResponse(AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); + addAiaResponse(AIA_ROOT_URL, rootCert.getEncoded()); - Certificate[] result = keyVaultClient.getCertificateChainForVersion(certificateVersion); + Certificate[] result = keyVaultClient.getCertificateChainForVersion(certificateVersion); - assertArrayEquals(new Certificate[] { leafCert, intermediateCert, rootCert }, result, - "The client must keep the AIA setting captured when it was constructed"); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(1)); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_ROOT_URL), Mockito.times(1)); - } + assertArrayEquals(new Certificate[] { leafCert, intermediateCert, rootCert }, result, + "The client must keep the AIA setting captured when it was constructed"); + assertEquals(1, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); + assertEquals(1, responseLoader.getCallCount(AIA_ROOT_URL)); } // ----------------------------------------------------------------------- @@ -499,59 +475,51 @@ void keyVaultClientKeepsAiaDownloadEnabledFromConstruction() throws Exception { @Test void loadCertificatesCompletesLeafOnlyChain() throws Exception { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); - mockAiaResponse(httpMock, AIA_ROOT_URL, rootCert.getEncoded()); + addAiaResponse(AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); + addAiaResponse(AIA_ROOT_URL, rootCert.getEncoded()); - Certificate[] result = CertificateUtil.loadCertificatesFromSecretBundleValue(toPem(leafCert), false); + Certificate[] result = CertificateUtil.loadCertificatesFromSecretBundleValue(toPem(leafCert), false); - assertArrayEquals(new Certificate[] { leafCert, intermediateCert, rootCert }, result, - "A leaf-only bundle must be completed up to the root CA"); - } + assertArrayEquals(new Certificate[] { leafCert, intermediateCert, rootCert }, result, + "A leaf-only bundle must be completed up to the root CA"); } @Test void loadCertificatesCompletesChainWithMissingIntermediate() throws Exception { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); + addAiaResponse(AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); - Certificate[] result - = CertificateUtil.loadCertificatesFromSecretBundleValue(toPem(leafCert) + toPem(rootCert), false); + Certificate[] result + = CertificateUtil.loadCertificatesFromSecretBundleValue(toPem(leafCert) + toPem(rootCert), false); - assertArrayEquals(new Certificate[] { leafCert, intermediateCert, rootCert }, result, - "An intermediate missing in the middle of the chain must still be downloaded"); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_ROOT_URL), Mockito.never()); - } + assertArrayEquals(new Certificate[] { leafCert, intermediateCert, rootCert }, result, + "An intermediate missing in the middle of the chain must still be downloaded"); + assertEquals(0, responseLoader.getCallCount(AIA_ROOT_URL)); } @Test void loadCertificatesCompletesChainWithoutRootAndCachesIssuer() throws Exception { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_ROOT_URL, rootCert.getEncoded()); - - Certificate[] firstResult = CertificateUtil - .loadCertificatesFromSecretBundleValue(toPem(leafCert) + toPem(intermediateCert), false); - Certificate[] secondResult = CertificateUtil - .loadCertificatesFromSecretBundleValue(toPem(leafCert) + toPem(intermediateCert), false); - - assertArrayEquals(new Certificate[] { leafCert, intermediateCert, rootCert }, firstResult, - "A contiguous chain must still be completed when its terminal certificate is not self-signed"); - assertArrayEquals(new Certificate[] { leafCert, intermediateCert, rootCert }, secondResult, - "A subsequent load must reuse the cached root certificate"); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_ROOT_URL), Mockito.times(1)); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.never()); - } + addAiaResponse(AIA_ROOT_URL, rootCert.getEncoded()); + + Certificate[] firstResult + = CertificateUtil.loadCertificatesFromSecretBundleValue(toPem(leafCert) + toPem(intermediateCert), false); + Certificate[] secondResult + = CertificateUtil.loadCertificatesFromSecretBundleValue(toPem(leafCert) + toPem(intermediateCert), false); + + assertArrayEquals(new Certificate[] { leafCert, intermediateCert, rootCert }, firstResult, + "A contiguous chain must still be completed when its terminal certificate is not self-signed"); + assertArrayEquals(new Certificate[] { leafCert, intermediateCert, rootCert }, secondResult, + "A subsequent load must reuse the cached root certificate"); + assertEquals(1, responseLoader.getCallCount(AIA_ROOT_URL)); + assertEquals(0, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test void loadCertificatesSkipsAiaForCompleteChain() throws Exception { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - Certificate[] result = CertificateUtil.loadCertificatesFromSecretBundleValue( - toPem(leafCert) + toPem(intermediateCert) + toPem(rootCert), false); + Certificate[] result = CertificateUtil + .loadCertificatesFromSecretBundleValue(toPem(leafCert) + toPem(intermediateCert) + toPem(rootCert), false); - assertArrayEquals(new Certificate[] { leafCert, intermediateCert, rootCert }, result); - httpMock.verifyNoInteractions(); - } + assertArrayEquals(new Certificate[] { leafCert, intermediateCert, rootCert }, result); + assertEquals(0, responseLoader.getTotalCallCount()); } @Test @@ -571,14 +539,12 @@ void loadCertificatesKeepsChainWithExpiredIssuerUntouched() throws Exception { X509Certificate leafOfExpiredCa = buildCertificate(leafKeyPair.getPublic(), "CN=Leaf Of Expired CA", "CN=Expired CA", expiredCaKeyPair.getPrivate(), false, AIA_INTERMEDIATE_URL); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - Certificate[] result = CertificateUtil - .loadCertificatesFromSecretBundleValue(toPem(leafOfExpiredCa) + toPem(expiredCaCert), false); + Certificate[] result = CertificateUtil + .loadCertificatesFromSecretBundleValue(toPem(leafOfExpiredCa) + toPem(expiredCaCert), false); - assertArrayEquals(new Certificate[] { leafOfExpiredCa, expiredCaCert }, result, - "An expired certificate already in the chain must not change how the chain is ordered"); - httpMock.verifyNoInteractions(); - } + assertArrayEquals(new Certificate[] { leafOfExpiredCa, expiredCaCert }, result, + "An expired certificate already in the chain must not change how the chain is ordered"); + assertEquals(0, responseLoader.getTotalCallCount()); } // ----------------------------------------------------------------------- @@ -591,14 +557,12 @@ void loadCertificatesKeepsChainWithExpiredIssuerUntouched() throws Exception { @Test void aiaResponseIsCachedAcrossDownloads() throws Exception { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); + addAiaResponse(AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); - assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); - assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); + assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); + assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(1)); - } + assertEquals(1, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test @@ -613,15 +577,13 @@ void cachedAiaResponseIsValidatedBeforeAndAfterForcedRefresh() throws Exception X509Certificate certSignedByAnotherKey = buildCertificate(subjectKeyPair.getPublic(), "CN=Other Leaf", "CN=Test Intermediate CA", impostorKeyPair.getPrivate(), false, AIA_INTERMEDIATE_URL); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); + addAiaResponse(AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); - assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); - assertNull(AiaCertificateChainUtil.downloadIssuerCertificateFromAia(certSignedByAnotherKey), - "A cache hit and its forced refresh must both reject a signature mismatch"); + assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); + assertNull(AiaCertificateChainUtil.downloadIssuerCertificateFromAia(certSignedByAnotherKey), + "A cache hit and its forced refresh must both reject a signature mismatch"); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(2)); - } + assertEquals(2, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test @@ -637,14 +599,12 @@ void expiredExtraCertificateDoesNotPreventCachingValidIssuer() throws Exception expiredExtraKeyPair.getPrivate(), true, null, KeyUsage.keyCertSign, expiredNotBefore, expiredNotAfter); String pemBundle = toPem(expiredExtra) + toPem(intermediateCert); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, pemBundle.getBytes(StandardCharsets.UTF_8)); + addAiaResponse(AIA_INTERMEDIATE_URL, pemBundle.getBytes(StandardCharsets.UTF_8)); - assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); - assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); + assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); + assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(1)); - } + assertEquals(1, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test @@ -659,15 +619,13 @@ void refreshesCachedResponseWhenIssuerRotates() throws Exception { X509Certificate rotatedLeaf = buildCertificate(rotatedLeafKeyPair.getPublic(), "CN=Rotated Leaf", "CN=Test Intermediate CA", rotatedIssuerKeyPair.getPrivate(), false, AIA_INTERMEDIATE_URL); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - httpMock.when(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL)) - .thenReturn(binaryResponse(intermediateCert.getEncoded()), binaryResponse(rotatedIssuer.getEncoded())); + responseLoader.addResponses(AIA_INTERMEDIATE_URL, binaryResponse(intermediateCert.getEncoded()), + binaryResponse(rotatedIssuer.getEncoded())); - assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); - assertEquals(rotatedIssuer, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(rotatedLeaf)); + assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); + assertEquals(rotatedIssuer, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(rotatedLeaf)); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(2)); - } + assertEquals(2, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test @@ -696,16 +654,13 @@ public void close() { logger.setUseParentHandlers(false); try { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - httpMock.when(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL)) - .thenReturn(binaryResponse(intermediateCert.getEncoded())); + addAiaResponse(AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); - assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); - assertNull(AiaCertificateChainUtil.downloadIssuerCertificateFromAia(rotatedLeaf)); - assertNull(AiaCertificateChainUtil.downloadIssuerCertificateFromAia(rotatedLeaf)); + assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); + assertNull(AiaCertificateChainUtil.downloadIssuerCertificateFromAia(rotatedLeaf)); + assertNull(AiaCertificateChainUtil.downloadIssuerCertificateFromAia(rotatedLeaf)); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(2)); - } + assertEquals(2, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } finally { logger.removeHandler(collector); logger.setLevel(originalLevel); @@ -731,123 +686,102 @@ void doesNotSuppressDifferentTarget() throws Exception { X509Certificate secondRotatedLeaf = buildCertificate(secondLeafKeyPair.getPublic(), "CN=Second Rotated Leaf", "CN=Test Intermediate CA", secondIssuerKeyPair.getPrivate(), false, AIA_INTERMEDIATE_URL); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - httpMock.when(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL)) - .thenReturn(binaryResponse(intermediateCert.getEncoded()), - binaryResponse(intermediateCert.getEncoded()), binaryResponse(secondIssuer.getEncoded())); + responseLoader.addResponses(AIA_INTERMEDIATE_URL, binaryResponse(intermediateCert.getEncoded()), + binaryResponse(intermediateCert.getEncoded()), binaryResponse(secondIssuer.getEncoded())); - assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); - assertNull(AiaCertificateChainUtil.downloadIssuerCertificateFromAia(firstRotatedLeaf)); - assertEquals(secondIssuer, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(secondRotatedLeaf)); + assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); + assertNull(AiaCertificateChainUtil.downloadIssuerCertificateFromAia(firstRotatedLeaf)); + assertEquals(secondIssuer, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(secondRotatedLeaf)); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(3)); - } + assertEquals(3, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test void failedRefreshDoesNotReplaceUsefulPositiveEntry() throws Exception { X509Certificate rotatedLeaf = buildRotatedLeafWithoutMatchingIssuer("CN=Failed Refresh Leaf"); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - httpMock.when(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL)) - .thenReturn(binaryResponse(intermediateCert.getEncoded()), binaryResponse(null)); + responseLoader.addResponses(AIA_INTERMEDIATE_URL, binaryResponse(intermediateCert.getEncoded()), + binaryResponse(null)); - assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); - assertNull(AiaCertificateChainUtil.downloadIssuerCertificateFromAia(rotatedLeaf)); - assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); + assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); + assertNull(AiaCertificateChainUtil.downloadIssuerCertificateFromAia(rotatedLeaf)); + assertEquals(intermediateCert, AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert)); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(2)); - } + assertEquals(2, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test void clearAiaCacheForcesNewDownload() throws Exception { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); + addAiaResponse(AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); - AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert); - AiaCertificateChainUtil.clearAiaCache(); - AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert); + AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert); + AiaCertificateChainUtil.clearAiaCache(); + AiaCertificateChainUtil.downloadIssuerCertificateFromAia(leafCert); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(2)); - } + assertEquals(2, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test void failedAiaResponseIsNegativelyCached() { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, null); + addAiaResponse(AIA_INTERMEDIATE_URL, null); - assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); - assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); + assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); + assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(1)); - } + assertEquals(1, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test void emptyAiaResponseIsNegativelyCached() { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, new byte[0]); + addAiaResponse(AIA_INTERMEDIATE_URL, new byte[0]); - assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); - assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); + assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); + assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(1)); - } + assertEquals(1, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test void unparseableAiaResponseIsNegativelyCached() { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, "not a certificate".getBytes(StandardCharsets.UTF_8)); + addAiaResponse(AIA_INTERMEDIATE_URL, "not a certificate".getBytes(StandardCharsets.UTF_8)); - assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); - assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); + assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); + assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(1)); - } + assertEquals(1, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test void noStoreAiaResponseIsNotCached() throws Exception { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - httpMock.when(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL)) - .thenReturn(binaryResponse(intermediateCert.getEncoded(), "no-store")); + responseLoader.addResponses(AIA_INTERMEDIATE_URL, binaryResponse(intermediateCert.getEncoded(), "no-store")); - assertEquals(intermediateCert, - AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).get(0)); - assertEquals(intermediateCert, - AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).get(0)); + assertEquals(intermediateCert, + AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).get(0)); + assertEquals(intermediateCert, + AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).get(0)); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(2)); - } + assertEquals(2, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test void noStoreFailedAiaResponseIsNotCached() { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - httpMock.when(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL)) - .thenReturn(binaryResponse(null, "no-store")); + responseLoader.addResponses(AIA_INTERMEDIATE_URL, binaryResponse(null, "no-store")); - assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); - assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); + assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); + assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(2)); - } + assertEquals(2, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test void noCacheUnparseableAiaResponseIsNotCached() { - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - httpMock.when(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL)) - .thenReturn(binaryResponse("not a certificate".getBytes(StandardCharsets.UTF_8), "no-cache")); + responseLoader.addResponses(AIA_INTERMEDIATE_URL, + binaryResponse("not a certificate".getBytes(StandardCharsets.UTF_8), "no-cache")); - assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); - assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); + assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); + assertTrue(AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(AIA_INTERMEDIATE_URL).isEmpty()); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(2)); - } + assertEquals(2, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); } @Test @@ -973,22 +907,19 @@ void responseExpirationIsIndependentOfCandidateValidity() { @Test void aiaCacheEvictsLeastRecentlyUsedEntryWhenFull() throws Exception { String firstUrl = "http://aia.example.com/cache-0.crt"; + HttpUtil.BinaryHttpResponse response = binaryResponse(intermediateCert.getEncoded()); + responseLoader.respondToAnyUrl(ignored -> response); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - httpMock.when(() -> HttpUtil.getBytesWithMetadata(Mockito.anyString())) - .thenReturn(binaryResponse(intermediateCert.getEncoded())); + AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(firstUrl); - AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(firstUrl); - - // Fill the cache past its maximum size so its first entry can no longer be retained. - for (int i = 1; i <= 128; i++) { - AiaCertificateChainUtil.fetchCertificatesFromAiaUrl("http://aia.example.com/cache-" + i + ".crt"); - } + // Fill the cache past its maximum size so its first entry can no longer be retained. + for (int i = 1; i <= 128; i++) { + AiaCertificateChainUtil.fetchCertificatesFromAiaUrl("http://aia.example.com/cache-" + i + ".crt"); + } - AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(firstUrl); + AiaCertificateChainUtil.fetchCertificatesFromAiaUrl(firstUrl); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(firstUrl), Mockito.times(2)); - } + assertEquals(2, responseLoader.getCallCount(firstUrl)); } @Test @@ -1017,16 +948,16 @@ public void close() { logger.setLevel(Level.FINE); logger.setUseParentHandlers(false); - try (MockedStatic httpMock = Mockito.mockStatic(HttpUtil.class)) { - mockAiaResponse(httpMock, AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); - mockAiaResponse(httpMock, AIA_ROOT_URL, rootCert.getEncoded()); + try { + addAiaResponse(AIA_INTERMEDIATE_URL, intermediateCert.getEncoded()); + addAiaResponse(AIA_ROOT_URL, rootCert.getEncoded()); // The second run resolves the same two issuers entirely from the cache. AiaCertificateChainUtil.completeChainViaAia(new Certificate[] { leafCert }, false); AiaCertificateChainUtil.completeChainViaAia(new Certificate[] { leafCert }, false); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_INTERMEDIATE_URL), Mockito.times(1)); - httpMock.verify(() -> HttpUtil.getBytesWithMetadata(AIA_ROOT_URL), Mockito.times(1)); + assertEquals(1, responseLoader.getCallCount(AIA_INTERMEDIATE_URL)); + assertEquals(1, responseLoader.getCallCount(AIA_ROOT_URL)); } finally { logger.removeHandler(collector); logger.setLevel(originalLevel); @@ -1072,8 +1003,8 @@ void completeChainViaAiaTerminatesOnCrossSignedIssuers() throws Exception { // Helper // ----------------------------------------------------------------------- - private static void mockAiaResponse(MockedStatic httpMock, String url, byte[] body) { - httpMock.when(() -> HttpUtil.getBytesWithMetadata(url)).thenReturn(binaryResponse(body)); + private void addAiaResponse(String url, byte[] body) { + responseLoader.addResponses(url, binaryResponse(body)); } private static HttpUtil.BinaryHttpResponse binaryResponse(byte[] body) { @@ -1139,4 +1070,40 @@ private static String toPem(X509Certificate certificate) throws Exception { String base64 = Base64.getMimeEncoder(64, new byte[] { '\n' }).encodeToString(certificate.getEncoded()); return "-----BEGIN CERTIFICATE-----\n" + base64 + "\n-----END CERTIFICATE-----\n"; } + + private static final class TestAiaResponseLoader implements AiaCertificateChainUtil.AiaResponseLoader { + private final Map> responses = new HashMap<>(); + private final Map callCounts = new HashMap<>(); + private Function fallback; + + private void addResponses(String url, HttpUtil.BinaryHttpResponse... configuredResponses) { + responses.put(url, new ArrayList<>(Arrays.asList(configuredResponses))); + } + + private void respondToAnyUrl(Function responder) { + fallback = responder; + } + + private int getCallCount(String url) { + AtomicInteger count = callCounts.get(url); + return count == null ? 0 : count.get(); + } + + private int getTotalCallCount() { + return callCounts.values().stream().mapToInt(AtomicInteger::get).sum(); + } + + @Override + public HttpUtil.BinaryHttpResponse load(String url) { + int invocation = callCounts.computeIfAbsent(url, ignored -> new AtomicInteger()).getAndIncrement(); + List configuredResponses = responses.get(url); + if (configuredResponses != null && !configuredResponses.isEmpty()) { + return configuredResponses.get(Math.min(invocation, configuredResponses.size() - 1)); + } + if (fallback != null) { + return fallback.apply(url); + } + throw new AssertionError("Unexpected AIA request: " + url); + } + } } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/AiaResponseCacheTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/AiaResponseCacheTest.java index 65f966438490..35d28e6cc7ef 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/AiaResponseCacheTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/AiaResponseCacheTest.java @@ -6,6 +6,10 @@ import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.Test; +import java.io.IOException; +import java.io.InputStream; +import java.security.cert.CertificateException; +import java.security.cert.CertificateFactory; import java.security.cert.X509Certificate; import java.util.ArrayList; import java.util.Collections; @@ -24,9 +28,9 @@ import static org.junit.jupiter.api.Assertions.assertSame; import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; -import static org.mockito.Mockito.mock; public class AiaResponseCacheTest { + private static final X509Certificate TEST_CERTIFICATE = loadTestCertificate(); private final AtomicLong clock = new AtomicLong(1_000L); private final ExecutorService executor = Executors.newFixedThreadPool(16); @@ -40,7 +44,7 @@ void shutdownExecutor() throws InterruptedException { void reusesSuccessfulResolutionBeforeExpiry() { AiaResponseCache cache = new AiaResponseCache(128, clock::get); AtomicInteger loads = new AtomicInteger(); - List certificates = Collections.singletonList(mock(X509Certificate.class)); + List certificates = Collections.singletonList(TEST_CERTIFICATE); assertSame(certificates, cache.getOrLoad("url", () -> entry(certificates, loads))); assertSame(certificates, cache.getOrLoad("url", () -> entry(certificates, loads))); @@ -63,7 +67,7 @@ void reusesNegativeResolutionBeforeExpiry() { void reloadsResolutionAfterExpiry() { AiaResponseCache cache = new AiaResponseCache(128, clock::get); AtomicInteger loads = new AtomicInteger(); - List certificates = Collections.singletonList(mock(X509Certificate.class)); + List certificates = Collections.singletonList(TEST_CERTIFICATE); cache.getOrLoad("url", () -> entry(certificates, loads)); clock.set(2_001L); @@ -75,7 +79,7 @@ void reloadsResolutionAfterExpiry() { @Test void lookupResultReportsSourceAndGeneration() { AiaResponseCache cache = new AiaResponseCache(128, clock::get); - List certificates = Collections.singletonList(mock(X509Certificate.class)); + List certificates = Collections.singletonList(TEST_CERTIFICATE); AiaResponseCache.LookupResult loaded = cache.getOrLoadResult("url", () -> new AiaResponseCache.Entry(certificates, 2_000L), () -> { @@ -95,8 +99,8 @@ void lookupResultReportsSourceAndGeneration() { void refreshIfUnchangedSkipsLoaderWhenEntryChanged() { AiaResponseCache cache = new AiaResponseCache(128, clock::get); AtomicInteger refreshLoads = new AtomicInteger(); - List firstCertificates = Collections.singletonList(mock(X509Certificate.class)); - List secondCertificates = Collections.singletonList(mock(X509Certificate.class)); + List firstCertificates = Collections.singletonList(TEST_CERTIFICATE); + List secondCertificates = Collections.singletonList(TEST_CERTIFICATE); AiaResponseCache.LookupResult first = cache.getOrLoadResult("url", () -> new AiaResponseCache.Entry(firstCertificates, 2_000L), () -> { @@ -119,8 +123,8 @@ void refreshIfUnchangedSkipsLoaderWhenEntryChanged() { @Test void coalescesConcurrentForcedRefreshes() throws Exception { AiaResponseCache cache = new AiaResponseCache(128, clock::get); - List oldCertificates = Collections.singletonList(mock(X509Certificate.class)); - List newCertificates = Collections.singletonList(mock(X509Certificate.class)); + List oldCertificates = Collections.singletonList(TEST_CERTIFICATE); + List newCertificates = Collections.singletonList(TEST_CERTIFICATE); AiaResponseCache.LookupResult initial = cache.getOrLoadResult("url", () -> new AiaResponseCache.Entry(oldCertificates, 2_000L), () -> { }); @@ -151,9 +155,9 @@ void coalescesConcurrentForcedRefreshes() throws Exception { void lateRefreshDoesNotOverwriteNewerNormalLoad() throws Exception { List messages = Collections.synchronizedList(new ArrayList<>()); AiaResponseCache cache = new AiaResponseCache(128, clock::get, (message, parameters) -> messages.add(message)); - List initialCertificates = Collections.singletonList(mock(X509Certificate.class)); - List refreshedCertificates = Collections.singletonList(mock(X509Certificate.class)); - List loadedCertificates = Collections.singletonList(mock(X509Certificate.class)); + List initialCertificates = Collections.singletonList(TEST_CERTIFICATE); + List refreshedCertificates = Collections.singletonList(TEST_CERTIFICATE); + List loadedCertificates = Collections.singletonList(TEST_CERTIFICATE); AiaResponseCache.LookupResult initial = cache.getOrLoadResult("url", () -> new AiaResponseCache.Entry(initialCertificates, 1_500L), () -> { }); @@ -189,10 +193,10 @@ void lateRefreshDoesNotOverwriteNewerNormalLoad() throws Exception { @Test void differentGenerationsDoNotShareForcedRefresh() throws Exception { AiaResponseCache cache = new AiaResponseCache(128, clock::get); - List initialCertificates = Collections.singletonList(mock(X509Certificate.class)); - List firstRefreshCertificates = Collections.singletonList(mock(X509Certificate.class)); - List loadedCertificates = Collections.singletonList(mock(X509Certificate.class)); - List secondRefreshCertificates = Collections.singletonList(mock(X509Certificate.class)); + List initialCertificates = Collections.singletonList(TEST_CERTIFICATE); + List firstRefreshCertificates = Collections.singletonList(TEST_CERTIFICATE); + List loadedCertificates = Collections.singletonList(TEST_CERTIFICATE); + List secondRefreshCertificates = Collections.singletonList(TEST_CERTIFICATE); AiaResponseCache.LookupResult initial = cache.getOrLoadResult("url", () -> new AiaResponseCache.Entry(initialCertificates, 1_500L), () -> { }); @@ -243,7 +247,7 @@ void differentGenerationsDoNotShareForcedRefresh() throws Exception { @Test void negativeForcedRefreshKeepsPositiveEntry() { AiaResponseCache cache = new AiaResponseCache(128, clock::get); - List positiveCertificates = Collections.singletonList(mock(X509Certificate.class)); + List positiveCertificates = Collections.singletonList(TEST_CERTIFICATE); AiaResponseCache.LookupResult initial = cache.getOrLoadResult("url", () -> new AiaResponseCache.Entry(positiveCertificates, 2_000L), () -> { }); @@ -293,7 +297,7 @@ void targetSuppressionExpires() { void reportsCacheAndSuppressionLifecycle() { List messages = new ArrayList<>(); AiaResponseCache cache = new AiaResponseCache(1, clock::get, (message, parameters) -> messages.add(message)); - List certificates = Collections.singletonList(mock(X509Certificate.class)); + List certificates = Collections.singletonList(TEST_CERTIFICATE); AiaResponseCache.LookupResult first = cache.getOrLoadResult("url-1", () -> new AiaResponseCache.Entry(certificates, 2_000L), () -> { @@ -327,7 +331,7 @@ void coalescesConcurrentMissesForSameUrl() throws Exception { AtomicInteger loads = new AtomicInteger(); CountDownLatch loaderStarted = new CountDownLatch(1); CountDownLatch releaseLoader = new CountDownLatch(1); - List certificates = Collections.singletonList(mock(X509Certificate.class)); + List certificates = Collections.singletonList(TEST_CERTIFICATE); List>> futures = new ArrayList<>(); for (int i = 0; i < 16; i++) { @@ -352,7 +356,7 @@ void allowsDifferentUrlsToLoadConcurrently() throws Exception { AiaResponseCache cache = new AiaResponseCache(128, clock::get); CountDownLatch bothLoadersStarted = new CountDownLatch(2); CountDownLatch releaseLoaders = new CountDownLatch(1); - List certificates = Collections.singletonList(mock(X509Certificate.class)); + List certificates = Collections.singletonList(TEST_CERTIFICATE); Future> first = executor.submit(() -> cache.getOrLoad("url-1", () -> { bothLoadersStarted.countDown(); @@ -375,7 +379,7 @@ void allowsDifferentUrlsToLoadConcurrently() throws Exception { void loaderFailureDoesNotLeaveInFlightEntry() { AiaResponseCache cache = new AiaResponseCache(128, clock::get); AtomicInteger loads = new AtomicInteger(); - List certificates = Collections.singletonList(mock(X509Certificate.class)); + List certificates = Collections.singletonList(TEST_CERTIFICATE); CompletionException exception = assertThrows(CompletionException.class, () -> cache.getOrLoad("url", () -> { loads.incrementAndGet(); @@ -391,7 +395,7 @@ void loaderFailureDoesNotLeaveInFlightEntry() { void loaderErrorPropagatesWithoutLeavingInFlightEntry() { AiaResponseCache cache = new AiaResponseCache(128, clock::get); AtomicInteger loads = new AtomicInteger(); - List certificates = Collections.singletonList(mock(X509Certificate.class)); + List certificates = Collections.singletonList(TEST_CERTIFICATE); AssertionError error = assertThrows(AssertionError.class, () -> cache.getOrLoad("url", () -> { loads.incrementAndGet(); @@ -407,7 +411,7 @@ void loaderErrorPropagatesWithoutLeavingInFlightEntry() { void evictsOnlyLeastRecentlyUsedEntry() { AiaResponseCache cache = new AiaResponseCache(2, clock::get); AtomicInteger loads = new AtomicInteger(); - List certificates = Collections.singletonList(mock(X509Certificate.class)); + List certificates = Collections.singletonList(TEST_CERTIFICATE); cache.getOrLoad("url-1", () -> entry(certificates, loads)); cache.getOrLoad("url-2", () -> entry(certificates, loads)); @@ -424,7 +428,7 @@ void evictsOnlyLeastRecentlyUsedEntry() { void removesExpiredEntriesBeforeLruEviction() { AiaResponseCache cache = new AiaResponseCache(2, clock::get); AtomicInteger loads = new AtomicInteger(); - List certificates = Collections.singletonList(mock(X509Certificate.class)); + List certificates = Collections.singletonList(TEST_CERTIFICATE); cache.getOrLoad("expired", () -> { loads.incrementAndGet(); @@ -453,7 +457,7 @@ void removesExpiredEntriesBeforeLruEviction() { void clearRemovesCachedEntries() { AiaResponseCache cache = new AiaResponseCache(128, clock::get); AtomicInteger loads = new AtomicInteger(); - List certificates = Collections.singletonList(mock(X509Certificate.class)); + List certificates = Collections.singletonList(TEST_CERTIFICATE); cache.getOrLoad("url", () -> entry(certificates, loads)); cache.clear(); @@ -468,7 +472,7 @@ void clearDuringLoadDoesNotRepopulateCache() throws Exception { AtomicInteger loads = new AtomicInteger(); CountDownLatch loadStarted = new CountDownLatch(1); CountDownLatch releaseLoad = new CountDownLatch(1); - List certificates = Collections.singletonList(mock(X509Certificate.class)); + List certificates = Collections.singletonList(TEST_CERTIFICATE); Future> first = executor.submit(() -> cache.getOrLoad("url", () -> { loads.incrementAndGet(); @@ -491,6 +495,17 @@ private AiaResponseCache.Entry entry(List certificates, AtomicI return new AiaResponseCache.Entry(certificates, 2_000L); } + private static X509Certificate loadTestCertificate() { + try (InputStream inputStream = AiaResponseCacheTest.class.getResourceAsStream("/well-known/sideload.pem")) { + if (inputStream == null) { + throw new IllegalStateException("Test certificate resource was not found."); + } + return (X509Certificate) CertificateFactory.getInstance("X.509").generateCertificate(inputStream); + } catch (IOException | CertificateException e) { + throw new IllegalStateException("Failed to load the test certificate.", e); + } + } + private static void await(CountDownLatch latch) { try { if (!latch.await(5, TimeUnit.SECONDS)) { diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/HttpUtilTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/HttpUtilTest.java index ffdd9f7ec42b..61ae85fdbc1b 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/HttpUtilTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/HttpUtilTest.java @@ -5,13 +5,30 @@ import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; -import org.apache.hc.core5.http.ContentType; -import org.apache.hc.core5.http.io.entity.ByteArrayEntity; -import org.apache.hc.core5.http.message.BasicClassicHttpResponse; + +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.io.InputStream; +import java.io.OutputStream; +import java.net.HttpURLConnection; +import java.net.URL; +import java.nio.charset.StandardCharsets; +import java.util.Arrays; +import java.util.Collections; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; import static com.azure.security.keyvault.jca.implementation.utils.HttpUtil.DEFAULT_USER_AGENT_VALUE_PREFIX; import static com.azure.security.keyvault.jca.implementation.utils.HttpUtil.VERSION; -import static org.junit.jupiter.api.Assertions.*; +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; public class HttpUtilTest { @@ -39,31 +56,124 @@ public void testHttpUtilGet1() { assertFalse(result.isEmpty()); } + @Test + void textGetReturnsSuccessfulResponseBody() throws Exception { + byte[] body = "response".getBytes(StandardCharsets.UTF_8); + TestHttpURLConnection connection + = new TestHttpURLConnection("https://example.test/value", 200, body, Collections.emptyMap()); + Map requestHeaders = new LinkedHashMap<>(); + requestHeaders.put("x-test", "value"); + requestHeaders.put(HttpUtil.USER_AGENT_KEY, "caller-value"); + + String result = HttpUtil.get("https://example.test/value", requestHeaders, ignored -> connection); + + assertEquals("response", result); + assertEquals("value", connection.getRequestProperty("x-test")); + assertEquals(HttpUtil.USER_AGENT_VALUE, connection.getRequestProperty(HttpUtil.USER_AGENT_KEY)); + assertEquals(HttpUtil.HTTP_TIMEOUT_IN_MILLISECONDS, connection.getConnectTimeout()); + assertEquals(HttpUtil.HTTP_TIMEOUT_IN_MILLISECONDS, connection.getReadTimeout()); + assertTrue(connection.disconnected); + } + + @Test + void textGetThrowsForNonSuccessfulResponse() throws Exception { + TestHttpURLConnection connection = new TestHttpURLConnection("https://example.test/value", 429, + "{\"error\":\"throttled\"}".getBytes(StandardCharsets.UTF_8), Collections.emptyMap()); + + RuntimeException exception = assertThrows(RuntimeException.class, + () -> HttpUtil.get("https://example.test/value", null, ignored -> connection)); + + assertTrue(exception.getMessage().contains("HTTP status code was 429")); + assertTrue(connection.disconnected); + } + + @Test + void postThrowsForNonSuccessfulResponse() throws Exception { + TestHttpURLConnection connection = new TestHttpURLConnection("https://example.test/value", 403, + "{\"error\":\"forbidden\"}".getBytes(StandardCharsets.UTF_8), Collections.emptyMap()); + + RuntimeException exception = assertThrows(RuntimeException.class, () -> HttpUtil + .post("https://example.test/value", null, "request", "application/json", ignored -> connection)); + + assertTrue(exception.getMessage().contains("403")); + assertEquals("request", new String(connection.requestBody.toByteArray(), StandardCharsets.UTF_8)); + assertEquals(HttpUtil.HTTP_TIMEOUT_IN_MILLISECONDS, connection.getConnectTimeout()); + assertEquals(HttpUtil.HTTP_TIMEOUT_IN_MILLISECONDS, connection.getReadTimeout()); + assertTrue(connection.disconnected); + } + + @Test + void authenticationChallengeReturnsCaseInsensitiveHeadersOnlyFor401() throws Exception { + Map> headers + = Collections.singletonMap("www-authenticate", Collections.singletonList("Bearer authorization=test")); + TestHttpURLConnection unauthorized + = new TestHttpURLConnection("https://example.test/challenge", 401, new byte[0], headers); + + Map> result + = HttpUtil.getWithOnlyResponseHeaders("https://example.test/challenge", ignored -> unauthorized); + + assertEquals("Bearer authorization=test", result.get("WWW-Authenticate").get(0)); + assertEquals(HttpUtil.HTTP_TIMEOUT_IN_MILLISECONDS, unauthorized.getConnectTimeout()); + assertEquals(HttpUtil.HTTP_TIMEOUT_IN_MILLISECONDS, unauthorized.getReadTimeout()); + assertTrue(unauthorized.disconnected); + + TestHttpURLConnection successful + = new TestHttpURLConnection("https://example.test/challenge", 200, new byte[0], headers); + assertNull(HttpUtil.getWithOnlyResponseHeaders("https://example.test/challenge", ignored -> successful)); + assertTrue(successful.disconnected); + } + @Test void binaryResponsePreservesBodyAndFreshnessHeaders() throws Exception { - BasicClassicHttpResponse response = new BasicClassicHttpResponse(200); byte[] body = new byte[] { 1, 2, 3 }; - response.setEntity(new ByteArrayEntity(body, ContentType.APPLICATION_OCTET_STREAM)); - response.addHeader("Cache-Control", "public, max-age=300"); - response.addHeader("Date", "Wed, 05 Aug 2026 10:00:00 GMT"); - response.addHeader("Age", "30"); - response.addHeader("Expires", "Wed, 05 Aug 2026 10:05:00 GMT"); + Map> headers = new LinkedHashMap<>(); + headers.put("Cache-Control", Collections.singletonList("public, max-age=300")); + headers.put("Date", Collections.singletonList("Wed, 05 Aug 2026 10:00:00 GMT")); + headers.put("Age", Collections.singletonList("30")); + headers.put("Expires", Collections.singletonList("Wed, 05 Aug 2026 10:05:00 GMT")); + TestHttpURLConnection connection = new TestHttpURLConnection(200, body, headers); - HttpUtil.BinaryHttpResponse result = HttpUtil.toBinaryResponse(response, "https://example.test/cert.crt"); + HttpUtil.BinaryHttpResponse result + = HttpUtil.getAiaBytesWithMetadata("https://example.test/cert.crt", ignored -> connection); assertArrayEquals(body, result.getBody()); assertEquals("public, max-age=300", result.getCacheControl()); assertEquals("Wed, 05 Aug 2026 10:00:00 GMT", result.getDate()); assertEquals("30", result.getAge()); assertEquals("Wed, 05 Aug 2026 10:05:00 GMT", result.getExpires()); + assertEquals(10_000, connection.getConnectTimeout()); + assertEquals(10_000, connection.getReadTimeout()); + assertEquals(HttpUtil.USER_AGENT_VALUE, connection.getRequestProperty(HttpUtil.USER_AGENT_KEY)); + assertTrue(connection.disconnected); } @Test - void binaryResponseForFailureHasNoBodyOrFreshnessMetadata() throws Exception { - BasicClassicHttpResponse response = new BasicClassicHttpResponse(503); - response.addHeader("Cache-Control", "max-age=3600"); + void connectionOpeningFailuresReturnNull() { + String unsupportedUrl = "unsupported://example.test"; - HttpUtil.BinaryHttpResponse result = HttpUtil.toBinaryResponse(response, "https://example.test/cert.crt"); + assertNull(HttpUtil.get(unsupportedUrl, null)); + assertNull(HttpUtil.post(unsupportedUrl, null, "request", "application/json")); + assertNull(HttpUtil.getWithOnlyResponseHeaders(unsupportedUrl)); + } + + @Test + void binaryResponseHandlesConnectionOpeningFailure() { + HttpUtil.BinaryHttpResponse result + = HttpUtil.getAiaBytesWithMetadata("https://example.test/cert.crt", ignored -> { + throw new IOException("Connection failed"); + }); + + assertNull(result.getBody()); + } + + @Test + void binaryResponseForFailureHasNoBodyAndPreservesFreshnessMetadata() throws Exception { + Map> headers + = Collections.singletonMap("Cache-Control", Collections.singletonList("max-age=3600")); + TestHttpURLConnection connection = new TestHttpURLConnection(503, new byte[] { 1 }, headers); + + HttpUtil.BinaryHttpResponse result + = HttpUtil.getAiaBytesWithMetadata("https://example.test/cert.crt", ignored -> connection); assertNull(result.getBody()); assertEquals("max-age=3600", result.getCacheControl()); @@ -74,13 +184,153 @@ void binaryResponseForFailureHasNoBodyOrFreshnessMetadata() throws Exception { @Test void binaryResponseCombinesMultipleCacheControlHeaders() throws Exception { - BasicClassicHttpResponse response = new BasicClassicHttpResponse(200); - response.setEntity(new ByteArrayEntity(new byte[] { 1 }, ContentType.APPLICATION_OCTET_STREAM)); - response.addHeader("Cache-Control", "public, max-age=300"); - response.addHeader("Cache-Control", "no-store"); + Map> headers = new LinkedHashMap<>(); + headers.put("cache-control", Arrays.asList("public, max-age=300", "no-store")); + TestHttpURLConnection connection = new TestHttpURLConnection(200, new byte[] { 1 }, headers); - HttpUtil.BinaryHttpResponse result = HttpUtil.toBinaryResponse(response, "https://example.test/cert.crt"); + HttpUtil.BinaryHttpResponse result + = HttpUtil.getAiaBytesWithMetadata("https://example.test/cert.crt", ignored -> connection); assertEquals("public, max-age=300, no-store", result.getCacheControl()); } + + @Test + void binaryResponseRejectsOversizedContentLength() throws Exception { + Map> headers = Collections.singletonMap("Content-Length", + Collections.singletonList(String.valueOf(HttpUtil.MAX_AIA_RESPONSE_SIZE_IN_BYTES + 1))); + TestHttpURLConnection connection = new TestHttpURLConnection(200, new byte[] { 1 }, headers); + + HttpUtil.BinaryHttpResponse result + = HttpUtil.getAiaBytesWithMetadata("https://example.test/cert.crt", ignored -> connection); + + assertNull(result.getBody()); + assertTrue(connection.disconnected); + } + + @Test + void binaryResponseRejectsStreamThatExceedsMaximumSize() throws Exception { + byte[] body = new byte[HttpUtil.MAX_AIA_RESPONSE_SIZE_IN_BYTES + 1]; + TestHttpURLConnection connection = new TestHttpURLConnection(200, body, Collections.emptyMap()); + + HttpUtil.BinaryHttpResponse result + = HttpUtil.getAiaBytesWithMetadata("https://example.test/cert.crt", ignored -> connection); + + assertNull(result.getBody()); + assertTrue(connection.disconnected); + } + + @Test + void binaryResponseFollowsHttpToHttpsRedirect() throws Exception { + String sourceUrl = "http://example.test/cert.crt"; + String targetUrl = "https://example.test/cert.crt"; + Map> redirectHeaders + = Collections.singletonMap("Location", Collections.singletonList(targetUrl)); + TestHttpURLConnection redirect = new TestHttpURLConnection(sourceUrl, 302, new byte[0], redirectHeaders); + byte[] body = new byte[] { 1, 2, 3 }; + TestHttpURLConnection response = new TestHttpURLConnection(targetUrl, 200, body, Collections.emptyMap()); + + HttpUtil.BinaryHttpResponse result = HttpUtil.getAiaBytesWithMetadata(sourceUrl, url -> { + if (sourceUrl.equals(url)) { + return redirect; + } + if (targetUrl.equals(url)) { + return response; + } + throw new AssertionError("Unexpected redirect URL: " + url); + }); + + assertArrayEquals(body, result.getBody()); + assertFalse(redirect.getInstanceFollowRedirects()); + assertTrue(redirect.disconnected); + assertTrue(response.disconnected); + } + + @Test + void binaryResponsePreservesPathForQueryOnlyRedirect() throws Exception { + String sourceUrl = "https://example.test/certificates/issuer.crt?v=1"; + String targetUrl = "https://example.test/certificates/issuer.crt?v=2"; + Map> redirectHeaders + = Collections.singletonMap("Location", Collections.singletonList("?v=2")); + TestHttpURLConnection redirect = new TestHttpURLConnection(sourceUrl, 302, new byte[0], redirectHeaders); + byte[] body = new byte[] { 1, 2, 3 }; + TestHttpURLConnection response = new TestHttpURLConnection(targetUrl, 200, body, Collections.emptyMap()); + + HttpUtil.BinaryHttpResponse result = HttpUtil.getAiaBytesWithMetadata(sourceUrl, url -> { + if (sourceUrl.equals(url)) { + return redirect; + } + if (targetUrl.equals(url)) { + return response; + } + throw new AssertionError("Unexpected redirect URL: " + url); + }); + + assertArrayEquals(body, result.getBody()); + assertTrue(redirect.disconnected); + assertTrue(response.disconnected); + } + + private static final class TestHttpURLConnection extends HttpURLConnection { + private final int status; + private final byte[] body; + private final Map> headers; + private final ByteArrayOutputStream requestBody = new ByteArrayOutputStream(); + private boolean disconnected; + + private TestHttpURLConnection(int status, byte[] body, Map> headers) throws Exception { + this("https://example.test/cert.crt", status, body, headers); + } + + private TestHttpURLConnection(String url, int status, byte[] body, Map> headers) + throws Exception { + super(new URL(url)); + this.status = status; + this.body = body; + this.headers = headers; + } + + @Override + public int getResponseCode() { + return status; + } + + @Override + public InputStream getInputStream() { + return new ByteArrayInputStream(body); + } + + @Override + public OutputStream getOutputStream() { + return requestBody; + } + + @Override + public Map> getHeaderFields() { + return headers; + } + + @Override + public String getHeaderField(String name) { + for (Map.Entry> entry : headers.entrySet()) { + if (entry.getKey().equalsIgnoreCase(name) && !entry.getValue().isEmpty()) { + return entry.getValue().get(0); + } + } + return null; + } + + @Override + public void disconnect() { + disconnected = true; + } + + @Override + public boolean usingProxy() { + return false; + } + + @Override + public void connect() { + } + } } diff --git a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/JsonConverterUtilTest.java b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/JsonConverterUtilTest.java index d0e11918ed04..b981cac2bbc6 100644 --- a/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/JsonConverterUtilTest.java +++ b/sdk/keyvault/azure-security-keyvault-jca/src/test/java/com/azure/security/keyvault/jca/implementation/utils/JsonConverterUtilTest.java @@ -16,6 +16,7 @@ import java.util.logging.LogRecord; import java.util.logging.Logger; +import static org.junit.jupiter.api.Assertions.assertDoesNotThrow; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertNotNull; import static org.junit.jupiter.api.Assertions.assertTrue; @@ -33,9 +34,9 @@ public class JsonConverterUtilTest { * Test the {@link JsonConverterUtil#fromJson(ReadValueCallback, String)} method. */ @Test - public void testFromJson() throws IOException { - String string = "{ \"cer\": \"cer\" }"; - CertificateBundle bundle = JsonConverterUtil.fromJson(CertificateBundle::fromJson, string); + public void testFromJson() { + CertificateBundle bundle + = assertDoesNotThrow(() -> JsonConverterUtil.fromJson(CertificateBundle::fromJson, "{\"cer\":\"cer\"}")); assertNotNull(bundle); assertEquals("cer", bundle.getCer()); @@ -47,23 +48,18 @@ public void testFromJson() throws IOException { @Test public void testToJson() { CertificateBundle bundle = new CertificateBundle(); - bundle.setCer("value"); String string = JsonConverterUtil.toJson(bundle); - assertTrue(string.contains("cer")); + assertTrue(string.contains("\"cer\"")); assertTrue(string.contains("\"value\"")); } @Test void testFromJsonWithTokenResponseBody() { - AccessToken accessToken = null; - try { - accessToken = JsonConverterUtil.fromJson(AccessToken::fromJson, DUMMY_TOKEN_RESPONSE_BODY); - } catch (IOException e) { - throw new RuntimeException(e); - } + AccessToken accessToken + = assertDoesNotThrow(() -> JsonConverterUtil.fromJson(AccessToken::fromJson, DUMMY_TOKEN_RESPONSE_BODY)); assertNotNull(accessToken); assertEquals("test_access_token_value", accessToken.getAccessToken()); }